wkea/blockdiffusion-api
0
1from __future__ import annotations2 3import os4from dataclasses import dataclass5from typing import Any, Dict, Optional6 7import torch8from torch import nn9from contextlib import contextmanager10from torch.utils.data import DataLoader11 12from blockdiffusion.diffusion.gaussian_diffusion import GaussianDiffusion13from blockdiffusion.train.guidance import CFGuidanceModel14from blockdiffusion.utils.checkpoint import load_checkpoint, save_checkpoint15from blockdiffusion.utils.ema import EMA16from blockdiffusion.utils.image_io import save_tensor_grid17 18 19def _make_grad_scaler(enabled: bool, device_type: str):20 """21 兼容不同 PyTorch 版本的 AMP GradScaler:22 - 新版推荐 torch.amp.GradScaler(device, ...)23 - 旧版使用 torch.cuda.amp.GradScaler(...)24 """25 enabled = bool(enabled and device_type == "cuda")26 try:27 from torch.amp import GradScaler as _GS # PyTorch 2.x28 29 if not enabled:30 return _GS(enabled=False)31 # 有的版本是 GradScaler(device='cuda', ...),有的接受位置参数32 try:33 return _GS("cuda", enabled=True)34 except TypeError:35 return _GS(enabled=True)36 except Exception:37 from torch.cuda.amp import GradScaler as _GS # 兼容旧版本38 39 return _GS(enabled=enabled)40 41 42@contextmanager43def _autocast(device_type: str, enabled: bool):44 """45 兼容不同 PyTorch 版本的 autocast:46 - torch.amp.autocast(device_type=..., enabled=...)47 - torch.cuda.amp.autocast(enabled=...)48 """49 enabled = bool(enabled and device_type == "cuda")50 try:51 from torch.amp import autocast as _AC52 53 with _AC(device_type=device_type, enabled=enabled):54 yield55 except Exception:56 from torch.cuda.amp import autocast as _AC57 58 with _AC(enabled=enabled):59 yield60 61 62@dataclass63class TrainerConfig:64 """65 训练器配置(尽量给出稳定且适合 16×16 的默认值)。66 """67 68 # 数据69 image_size: int = 1670 channels: int = 371 batch_size: int = 51272 num_workers: int = 473 74 # 训练75 lr: float = 2e-476 weight_decay: float = 1e-477 total_steps: int = 200_00078 grad_clip: float = 1.079 amp: bool = True80 cond_drop_prob: float = 0.1 # 仅在有条件训练时生效(CFG 训练)81 grad_accum_steps: int = 1 # 梯度累积:提高有效 batch(质量更稳),也可用于控制显存82 lr_warmup_steps: int = 2000 # 学习率 warmup 步数(0 表示关闭)83 lr_min_ratio: float = 0.1 # 余弦退火最低 lr 比例(相对 lr)84 snr_gamma: float = 5.0 # SNR 加权 MSE 的 gamma(<=0 表示关闭)85 channels_last: bool = True # Conv 的内存布局优化,通常能略提速86 87 # EMA88 ema_decay: float = 0.999989 90 # 日志/保存91 log_every: int = 10092 save_every: int = 5_00093 sample_every: int = 5_00094 sample_batch: int = 6495 sample_use_ddim: bool = True96 sample_ddim_steps: int = 5097 sample_ddim_eta: float = 0.098 sample_cfg_scale: float = 3.099 100 out_dir: str = "outputs"101 run_name: str = "default"102 resume_path: Optional[str] = None103 104 105class Trainer:106 """107 训练器只负责:108 - dataloader109 - 优化器/AMP/EMA110 - checkpoint 与训练过程采样111 扩散公式与采样过程由 GaussianDiffusion 负责,保证解耦。112 """113 114 def __init__(115 self,116 model: nn.Module,117 diffusion: GaussianDiffusion,118 dataloader: DataLoader,119 device: torch.device,120 cfg: TrainerConfig,121 meta: Optional[dict] = None,122 sample_model_kwargs: Optional[Dict[str, Any]] = None,123 sample_uncond_model_kwargs: Optional[Dict[str, Any]] = None,124 ):125 self.model = model.to(device)126 self.diffusion = diffusion127 self.dataloader = dataloader128 self.device = device129 self.cfg = cfg130 self.meta = meta or {}131 self.sample_model_kwargs = sample_model_kwargs132 self.sample_uncond_model_kwargs = sample_uncond_model_kwargs133 134 # 可选:channels_last 提升卷积吞吐(对小图也常有帮助)135 if cfg.channels_last and device.type == "cuda":136 self.model = self.model.to(memory_format=torch.channels_last)137 138 self.optimizer = torch.optim.AdamW(self.model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)139 use_amp = bool(cfg.amp and device.type == "cuda")140 self.scaler = _make_grad_scaler(enabled=use_amp, device_type=device.type)141 self.ema = EMA(self.model, decay=cfg.ema_decay)142 143 self.step = 0144 self._prepare_dirs()145 146 if cfg.resume_path:147 self._resume(cfg.resume_path)148 149 # 学习率调度(warmup + cosine decay),提升稳定性/质量150 self.base_lr = float(cfg.lr)151 self.min_lr = float(cfg.lr) * float(cfg.lr_min_ratio)152 153 def _prepare_dirs(self) -> None:154 self.run_dir = os.path.join(self.cfg.out_dir, self.cfg.run_name)155 self.ckpt_dir = os.path.join(self.run_dir, "checkpoints")156 self.sample_dir = os.path.join(self.run_dir, "samples")157 os.makedirs(self.ckpt_dir, exist_ok=True)158 os.makedirs(self.sample_dir, exist_ok=True)159 160 def _resume(self, path: str) -> None:161 ckpt = load_checkpoint(path, map_location="cpu")162 self.model.load_state_dict(ckpt["model"], strict=True)163 if "ema" in ckpt:164 self.ema.load_state_dict(ckpt["ema"])165 # 续训:确保 EMA shadow 跟随当前训练设备166 self.ema.to(self.device)167 if "optimizer" in ckpt:168 self.optimizer.load_state_dict(ckpt["optimizer"])169 # 续训:optimizer state 里的 tensor 默认仍在 CPU,需要搬到当前设备170 for state in self.optimizer.state.values():171 for kk, vv in list(state.items()):172 if torch.is_tensor(vv):173 state[kk] = vv.to(self.device)174 if "scaler" in ckpt:175 self.scaler.load_state_dict(ckpt["scaler"])176 self.step = int(ckpt.get("step", 0))177 178 def _save(self) -> None:179 path = os.path.join(self.ckpt_dir, f"step_{self.step}.pt")180 payload = {181 "step": self.step,182 "meta": self.meta,183 "trainer_cfg": self.cfg.__dict__,184 "model": self.model.state_dict(),185 "ema": self.ema.state_dict(),186 "optimizer": self.optimizer.state_dict(),187 "scaler": self.scaler.state_dict(),188 }189 save_checkpoint(path, payload)190 191 @torch.no_grad()192 def _sample(self) -> None:193 # 采样时使用 EMA 参数,避免训练噪声导致的质量抖动194 tmp = {k: v.detach().clone() for k, v in self.model.state_dict().items()}195 self.ema.copy_to(self.model)196 self.model.eval()197 198 shape = (self.cfg.sample_batch, self.cfg.channels, self.cfg.image_size, self.cfg.image_size)199 sampler_model: nn.Module = self.model200 if self.sample_model_kwargs is not None and self.sample_uncond_model_kwargs is not None:201 sampler_model = CFGuidanceModel(202 base_model=self.model,203 cond_kwargs=self.sample_model_kwargs,204 uncond_kwargs=self.sample_uncond_model_kwargs,205 scale=self.cfg.sample_cfg_scale,206 )207 if self.cfg.sample_use_ddim:208 x = self.diffusion.ddim_sample_loop(209 sampler_model,210 shape=shape,211 device=self.device,212 steps=self.cfg.sample_ddim_steps,213 eta=self.cfg.sample_ddim_eta,214 model_kwargs=None,215 progress=False,216 )217 else:218 x = self.diffusion.p_sample_loop(sampler_model, shape=shape, device=self.device, model_kwargs=None, progress=False)219 220 # [-1,1] -> [0,1] 保存221 x = (x.clamp(-1, 1) + 1) * 0.5222 out_path = os.path.join(self.sample_dir, f"step_{self.step}.png")223 save_tensor_grid(x, out_path, nrow=int(self.cfg.sample_batch**0.5), padding=1)224 225 # 恢复训练参数226 self.model.load_state_dict(tmp, strict=True)227 self.model.train()228 229 def train(self) -> None:230 self.model.train()231 data_iter = iter(self.dataloader)232 233 while self.step < self.cfg.total_steps:234 # 梯度累积:每 grad_accum_steps 次 forward 才做一次 optimizer step235 accum_steps = max(1, int(self.cfg.grad_accum_steps))236 self.optimizer.zero_grad(set_to_none=True)237 238 total_loss = 0.0239 for micro in range(accum_steps):240 try:241 batch = next(data_iter)242 except StopIteration:243 data_iter = iter(self.dataloader)244 batch = next(data_iter)245 246 model_kwargs: Dict[str, Any] = {}247 if isinstance(batch, (tuple, list)) and len(batch) == 3:248 x0, tokens, mask = batch249 model_kwargs["tokens"] = tokens.to(self.device, non_blocking=True)250 model_kwargs["mask"] = mask.to(self.device, non_blocking=True)251 else:252 x0 = batch253 254 # channels_last:数据也转换同布局255 x0 = x0.to(self.device, non_blocking=True)256 if self.cfg.channels_last and self.device.type == "cuda":257 x0 = x0.contiguous(memory_format=torch.channels_last)258 259 b = x0.shape[0]260 t = torch.randint(0, self.diffusion.num_timesteps, (b,), device=self.device, dtype=torch.long)261 262 with _autocast(device_type=self.device.type, enabled=self.scaler.is_enabled()):263 loss = self.diffusion.training_losses(264 self.model,265 x0,266 t,267 model_kwargs=model_kwargs,268 cond_drop_prob=self.cfg.cond_drop_prob,269 snr_gamma=float(self.cfg.snr_gamma),270 )["loss"]271 # 累积时按比例缩放,保持有效学习率不变272 loss = loss / accum_steps273 274 self.scaler.scale(loss).backward()275 total_loss += float(loss.detach().item())276 277 # 学习率 warmup + cosine decay(基于全局 step)278 if self.cfg.lr_warmup_steps and self.cfg.lr_warmup_steps > 0 and self.step < self.cfg.lr_warmup_steps:279 lr = self.min_lr + (self.base_lr - self.min_lr) * (float(self.step + 1) / float(self.cfg.lr_warmup_steps))280 else:281 # cosine 从 warmup 结束后开始282 warm = max(0, int(self.cfg.lr_warmup_steps))283 prog = (self.step - warm) / max(1.0, float(self.cfg.total_steps - warm))284 prog = min(max(prog, 0.0), 1.0)285 import math286 287 lr = self.min_lr + 0.5 * (self.base_lr - self.min_lr) * (1.0 + math.cos(math.pi * prog))288 for pg in self.optimizer.param_groups:289 pg["lr"] = lr290 291 if self.cfg.grad_clip and self.cfg.grad_clip > 0:292 self.scaler.unscale_(self.optimizer)293 torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.cfg.grad_clip)294 295 self.scaler.step(self.optimizer)296 self.scaler.update()297 298 self.ema.update(self.model)299 self.step += 1300 301 if self.cfg.log_every > 0 and (self.step % self.cfg.log_every == 0):302 # total_loss 已按 accum_steps 平均后的值(每 micro 先 /accum)303 print(f"[step {self.step}] loss={total_loss:.6f} lr={lr:.2e}")304 305 if self.cfg.save_every > 0 and (self.step % self.cfg.save_every == 0):306 self._save()307 308 if self.cfg.sample_every > 0 and (self.step % self.cfg.sample_every == 0):309 self._sample()310 311 312 