#!/usr/bin/env python """S1: LoRA 像素域监督适配(单卡)。 - 数据: synthetic(在线 Real-ESRGAN 退化) + real pairs(manifest) 混合(--real_prob) - 模型: 官方 Net+halfDecoder 全链(输入 LR 128 -> 输出 RGB 512) - 训练: 仅 LoRA 参数(手工注入, rank 可设), bf16, grad accum, save net/full state 用法示例见 scripts/run_stage1.sh """ import argparse, json, math, os, random, sys, time, copy from pathlib import Path REPO = Path(__file__).resolve().parents[1] sys.path.insert(0, str(REPO)) sys.path.insert(0, str(REPO / "src")) sys.path.insert(0, str(REPO / "official")) import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.data import Dataset, DataLoader from omegaconf import OmegaConf from common import (ensure_official, load_diffusers_sd, load_pruned_decoder, assemble_full_student, inject_lora, lora_params, count_params, build_net, is_finite, check_tensor, clip_and_check_grads, EMA, preview_grid) ensure_official() from dataset import RealESRGANDataset, RealESRGANDegrader # official # --------------------------------------------------------------------------- # 小波高频 mask(纯 torch 单层 Haar,无额外依赖) # --------------------------------------------------------------------------- def haar_decomp(x): """x: [B,C,H,W]; 返回 dict(ll,lh,hl,hh) 每个 [B,C,H/2,W/2]""" B, C, H, W = x.shape H2, W2 = H // 2, W // 2 if H % 2 or W % 2: x = F.pad(x, (0, W % 2, 0, H % 2)) a = x.view(B, C, H2, 2, W2, 2) ll = (a[:, :, :, 0, :, 0] + a[:, :, :, 0, :, 1] + a[:, :, :, 1, :, 0] + a[:, :, :, 1, :, 1]) / 4 lh = (a[:, :, :, 0, :, 0] - a[:, :, :, 0, :, 1] + a[:, :, :, 1, :, 0] - a[:, :, :, 1, :, 1]) / 4 hl = (a[:, :, :, 0, :, 0] + a[:, :, :, 0, :, 1] - a[:, :, :, 1, :, 0] - a[:, :, :, 1, :, 1]) / 4 hh = (a[:, :, :, 0, :, 0] - a[:, :, :, 0, :, 1] - a[:, :, :, 1, :, 0] + a[:, :, :, 1, :, 1]) / 4 return ll, lh, hl, hh def highfreq_mask(x): with torch.no_grad(): _, lh, hl, hh = haar_decomp(x.detach().float()) e = (lh ** 2 + hl ** 2 + hh ** 2).sqrt() m = e / (e.flatten(2).mean(dim=2, keepdim=True) + 1e-6).unsqueeze(-1) m = F.interpolate(m, size=x.shape[-2:], mode="bilinear", align_corners=False) return m # --------------------------------------------------------------------------- # Real pairs dataset: manifest {"real_pairs":[{lr,hr}]} # --------------------------------------------------------------------------- class RealPairDataset(Dataset): """?? x4 ?(RealSR/DRealSR): ?? HR 512 crop + ??? LR 128 crop? ??? LQ~1K??/GT?4K ? 4x ??????????? 128 LR -> 512 HR? manifest: {"real_pairs":[{lr,hr}]}?lr/hr ??? 4x ??lr ???hr/4?? """ def __init__(self, manifest_path, patch=512, scale=4, seed=0): with open(manifest_path, encoding="utf-8") as fh: m = json.load(fh) self.pairs = m.get("real_pairs", []) self.patch = patch self.scale = scale self.lr_patch = patch // scale self.rng = random.Random(seed) self._cache = {} def __len__(self): return max(1, len(self.pairs) * 40) def __getitem__(self, idx): from PIL import Image from torchvision import transforms p = self.pairs[idx % len(self.pairs)] key = p["hr"] if key not in self._cache: hr = Image.open(p["hr"]).convert("RGB") lr = Image.open(p["lr"]).convert("RGB") self._cache[key] = (lr, hr) lr, hr = self._cache[key] w, h = hr.size if w < self.patch or h < self.patch: raise RuntimeError(f"HR ?? patch: {p['hr']} {hr.size}") x = self.rng.randint(0, w - self.patch) y = self.rng.randint(0, h - self.patch) hr_c = hr.crop((x, y, x + self.patch, y + self.patch)) # LR ??????: ??? lr/hr ????, ????? 128 sc_w, sc_h = lr.width / w, lr.height / h lx0, ly0 = int(x * sc_w), int(y * sc_h) lx1, ly1 = int((x + self.patch) * sc_w), int((y + self.patch) * sc_h) lx1 = min(lx1, lr.width); ly1 = min(ly1, lr.height) lr_c = lr.crop((lx0, ly0, lx1, ly1)).resize((self.lr_patch, self.lr_patch), Image.BICUBIC) to_t = transforms.ToTensor() lr_t = to_t(lr_c) * 2 - 1 hr_t = to_t(hr_c) * 2 - 1 if self.rng.random() < 0.5: lr_t = torch.flip(lr_t, dims=[2]); hr_t = torch.flip(hr_t, dims=[2]) return lr_t, hr_t # --------------------------------------------------------------------------- def main(): ap = argparse.ArgumentParser() ap.add_argument("--config", default="configs/config_s1_lora.yml") ap.add_argument("--manifest", default="data/manifest_train.json", help="含 real_pairs 的训练清单") ap.add_argument("--real_prob", type=float, default=0.35) ap.add_argument("--model_id", default="models/stable-diffusion-2-1-base") ap.add_argument("--half_decoder", default="weight/pretrained/halfDecoder.ckpt") ap.add_argument("--init_net", default="weight/net_params_200.pkl", help="官方学生权重(Net state)") ap.add_argument("--out", default="weight/s1") ap.add_argument("--log_dir", default="logs/s1") ap.add_argument("--steps", type=int, default=20000) ap.add_argument("--batch_size", type=int, default=8) ap.add_argument("--grad_accum", type=int, default=2) ap.add_argument("--lr", type=float, default=5e-5) ap.add_argument("--lora_rank", type=int, default=64) ap.add_argument("--lora_alpha", type=float, default=1.0) ap.add_argument("--save_every", type=int, default=2000) ap.add_argument("--w_l1", type=float, default=1.0) ap.add_argument("--w_lpips", type=float, default=1.0) ap.add_argument("--w_dists", type=float, default=0.3) ap.add_argument("--w_wave", type=float, default=0.5) ap.add_argument("--w_color", type=float, default=0.2) ap.add_argument("--bf16", action="store_true", default=True) ap.add_argument("--no_bf16", dest="bf16", action="store_false") ap.add_argument("--seed", type=int, default=123) ap.add_argument("--num_workers", type=int, default=8) ap.add_argument("--clip_grad", type=float, default=1.0, help="??????; 0=??") ap.add_argument("--ema_decay", type=float, default=0.999, help="EMA ??; 0=??") ap.add_argument("--vis_every", type=int, default=500, help="? N ?? LR/HR/?????") args = ap.parse_args() random.seed(args.seed); torch.manual_seed(args.seed) device = "cuda" if torch.cuda.is_available() else "cpu" cfg = OmegaConf.load(args.config) os.makedirs(args.out, exist_ok=True); os.makedirs(args.log_dir, exist_ok=True) log_path = os.path.join(args.log_dir, "train_lora.log") logf = open(log_path, "a", encoding="utf-8") def log(msg): print(msg, flush=True); logf.write(msg + "\n"); logf.flush() # ---- data ---- syn_ds = RealESRGANDataset(cfg, args.batch_size) syn_dl = DataLoader(syn_ds, batch_size=args.batch_size, num_workers=args.num_workers, shuffle=True) degrader = RealESRGANDegrader(cfg, device) real_ds = RealPairDataset(args.manifest) if args.real_prob > 0 else None real_dl = DataLoader(real_ds, batch_size=args.batch_size, num_workers=args.num_workers, shuffle=True) if real_ds else None # ---- model ---- dtype = torch.bfloat16 if (args.bf16 and device == "cuda") else torch.float32 vae, unet, text_encoder, tokenizer = load_diffusers_sd(args.model_id, dtype=torch.float32, device="cpu") del text_encoder, tokenizer, vae decoder = load_pruned_decoder(args.half_decoder, device="cpu", dtype=torch.float32) full = assemble_full_student(unet, decoder, net_weights=args.init_net, device=device, dtype=dtype) inject_lora(full, rank=args.lora_rank, alpha=args.lora_alpha) params = list(lora_params(full)) log(f"trainable params: {count_params(full, only_trainable=True)/1e6:.2f}M") optimizer = torch.optim.AdamW(params, lr=args.lr) sched = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.steps) scaler = torch.amp.GradScaler("cuda", enabled=(dtype == torch.float16)) if (device == "cuda" and dtype == torch.float16) else None # lpips / dists 可选 lpips_fn = None if args.w_lpips > 0: try: import lpips lpips_fn = lpips.LPIPS(net="alex").to(device).eval() for p in lpips_fn.parameters(): p.requires_grad_(False) except Exception as e: log(f"[warn] lpips 不可用: {e}; 该损失置 0") dists_fn = None if args.w_dists > 0: try: import pyiqa dists_fn = pyiqa.create_metric("dists", device=device) for p in dists_fn.parameters(): p.requires_grad_(False) except Exception as e: log(f"[warn] pyiqa dists 不可用: {e}; 该损失置 0") # ---- train loop (????/????/EMA/???) ---- syn_iter = iter(syn_dl); real_iter = iter(real_dl) if real_dl else None step = 0; skip_streak = 0 ema = EMA(params, args.ema_decay) if args.ema_decay > 0 else None name_of = {id(p): n for n, p in full.named_parameters() if p.requires_grad} full.train() optimizer.zero_grad(set_to_none=True) while step < args.steps: use_real = real_dl is not None and random.random() < args.real_prob try: if use_real: lr_t, hr_t = next(real_iter) else: batch = next(syn_iter) lr_t, hr_t = degrader.degrade(batch) except StopIteration: syn_iter = iter(syn_dl) real_iter = iter(real_dl) if real_dl else None continue lr_t, hr_t = lr_t.to(device), hr_t.to(device) if check_tensor(lr_t, "lr", log) or check_tensor(hr_t, "hr", log): skip_streak += 1 if skip_streak > 20: log("[anomaly] too many bad batches, abort"); break continue if dtype == torch.bfloat16: with torch.autocast("cuda", dtype=torch.bfloat16): out = full(lr_t) loss, items = _losses(out.float(), hr_t.float(), full, args, lpips_fn, dists_fn) else: out = full(lr_t) loss, items = _losses(out, hr_t, full, args, lpips_fn, dists_fn) if check_tensor(out, "output", log): skip_streak += 1 if skip_streak > 20: log("[anomaly] too many bad outputs, abort"); break optimizer.zero_grad(set_to_none=True) continue if not is_finite(loss): log(f"[anomaly] loss NaN/Inf at step {step+1}; skip step") skip_streak += 1 optimizer.zero_grad(set_to_none=True) if skip_streak > 20: log("[anomaly] too many bad losses, abort"); break continue skip_streak = 0 (loss / args.grad_accum).backward() if (step + 1) % args.grad_accum == 0: if clip_and_check_grads(params, args.clip_grad, log): optimizer.zero_grad(set_to_none=True) else: if scaler is not None: scaler.step(optimizer); scaler.update() else: optimizer.step() optimizer.zero_grad(set_to_none=True) sched.step() if ema is not None: ema.update(params) if (step + 1) % 50 == 0: log(f"step {step+1}/{args.steps} loss {loss.item():.4f} " + " ".join(f"{k}:{v:.4f}" for k, v in items.items())) if args.vis_every > 0 and (step + 1) % args.vis_every == 0: try: preview_grid([lr_t.float()[:1], out.float()[:1], hr_t.float()[:1]], os.path.join(args.log_dir, f"step_{step+1:06d}.png")) except Exception as e: log(f"[warn] preview fail: {e}") if (step + 1) % args.save_every == 0: _save(full, args.out, step + 1) if ema is not None: _save_ema(ema, name_of, args.out, step + 1) step += 1 _save(full, args.out, step) if ema is not None: _save_ema(ema, name_of, args.out, step) log("S1 done") def _losses(out, hr, full, args, lpips_fn, dists_fn): out = out.float(); hr = hr.float() l1 = F.l1_loss(out, hr) items = {"l1": l1.item()} total = args.w_l1 * l1 if lpips_fn is not None: try: lp = lpips_fn(out.clamp(-1, 1), hr.clamp(-1, 1)).mean() total = total + args.w_lpips * lp; items["lpips"] = lp.item() except Exception: pass if dists_fn is not None: try: d = dists_fn((out.clamp(-1,1)+1)/2, (hr.clamp(-1,1)+1)/2).mean() total = total + args.w_dists * d; items["dists"] = d.item() except Exception: pass if args.w_wave > 0: mask = highfreq_mask(out.detach()) wav = (mask * (out - hr).abs()).mean() total = total + args.w_wave * wav; items["wave"] = wav.item() if args.w_color > 0: mo, so = out.mean(dim=(2,3)), out.std(dim=(2,3)) mh, sh = hr.mean(dim=(2,3)), hr.std(dim=(2,3)) col = (mo - mh).abs().mean() + (so - sh).abs().mean() total = total + args.w_color * col; items["color"] = col.item() return total, items def _save(full, out_dir, step): net = full[0] # Net (first module of full chain) torch.save(net.state_dict(), os.path.join(out_dir, f"net_params_{step}.pkl")) torch.save(full.state_dict(), os.path.join(out_dir, f"full_params_{step}.pkl")) def _save_ema(ema, name_of, out_dir, step): sd = {} for pid, val in ema.shadow.items(): nm = name_of.get(pid) if nm: sd[nm] = val.detach().cpu().clone() if sd: torch.save(sd, os.path.join(out_dir, f"lora_ema_{step}.pkl")) if __name__ == "__main__": main()