CSIGv3_train_script / src /train_lora.py
XenderYang's picture
fix smoke bugs + anomaly guards; runbook update
ca2409a verified
Raw
History Blame Contribute Delete
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()