wkea/blockdiffusion-api
0
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 