wkea/blockdiffusion-api
0
1from __future__ import annotations2 3import argparse4import os5from typing import Optional6 7import torch8 9from blockdiffusion.diffusion.gaussian_diffusion import GaussianDiffusion10from blockdiffusion.diffusion.schedules import get_beta_schedule11from blockdiffusion.models.unet16 import UNet16, UNet16Config12from blockdiffusion.models.text_unet16 import TextCondConfig, TextCondUNet1613from blockdiffusion.text.simple_tokenizer import SimpleVocab14from blockdiffusion.train.guidance import CFGuidanceModel15from blockdiffusion.utils.checkpoint import load_checkpoint16from blockdiffusion.utils.ema import EMA17from blockdiffusion.utils.image_io import save_tensor_grid18from blockdiffusion.utils.seed import get_device, seed_everything19 20 21def parse_args() -> argparse.Namespace:22 p = argparse.ArgumentParser(description="BlockDiffusion 16×16 纹理采样入口")23 p.add_argument("--ckpt", type=str, required=True, help="训练生成的 checkpoint 路径(.pt)")24 p.add_argument("--out_dir", type=str, default="samples", help="输出目录")25 p.add_argument("--n", type=int, default=256, help="生成张数")26 p.add_argument("--batch", type=int, default=256, help="单次采样 batch(受显存影响)")27 28 p.add_argument("--use_ema", action="store_true", help="使用 EMA 权重采样(推荐)")29 p.add_argument("--use_ddim", action="store_true", help="使用 DDIM 快速采样")30 p.add_argument("--ddim_steps", type=int, default=50)31 p.add_argument("--ddim_eta", type=float, default=0.0)32 p.add_argument("--prompt", type=str, default=None, help="提示词(仅条件模型可用)")33 p.add_argument("--cfg_scale", type=float, default=3.0, help="CFG scale(仅条件模型可用)")34 35 p.add_argument("--seed", type=int, default=123)36 p.add_argument("--device", type=str, default=None)37 return p.parse_args()38 39 40@torch.no_grad()41def main() -> None:42 args = parse_args()43 seed_everything(args.seed, deterministic=False)44 device = get_device(args.device)45 46 ckpt = load_checkpoint(args.ckpt, map_location="cpu")47 meta = ckpt.get("meta", {})48 49 model_type = str(meta.get("model_type", "uncond"))50 unet_cfg_dict = meta.get("unet_cfg", None)51 if unet_cfg_dict is None:52 # 兼容:若旧 checkpoint 未保存配置,则使用默认 16×16 RGB 配置53 unet_cfg = UNet16Config()54 timesteps = 100055 beta_schedule = "cosine"56 image_size = 1657 channels = 358 model_type = "uncond"59 else:60 unet_cfg = UNet16Config(**unet_cfg_dict)61 diff_meta = meta.get("diffusion", {})62 timesteps = int(diff_meta.get("timesteps", 1000))63 beta_schedule = str(diff_meta.get("beta_schedule", "cosine"))64 data_meta = meta.get("data", {})65 image_size = int(data_meta.get("image_size", 16))66 channels = int(data_meta.get("channels", unet_cfg.in_channels))67 68 vocab = None69 text_cfg = None70 if model_type == "text_cond":71 tok_dict = meta.get("tokenizer", None)72 text_cfg_dict = meta.get("text_cfg", None)73 if tok_dict is None or text_cfg_dict is None:74 raise ValueError("该 checkpoint 标记为 text_cond,但缺少 tokenizer/text_cfg")75 vocab = SimpleVocab.from_dict(tok_dict)76 text_cfg = TextCondConfig(**text_cfg_dict)77 model = TextCondUNet16(unet_cfg=unet_cfg, vocab=vocab, text_cfg=text_cfg).to(device)78 else:79 model = UNet16(unet_cfg).to(device)80 model.load_state_dict(ckpt["model"], strict=True)81 82 if args.use_ema and ("ema" in ckpt):83 ema = EMA(model, decay=0.9999)84 ema.load_state_dict(ckpt["ema"])85 ema.copy_to(model)86 87 betas = get_beta_schedule(beta_schedule, timesteps).to(device)88 diffusion = GaussianDiffusion(betas)89 90 os.makedirs(args.out_dir, exist_ok=True)91 92 # 条件采样:如果提供 prompt 且模型支持,则启用 CFG93 sampler_model = model94 model_kwargs = None95 if args.prompt and model_type == "text_cond":96 if vocab is None or text_cfg is None:97 raise ValueError("内部错误:条件模型缺少 vocab/text_cfg")98 tok = torch.tensor([vocab.encode(args.prompt, max_len=text_cfg.max_text_len)], dtype=torch.long, device=device)99 m = (tok != vocab.pad_id)100 cond_kwargs = {"tokens": tok, "mask": m}101 un_tok = torch.full_like(tok, fill_value=vocab.pad_id)102 un_m = torch.zeros_like(m)103 uncond_kwargs = {"tokens": un_tok, "mask": un_m}104 sampler_model = CFGuidanceModel(model, cond_kwargs=cond_kwargs, uncond_kwargs=uncond_kwargs, scale=args.cfg_scale).to(device)105 model_kwargs = None106 elif args.prompt and model_type != "text_cond":107 raise ValueError("当前 checkpoint 是无条件模型,不支持 --prompt;请用带 captions 训练的模型")108 109 all_imgs = []110 remaining = args.n111 while remaining > 0:112 b = min(args.batch, remaining)113 shape = (b, channels, image_size, image_size)114 if args.use_ddim:115 x = diffusion.ddim_sample_loop(116 sampler_model,117 shape=shape,118 device=device,119 steps=args.ddim_steps,120 eta=args.ddim_eta,121 model_kwargs=model_kwargs,122 progress=False,123 )124 else:125 x = diffusion.p_sample_loop(sampler_model, shape=shape, device=device, model_kwargs=model_kwargs, progress=False)126 x = (x.clamp(-1, 1) + 1) * 0.5127 all_imgs.append(x.cpu())128 remaining -= b129 130 imgs = torch.cat(all_imgs, dim=0)131 out_path = os.path.join(args.out_dir, "samples.png")132 save_tensor_grid(imgs, out_path, nrow=int(max(1, args.n) ** 0.5), padding=1)133 134 135if __name__ == "__main__":136 main()137 138 139 