wkea/blockdiffusion-api
0
1from __future__ import annotations2 3import argparse4import dataclasses5import os6from typing import Optional7 8import torch9from torch.utils.data import DataLoader10 11from blockdiffusion.data.texture_folder import TextureFolderDataset12from blockdiffusion.data.captions import load_captions13from blockdiffusion.data.texture_caption_dataset import TextureCaptionDataset14from blockdiffusion.diffusion.gaussian_diffusion import GaussianDiffusion15from blockdiffusion.diffusion.schedules import get_beta_schedule16from blockdiffusion.models.unet16 import UNet16, UNet16Config17from blockdiffusion.models.text_unet16 import TextCondConfig, TextCondUNet1618from blockdiffusion.text.simple_tokenizer import SimpleVocab, VocabConfig19from blockdiffusion.train.trainer import Trainer, TrainerConfig20from blockdiffusion.utils.seed import get_device, seed_everything21 22 23def parse_args() -> argparse.Namespace:24 p = argparse.ArgumentParser(description="BlockDiffusion 16×16 纹理训练入口")25 26 # 数据27 p.add_argument("--data_dir", type=str, required=True, help="纹理图片目录(可含子目录)")28 p.add_argument("--image_size", type=int, default=16, help="纹理边长(默认 16)")29 p.add_argument("--channels", type=int, default=3, help="纹理通道数(默认 RGB=3)")30 p.add_argument("--captions", type=str, default=None, help="图文配对文件(.jsonl/.csv),提供则启用提示词条件训练")31 p.add_argument("--max_text_len", type=int, default=32, help="文本最大长度(中文按字、英文按词)")32 p.add_argument("--vocab_max", type=int, default=8000, help="词表最大大小")33 34 # 扩散35 p.add_argument("--timesteps", type=int, default=1000, help="扩散步数")36 p.add_argument("--beta_schedule", type=str, default="cosine", choices=["linear", "cosine"], help="噪声调度")37 38 # UNet39 p.add_argument("--base_channels", type=int, default=64, help="UNet 基础通道数")40 p.add_argument("--num_res_blocks", type=int, default=2, help="每个 level 的 ResBlock 数")41 p.add_argument("--dropout", type=float, default=0.0, help="ResBlock dropout(一般可为 0)")42 43 # 训练44 p.add_argument("--batch_size", type=int, default=512)45 # Windows 下多进程 DataLoader 可能更挑剔,默认给 0 更稳;需要更快可手动调大。46 p.add_argument("--num_workers", type=int, default=(0 if os.name == "nt" else 4))47 p.add_argument("--lr", type=float, default=2e-4)48 p.add_argument("--weight_decay", type=float, default=1e-4)49 p.add_argument("--total_steps", type=int, default=200_000)50 p.add_argument("--grad_clip", type=float, default=1.0)51 p.add_argument("--no_amp", action="store_true", help="关闭混合精度")52 p.add_argument("--ema_decay", type=float, default=0.9999)53 p.add_argument("--cond_drop_prob", type=float, default=0.1, help="CFG 训练:丢条件概率(0~1)")54 p.add_argument("--sample_prompt", type=str, default=None, help="训练中采样使用的提示词(仅条件训练时生效)")55 p.add_argument("--sample_cfg_scale", type=float, default=3.0, help="采样 CFG scale(仅条件训练时生效)")56 57 # 输出58 p.add_argument("--out_dir", type=str, default="outputs")59 p.add_argument("--run_name", type=str, default="default")60 p.add_argument("--resume", type=str, default=None, help="从 checkpoint 恢复训练")61 62 # 其他63 p.add_argument("--seed", type=int, default=42)64 p.add_argument("--device", type=str, default=None, help='如 "cuda" / "cuda:0" / "cpu"')65 66 return p.parse_args()67 68 69def main() -> None:70 args = parse_args()71 72 seed_everything(args.seed, deterministic=False)73 device = get_device(args.device)74 75 # 性能优化:适用于 RTX 40 系列(不影响数值正确性)76 if hasattr(torch, "set_float32_matmul_precision"):77 torch.set_float32_matmul_precision("high")78 if device.type == "cuda":79 torch.backends.cuda.matmul.allow_tf32 = True80 torch.backends.cudnn.allow_tf32 = True81 82 sample_model_kwargs = None83 sample_uncond_model_kwargs = None84 85 # 选择数据集:无条件/条件(captions)86 if args.captions:87 cap_items = load_captions(args.captions, data_root=args.data_dir)88 # 可选:强制加入透明/不透明 token,方便学习 alpha 与提示词关联(需你在 captions 里实际使用这些词)89 vocab = SimpleVocab.build(90 (it.text for it in cap_items),91 VocabConfig(max_vocab=args.vocab_max, min_freq=1),92 extra_tokens=["transparent", "opaque"],93 )94 dataset = TextureCaptionDataset(95 items=cap_items,96 vocab=vocab,97 image_size=args.image_size,98 channels=args.channels,99 max_text_len=args.max_text_len,100 random_flip=True,101 )102 else:103 vocab = None104 dataset = TextureFolderDataset(105 root_dir=args.data_dir,106 image_size=args.image_size,107 channels=args.channels,108 random_flip=True,109 )110 dataloader = DataLoader(111 dataset,112 batch_size=args.batch_size,113 shuffle=True,114 num_workers=args.num_workers,115 pin_memory=(device.type == "cuda"),116 drop_last=True,117 )118 119 unet_cfg = UNet16Config(120 in_channels=args.channels,121 out_channels=args.channels,122 base_channels=args.base_channels,123 channel_mults=(1, 2, 2), # 16->8->4124 num_res_blocks=args.num_res_blocks,125 attn_resolutions=(8, 4),126 dropout=args.dropout,127 time_emb_dim=256,128 cond_dim=256,129 )130 model_type = "uncond"131 text_cfg = None132 if args.captions:133 # 条件模型:TextEncoder 输出维度与 time_emb_dim 对齐134 text_cfg = TextCondConfig(max_text_len=args.max_text_len, emb_dim=128, dropout=0.0)135 model = TextCondUNet16(unet_cfg=unet_cfg, vocab=vocab, text_cfg=text_cfg)136 model_type = "text_cond"137 138 # 训练中采样(可选):提供 sample_prompt 才启用 CFG 采样,否则走无条件采样139 if args.sample_prompt:140 tok = torch.tensor([vocab.encode(args.sample_prompt, max_len=args.max_text_len)], dtype=torch.long)141 m = (tok != vocab.pad_id)142 sample_model_kwargs = {"tokens": tok.to(device), "mask": m.to(device)}143 # 无条件分支:mask 全 0(tokens 可以随意,统一用 pad)144 un_tok = torch.full_like(tok, fill_value=vocab.pad_id)145 un_m = torch.zeros_like(m)146 sample_uncond_model_kwargs = {"tokens": un_tok.to(device), "mask": un_m.to(device)}147 else:148 model = UNet16(unet_cfg)149 150 betas = get_beta_schedule(args.beta_schedule, args.timesteps).to(device)151 diffusion = GaussianDiffusion(betas)152 153 trainer_cfg = TrainerConfig(154 image_size=args.image_size,155 channels=args.channels,156 batch_size=args.batch_size,157 num_workers=args.num_workers,158 lr=args.lr,159 weight_decay=args.weight_decay,160 total_steps=args.total_steps,161 grad_clip=args.grad_clip,162 amp=not args.no_amp,163 ema_decay=args.ema_decay,164 cond_drop_prob=args.cond_drop_prob,165 sample_cfg_scale=args.sample_cfg_scale,166 out_dir=args.out_dir,167 run_name=args.run_name,168 resume_path=args.resume,169 )170 171 meta = {172 "model_type": model_type,173 "unet_cfg": dataclasses.asdict(unet_cfg),174 "diffusion": {"timesteps": args.timesteps, "beta_schedule": args.beta_schedule},175 "data": {"image_size": args.image_size, "channels": args.channels},176 }177 if args.captions and vocab is not None and text_cfg is not None:178 meta["tokenizer"] = vocab.to_dict()179 meta["text_cfg"] = dataclasses.asdict(text_cfg)180 181 trainer = Trainer(182 model=model,183 diffusion=diffusion,184 dataloader=dataloader,185 device=device,186 cfg=trainer_cfg,187 meta=meta,188 sample_model_kwargs=sample_model_kwargs,189 sample_uncond_model_kwargs=sample_uncond_model_kwargs,190 )191 trainer.train()192 193 194if __name__ == "__main__":195 main()196 197 198 