Team Ai
Apppublic

wkea/blockdiffusion-api

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