Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
sample_cpu.py195 linesDownload Raw Back to root
1from __future__ import annotations
2
3"""
4BlockDiffusion:CPU 推理/采样脚本(强制使用 CPU)
5
6设计目标:
7- 即使机器有 GPU,也只使用 CPU(避免 CUDA 初始化与显存占用)
8- 复用项目现有的 ckpt 元信息/模型构建/EMA/CFG/扩散采样逻辑
9
10注意:
11- CPU 采样会明显更慢,建议配合 --use_ddim 与较小的 --ddim_steps / --batch
12"""
13
14import argparse
15import os
16from typing import Optional
17
18# 必须在 import torch 之前设置,否则可能已经触发 CUDA 初始化
19os.environ.setdefault("CUDA_VISIBLE_DEVICES", "")
20
21import torch
22
23from blockdiffusion.diffusion.gaussian_diffusion import GaussianDiffusion
24from blockdiffusion.diffusion.schedules import get_beta_schedule
25from blockdiffusion.models.text_unet16 import TextCondConfig, TextCondUNet16
26from blockdiffusion.models.unet16 import UNet16, UNet16Config
27from blockdiffusion.text.simple_tokenizer import SimpleVocab
28from blockdiffusion.train.guidance import CFGuidanceModel
29from blockdiffusion.utils.checkpoint import load_checkpoint
30from blockdiffusion.utils.ema import EMA
31from blockdiffusion.utils.image_io import save_tensor_grid
32from blockdiffusion.utils.seed import seed_everything
33
34
35def parse_args() -> argparse.Namespace:
36    p = argparse.ArgumentParser(description="BlockDiffusion CPU 采样入口(16×16 纹理)")
37    p.add_argument("--ckpt", type=str, required=True, help="训练生成的 checkpoint 路径(.pt)")
38    p.add_argument("--out_dir", type=str, default="samples_cpu", help="输出目录")
39    p.add_argument("--out_name", type=str, default="samples.png", help="输出文件名(建议 .png)")
40
41    p.add_argument("--n", type=int, default=16, help="生成张数(CPU 建议较小)")
42    p.add_argument("--batch", type=int, default=16, help="单次采样 batch(CPU 建议较小)")
43    p.add_argument("--threads", type=int, default=0, help="torch CPU 线程数(0 表示不设置)")
44
45    p.add_argument("--use_ema", action="store_true", help="使用 EMA 权重采样(推荐)")
46    p.add_argument("--use_ddim", action="store_true", help="使用 DDIM 快速采样")
47    p.add_argument("--ddim_steps", type=int, default=50, help="DDIM steps(越小越快)")
48    p.add_argument("--ddim_eta", type=float, default=0.0, help="DDIM eta(0=确定性)")
49
50    p.add_argument("--prompt", type=str, default=None, help="提示词(仅 text_cond 模型可用)")
51    p.add_argument("--cfg_scale", type=float, default=3.0, help="CFG scale(仅 text_cond 模型可用)")
52
53    p.add_argument("--seed", type=int, default=123, help="随机种子")
54    return p.parse_args()
55
56
57def _set_cpu_threads(threads: int) -> None:
58    """
59    可选设置 CPU 线程数,便于在不同机器上控制推理占用。
60    """
61    threads = int(threads)
62    if threads <= 0:
63        return
64    try:
65        torch.set_num_threads(threads)
66    except Exception:
67        # 某些环境可能限制 set_num_threads;这里静默忽略即可
68        pass
69    try:
70        # inter-op 线程在部分版本/平台上可能不可设置
71        torch.set_num_interop_threads(threads)
72    except Exception:
73        pass
74
75
76@torch.no_grad()
77def main() -> None:
78    args = parse_args()
79    _set_cpu_threads(args.threads)
80    seed_everything(args.seed, deterministic=False)
81
82    # 强制 CPU
83    device = torch.device("cpu")
84
85    ckpt = load_checkpoint(args.ckpt, map_location="cpu")
86    meta = ckpt.get("meta", {}) or {}
87
88    model_type = str(meta.get("model_type", "uncond"))
89    unet_cfg_dict = meta.get("unet_cfg", None)
90    if unet_cfg_dict is None:
91        # 兼容:旧 checkpoint 未保存配置
92        unet_cfg = UNet16Config()
93        timesteps = 1000
94        beta_schedule = "cosine"
95        image_size = 16
96        channels = 3
97        model_type = "uncond"
98    else:
99        unet_cfg = UNet16Config(**unet_cfg_dict)
100        diff_meta = meta.get("diffusion", {}) or {}
101        timesteps = int(diff_meta.get("timesteps", 1000))
102        beta_schedule = str(diff_meta.get("beta_schedule", "cosine"))
103        data_meta = meta.get("data", {}) or {}
104        image_size = int(data_meta.get("image_size", 16))
105        channels = int(data_meta.get("channels", unet_cfg.in_channels))
106
107    vocab: Optional[SimpleVocab] = None
108    text_cfg: Optional[TextCondConfig] = None
109    if model_type == "text_cond":
110        tok_dict = meta.get("tokenizer", None)
111        text_cfg_dict = meta.get("text_cfg", None)
112        if tok_dict is None or text_cfg_dict is None:
113            raise ValueError("该 checkpoint 标记为 text_cond,但缺少 tokenizer/text_cfg")
114        vocab = SimpleVocab.from_dict(tok_dict)
115        text_cfg = TextCondConfig(**text_cfg_dict)
116        model = TextCondUNet16(unet_cfg=unet_cfg, vocab=vocab, text_cfg=text_cfg).to(device)
117    else:
118        model = UNet16(unet_cfg).to(device)
119
120    model.load_state_dict(ckpt["model"], strict=True)
121    model.eval()
122
123    if args.use_ema and ("ema" in ckpt):
124        ema = EMA(model, decay=0.9999)
125        ema.load_state_dict(ckpt["ema"])
126        ema.copy_to(model)
127        model.eval()
128
129    betas = get_beta_schedule(beta_schedule, timesteps).to(device)
130    diffusion = GaussianDiffusion(betas)
131
132    # 条件采样:如果提供 prompt 且模型支持,则启用 CFG
133    sampler_model = model
134    model_kwargs = None
135    if args.prompt and model_type == "text_cond":
136        if vocab is None or text_cfg is None:
137            raise ValueError("内部错误:条件模型缺少 vocab/text_cfg")
138        tok = torch.tensor([vocab.encode(args.prompt, max_len=text_cfg.max_text_len)], dtype=torch.long, device=device)
139        m = tok != vocab.pad_id
140        cond_kwargs = {"tokens": tok, "mask": m}
141
142        un_tok = torch.full_like(tok, fill_value=vocab.pad_id)
143        un_m = torch.zeros_like(m)
144        uncond_kwargs = {"tokens": un_tok, "mask": un_m}
145
146        sampler_model = CFGuidanceModel(
147            model,
148            cond_kwargs=cond_kwargs,
149            uncond_kwargs=uncond_kwargs,
150            scale=float(args.cfg_scale),
151        ).to(device)
152        model_kwargs = None
153    elif args.prompt and model_type != "text_cond":
154        raise ValueError("当前 checkpoint 是无条件模型,不支持 --prompt;请用带 captions 训练的模型")
155
156    os.makedirs(args.out_dir, exist_ok=True)
157    out_path = os.path.join(args.out_dir, args.out_name)
158
159    all_imgs = []
160    remaining = int(args.n)
161    while remaining > 0:
162        b = min(int(args.batch), remaining)
163        shape = (b, int(channels), int(image_size), int(image_size))
164        if args.use_ddim:
165            x = diffusion.ddim_sample_loop(
166                sampler_model,
167                shape=shape,
168                device=device,
169                steps=int(args.ddim_steps),
170                eta=float(args.ddim_eta),
171                model_kwargs=model_kwargs,
172                progress=True,
173            )
174        else:
175            x = diffusion.p_sample_loop(
176                sampler_model,
177                shape=shape,
178                device=device,
179                model_kwargs=model_kwargs,
180                progress=True,
181            )
182        x = (x.clamp(-1, 1) + 1) * 0.5
183        all_imgs.append(x.cpu())
184        remaining -= b
185
186    imgs = torch.cat(all_imgs, dim=0)
187    nrow = int(max(1, int(args.n)) ** 0.5)
188    save_tensor_grid(imgs, out_path, nrow=nrow, padding=1)
189
190
191if __name__ == "__main__":
192    main()
193
194
195