Team Ai
Apppublic

wkea/blockdiffusion-api

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