Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
train_oneclick.py312 linesDownload Raw Back to root
1from __future__ import annotations
2
3"""
4一键训练脚本(无命令行参数):
5 - 直接修改本文件顶部常量即可
6 - 数据目录:DATA_DIR
7 - 提示词来自文件名(HASH-xxx_yyy.png => "xxx yyy")
8
9运行方式(推荐 Python 3.10):
10 - py -3.10 train_oneclick.py
11"""
12
13import dataclasses
14import os
15
16import torch
17from torch.utils.data import DataLoader
18
19from blockdiffusion.data.filename_captions import build_caption_items_from_filenames
20from blockdiffusion.data.texture_caption_dataset import TextureCaptionDataset
21from blockdiffusion.diffusion.gaussian_diffusion import GaussianDiffusion
22from blockdiffusion.diffusion.schedules import get_beta_schedule
23from blockdiffusion.models.text_unet16 import TextCondConfig, TextCondUNet16
24from blockdiffusion.models.unet16 import UNet16Config
25from blockdiffusion.text.simple_tokenizer import SimpleVocab, VocabConfig
26from blockdiffusion.train.trainer import Trainer, TrainerConfig
27from blockdiffusion.utils.checkpoint import load_checkpoint
28from blockdiffusion.utils.seed import get_device, seed_everything
29
30# =========================
31# 可修改常量(按需改这里)
32# =========================
33
34# 数据目录(你提供的路径)
35DATA_DIR = r"G:\PycharmProjects\modrinth_downloads\Training_v2"
36
37# 纹理尺寸(本项目目标 16)
38IMAGE_SIZE = 16
39
40# 通道数:3=RGB,4=RGBA(需要保留透明就用 4)
41CHANNELS = 4
42
43# 扩散步数与调度
44TIMESTEPS = 1000
45BETA_SCHEDULE = "cosine"  # "linear" / "cosine"
46
47# UNet 规模(16×16 不要太深)
48BASE_CHANNELS = 64
49NUM_RES_BLOCKS = 2
50DROPOUT = 0.0
51
52# 文本(提示词)配置 
53MAX_TEXT_LEN = 32
54VOCAB_MAX = 8000
55COND_DROP_PROB = 0.1  # CFG 训练丢条件概率(0~1)
56PRETOKENIZE = False  # 大数据集建议 False:启动更快、占用更小
57
58# 透明度提示词(可选):
59# - 仅在 CHANNELS=4 且“新训练(非续训)”时建议开启
60# - 作用:自动根据图片 alpha 是否含透明像素,为文本追加 "transparent"/"opaque"
61#   这样模型才能学到“透明度与提示词相关”的监督信号
62APPEND_ALPHA_TAG = False
63ALPHA_TAG_TRANSPARENT = "transparent"
64ALPHA_TAG_OPAQUE = "opaque"
65
66# 训练配置(4070Ti 12G:RGBA 建议先从 128/256 试起)
67BATCH_SIZE = 128
68# 数据加载并行:Windows 若不稳定可改回 0;稳定时建议 4~8
69NUM_WORKERS = 4
70PREFETCH_FACTOR = 4  # NUM_WORKERS>0 时生效
71PERSISTENT_WORKERS = True  # NUM_WORKERS>0 时生效
72LR = 2e-4
73WEIGHT_DECAY = 1e-4
74TOTAL_STEPS = 700_000
75GRAD_CLIP = 1.0
76AMP = True
77EMA_DECAY = 0.9999
78GRAD_ACCUM_STEPS = 1  # 不想OOM但想更大有效batch:例如设为 2/4
79LR_WARMUP_STEPS = 2000
80LR_MIN_RATIO = 0.1
81SNR_GAMMA = 5.0  # <=0 关闭;5 通常比较稳
82CHANNELS_LAST = True
83LOG_EVERY = 10
84SAVE_EVERY = 5000
85SAMPLE_EVERY = 5000
86
87# 输出
88OUT_DIR = "outputs"
89RUN_NAME = "oneclick_filename_caption"
90# 断点续训:
91# - 若你明确指定某个 checkpoint,就填 RESUME_PATH
92# - 若 RESUME_PATH=None 且 AUTO_RESUME=True,会自动扫描 outputs/<RUN_NAME>/checkpoints 下最新 step_*.pt 继续训练
93RESUME_PATH = None  # 例如 r"outputs/oneclick_filename_caption/checkpoints/step_5000.pt"
94AUTO_RESUME = True
95
96# 训练中采样(可选):用这个提示词做固定可视化
97SAMPLE_PROMPT = "block"  # 改成你想看的关键词;不想要可设为 None
98SAMPLE_CFG_SCALE = 3.0
99
100# 随机种子
101SEED = 42
102DEVICE = None  # 例如 "cuda:0";None 表示自动
103
104
105def main() -> None:
106    seed_everything(SEED, deterministic=False)
107    device = get_device(DEVICE)
108    print(f"[init] device={device}  channels={CHANNELS}  image_size={IMAGE_SIZE}")
109
110    # RTX 40 系列性能优化(不影响数值正确性)
111    if hasattr(torch, "set_float32_matmul_precision"):
112        torch.set_float32_matmul_precision("high")
113    if device.type == "cuda":
114        torch.backends.cuda.matmul.allow_tf32 = True
115        torch.backends.cudnn.allow_tf32 = True
116
117    # =====
118    # 续训注意:
119    # - 如果设置了 RESUME_PATH,为保证词表/模型结构不变,本脚本会优先从 checkpoint 读取并复用配置;
120    # - 新增数据会自动被扫描并参与训练;新出现的词如果不在旧词表里,会被编码为 <unk>(可继续训练,但不会学到新词语义)。
121    # =====
122
123    def _auto_find_latest_ckpt() -> str | None:
124        ckpt_dir = os.path.join(OUT_DIR, RUN_NAME, "checkpoints")
125        if not os.path.isdir(ckpt_dir):
126            return None
127        best_step = -1
128        best_path = None
129        for fn in os.listdir(ckpt_dir):
130            if not (fn.startswith("step_") and fn.endswith(".pt")):
131                continue
132            # 形如 step_5000.pt
133            core = fn[len("step_") : -len(".pt")]
134            try:
135                step = int(core)
136            except Exception:
137                continue
138            if step > best_step:
139                best_step = step
140                best_path = os.path.join(ckpt_dir, fn)
141        return best_path
142
143    effective_resume_path = RESUME_PATH
144    if (effective_resume_path is None) and AUTO_RESUME:
145        effective_resume_path = _auto_find_latest_ckpt()
146
147    resume_meta = None
148    if effective_resume_path:
149        print(f"[resume] 使用 checkpoint: {effective_resume_path}")
150        ckpt = load_checkpoint(effective_resume_path, map_location="cpu")
151        resume_meta = ckpt.get("meta", {}) or {}
152
153    # 从 checkpoint 复用关键配置(避免“重建词表”导致继续训练变差/不可用)
154    if resume_meta:
155        data_meta = resume_meta.get("data", {}) or {}
156        diff_meta = resume_meta.get("diffusion", {}) or {}
157        unet_cfg_dict = resume_meta.get("unet_cfg", None)
158        text_cfg_dict = resume_meta.get("text_cfg", None)
159        tok_dict = resume_meta.get("tokenizer", None)
160
161        cfg_image_size = int(data_meta.get("image_size", IMAGE_SIZE))
162        cfg_channels = int(data_meta.get("channels", CHANNELS))
163        cfg_timesteps = int(diff_meta.get("timesteps", TIMESTEPS))
164        cfg_beta_schedule = str(diff_meta.get("beta_schedule", BETA_SCHEDULE))
165
166        if unet_cfg_dict is None or text_cfg_dict is None or tok_dict is None:
167            raise ValueError("RESUME_PATH 指向的 checkpoint 缺少 tokenizer/text_cfg/unet_cfg,无法安全续训")
168        vocab = SimpleVocab.from_dict(tok_dict)
169        text_cfg = TextCondConfig(**text_cfg_dict)
170        unet_cfg = UNet16Config(**unet_cfg_dict)
171    else:
172        cfg_image_size = IMAGE_SIZE
173        cfg_channels = CHANNELS
174        cfg_timesteps = TIMESTEPS
175        cfg_beta_schedule = BETA_SCHEDULE
176        text_cfg = TextCondConfig(max_text_len=MAX_TEXT_LEN, emb_dim=128, dropout=0.0)
177        unet_cfg = UNet16Config(
178            in_channels=cfg_channels,
179            out_channels=cfg_channels,
180            base_channels=BASE_CHANNELS,
181            channel_mults=(1, 2, 2),  # 16->8->4
182            num_res_blocks=NUM_RES_BLOCKS,
183            attn_resolutions=(8, 4),
184            dropout=DROPOUT,
185            time_emb_dim=256,
186            cond_dim=256,
187        )
188
189        # 非续训:词表会基于当前数据构建
190        vocab = None
191
192    # 扫描目录(只扫描一次):用于训练集;非续训时也用于构建词表
193    print("[data] 扫描图片并从文件名提取提示词...")
194    cap_items = build_caption_items_from_filenames(DATA_DIR, verbose=True)
195    if vocab is None:
196        # 新训练:可强制加入额外 token,避免采样时提示词变成 <unk> 而“看起来不遵从”
197        extra_tokens = None
198        append_alpha_tag = bool(APPEND_ALPHA_TAG and (cfg_channels == 4))
199        if append_alpha_tag:
200            extra_tokens = [ALPHA_TAG_TRANSPARENT, ALPHA_TAG_OPAQUE]
201        vocab = SimpleVocab.build(
202            (it.text for it in cap_items),
203            VocabConfig(max_vocab=VOCAB_MAX, min_freq=1),
204            extra_tokens=extra_tokens,
205        )
206    else:
207        # 续训:vocab 固定,不能改变;否则 embedding 尺寸不匹配
208        append_alpha_tag = False
209
210    dataset = TextureCaptionDataset(
211        items=cap_items,
212        vocab=vocab,
213        image_size=cfg_image_size,
214        channels=cfg_channels,
215        max_text_len=text_cfg.max_text_len,
216        random_flip=True,
217        pretokenize=PRETOKENIZE,
218        append_alpha_tag=append_alpha_tag,
219        alpha_tag_transparent=ALPHA_TAG_TRANSPARENT,
220        alpha_tag_opaque=ALPHA_TAG_OPAQUE,
221        verbose=True,
222    )
223    print(f"[data] dataset_size={len(dataset)}  vocab_size={vocab.size}  pretokenize={PRETOKENIZE}")
224    dl_kwargs = dict(
225        batch_size=BATCH_SIZE,
226        shuffle=True,
227        num_workers=NUM_WORKERS,
228        pin_memory=(device.type == "cuda"),
229        drop_last=True,
230    )
231    if NUM_WORKERS > 0:
232        dl_kwargs["prefetch_factor"] = int(PREFETCH_FACTOR)
233        dl_kwargs["persistent_workers"] = bool(PERSISTENT_WORKERS)
234    dataloader = DataLoader(dataset, **dl_kwargs)
235    print(f"[data] dataloader_workers={NUM_WORKERS}")
236
237    # 2) 模型(文本条件)
238    model = TextCondUNet16(unet_cfg=unet_cfg, vocab=vocab, text_cfg=text_cfg)
239
240    # 3) Diffusion
241    betas = get_beta_schedule(cfg_beta_schedule, cfg_timesteps).to(device)
242    diffusion = GaussianDiffusion(betas)
243
244    # 4) 训练配置
245    trainer_cfg = TrainerConfig(
246        image_size=cfg_image_size,
247        channels=cfg_channels,
248        batch_size=BATCH_SIZE,
249        num_workers=NUM_WORKERS,
250        lr=LR,
251        weight_decay=WEIGHT_DECAY,
252        total_steps=TOTAL_STEPS,
253        grad_clip=GRAD_CLIP,
254        amp=AMP,
255        ema_decay=EMA_DECAY,
256        cond_drop_prob=COND_DROP_PROB,
257        grad_accum_steps=GRAD_ACCUM_STEPS,
258        lr_warmup_steps=LR_WARMUP_STEPS,
259        lr_min_ratio=LR_MIN_RATIO,
260        snr_gamma=SNR_GAMMA,
261        channels_last=CHANNELS_LAST,
262        log_every=LOG_EVERY,
263        save_every=SAVE_EVERY,
264        sample_every=SAMPLE_EVERY,
265        sample_cfg_scale=SAMPLE_CFG_SCALE,
266        out_dir=OUT_DIR,
267        run_name=RUN_NAME,
268        resume_path=effective_resume_path,
269    )
270
271    # 训练中采样 CFG(可选)
272    sample_model_kwargs = None
273    sample_uncond_model_kwargs = None
274    if SAMPLE_PROMPT:
275        tok = torch.tensor([vocab.encode(SAMPLE_PROMPT, max_len=text_cfg.max_text_len)], dtype=torch.long)
276        m = (tok != vocab.pad_id)
277        sample_model_kwargs = {"tokens": tok.to(device), "mask": m.to(device)}
278        un_tok = torch.full_like(tok, fill_value=vocab.pad_id)
279        un_m = torch.zeros_like(m)
280        sample_uncond_model_kwargs = {"tokens": un_tok.to(device), "mask": un_m.to(device)}
281
282    meta = {
283        "model_type": "text_cond",
284        "unet_cfg": dataclasses.asdict(unet_cfg),
285        "diffusion": {"timesteps": cfg_timesteps, "beta_schedule": cfg_beta_schedule},
286        "data": {"image_size": cfg_image_size, "channels": cfg_channels},
287        "tokenizer": vocab.to_dict(),
288        "text_cfg": dataclasses.asdict(text_cfg),
289        # 额外信息:方便回溯
290        "data_dir": os.path.abspath(DATA_DIR),
291        "caption_source": "filename",
292    }
293
294    trainer = Trainer(
295        model=model,
296        diffusion=diffusion,
297        dataloader=dataloader,
298        device=device,
299        cfg=trainer_cfg,
300        meta=meta,
301        sample_model_kwargs=sample_model_kwargs,
302        sample_uncond_model_kwargs=sample_uncond_model_kwargs,
303    )
304    print("[train] start...")
305    trainer.train()
306
307
308if __name__ == "__main__":
309    main()
310
311
312