Download src/train_lora.py from XenderYang/CSIGv3_train_script: direct link, hf CLI and curl.
- Browser
- Download file 14.3 kB
-
https://huggingface.co/XenderYang/CSIGv3_train_script/resolve/main/src/train_lora.py
- Command line
-
hf download hf://XenderYang/CSIGv3_train_script/src/train_lora.py
-
curl -L -o train_lora.py https://huggingface.co/XenderYang/CSIGv3_train_script/resolve/main/src/train_lora.py
14.3 kB
| #!/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() | |