wkea/blockdiffusion-api
0
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 