Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
sample.py139 linesDownload Raw Back to root
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