|
|
| """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
|
|
|
|
|
|
|
|
|
| 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
|
|
|
|
|
|
|
|
|
| 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))
|
|
|
| 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()
|
|
|
|
|
| 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
|
|
|
|
|
| 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_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")
|
|
|
|
|
| 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]
|
| 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()
|
|
|