CSIGv3_train_script / src /common.py
XenderYang's picture
fix smoke bugs + anomaly guards; runbook update
ca2409a verified
Raw
History Blame Contribute Delete
13.2 kB
#!/usr/bin/env python
"""AdcSR 工程公共工具:路径、模型装配、手工 LoRA 注入、教师加载、GDPO probe。
训练脚本统一从这里 import,禁止各自重复实现装配逻辑。
"""
import os, sys, copy, json, types, math
from pathlib import Path
REPO = Path(__file__).resolve().parents[1]
OFFICIAL = REPO / "official"
def ensure_official():
if str(OFFICIAL) not in sys.path:
sys.path.insert(0, str(OFFICIAL))
ensure_official()
import torch
import torch.nn as nn
import torch.nn.functional as F
# ---------------------------------------------------------------------------
# 模型装配(与 official/test.py 全链一致)
# ---------------------------------------------------------------------------
def load_diffusers_sd(model_id, dtype=torch.float32, device="cpu", variant=None):
from diffusers import StableDiffusionPipeline
if variant is None:
# ???? fp16 ???????; ??? variant="" ???
import os as _os
if _os.path.isdir(model_id) and _os.path.exists(_os.path.join(model_id, "unet", "diffusion_pytorch_model.fp16.safetensors")):
variant = "fp16"
pipe = StableDiffusionPipeline.from_pretrained(model_id, torch_dtype=dtype,
variant=variant).to(device)
return pipe.vae, pipe.unet, pipe.text_encoder, pipe.tokenizer
def load_pruned_decoder(half_decoder_ckpt, device="cpu", dtype=torch.float32):
from diffusers.models.autoencoders.vae import Decoder
decoder = Decoder(in_channels=4, out_channels=3,
up_block_types=["UpDecoderBlock2D"] * 4,
block_out_channels=[64, 128, 256, 256], layers_per_block=2,
norm_num_groups=32, act_fn="silu", norm_type="group",
mid_block_add_attention=True).to(device=device, dtype=dtype)
ckpt = torch.load(half_decoder_ckpt, map_location="cpu", weights_only=False)
sd = {k.replace("decoder.", ""): v for k, v in ckpt["state_dict"].items() if k.startswith("decoder.")}
decoder.load_state_dict(sd, strict=True)
return decoder
def build_net(unet, decoder):
from model import Net # official
return Net(unet, copy.deepcopy(decoder))
def assemble_full_student(unet, decoder, net_weights=None, device="cuda", dtype=torch.float32):
"""???? 512 ???: Net(unet,decoder) + decoder ??(up_blocks...conv_out)?"""
net = build_net(unet, decoder)
if net_weights is not None:
sd = torch.load(net_weights, map_location="cpu", weights_only=False)
if any(k.startswith("module.") for k in sd):
sd = {k.replace("module.", "", 1): v for k, v in sd.items()}
net.load_state_dict(sd, strict=True)
net.to(device=device, dtype=dtype)
tail_mods = [*decoder.up_blocks, decoder.conv_norm_out, decoder.conv_act, decoder.conv_out]
for m in tail_mods:
m.to(device=device, dtype=dtype)
full = nn.Sequential(net, *tail_mods)
return full
def build_discriminator_unet(unet_copy, rank=4, dtype=torch.float32, device="cuda"):
"""official 判别器: conv_in 4->256 + LoRA(unet)。"""
from utils import add_lora_to_unet
unet_D = copy.deepcopy(unet_copy).to(device=device, dtype=dtype)
cin = unet_D.conv_in
new_conv_in = nn.Conv2d(256, cin.out_channels, 3, padding=1).to(device=device, dtype=dtype)
new_conv_in.weight.data = cin.weight.data.repeat(1, 64, 1, 1) / 64
new_conv_in.bias.data = cin.bias.data
unet_D.conv_in = new_conv_in
unet_D = add_lora_to_unet(unet_D, rank=rank)
unet_D.set_adapters(["default_encoder", "default_decoder", "default_others"])
return unet_D
# ---------------------------------------------------------------------------
# 教师加载
# ---------------------------------------------------------------------------
def load_osediff_teacher(osediff_pkl, device="cuda", dtype=torch.float32):
ckpt = torch.load(osediff_pkl, map_location="cpu", weights_only=False)
return ckpt # {"vae":..., "unet":...}
def load_gdpo_teacher(gdpo_dir, device="cuda", dtype=torch.float32):
"""GDPO ???????????????????????? probe_gdpo?
?? diffusers UNet2DConditionModel?state dict ????? dict?"""
from diffusers import UNet2DConditionModel
if os.path.isdir(os.path.join(gdpo_dir, "unet")) and os.path.exists(
os.path.join(gdpo_dir, "unet", "diffusion_pytorch_model.safetensors")):
return UNet2DConditionModel.from_pretrained(os.path.join(gdpo_dir, "unet"),
torch_dtype=dtype).to(device)
if os.path.isdir(gdpo_dir):
if os.path.exists(os.path.join(gdpo_dir, "diffusion_pytorch_model.safetensors")):
p = os.path.join(gdpo_dir, "diffusion_pytorch_model.safetensors")
else:
p = os.path.join(gdpo_dir, "ckp", "diffusion_pytorch_model.safetensors")
if os.path.exists(p):
try:
return UNet2DConditionModel.from_pretrained(os.path.dirname(p),
torch_dtype=dtype).to(device)
except Exception:
return _load_raw(p)
if os.path.isfile(gdpo_dir):
return _load_raw(gdpo_dir)
raise RuntimeError("GDPO ??????????? python -m src.common --probe_gdpo <path> ?????"
"??? --teacher osediff")
def _load_raw(p):
if p.endswith(".safetensors"):
from safetensors.torch import load_file
return load_file(p)
return torch.load(p, map_location="cpu", weights_only=False)
def probe_gdpo(gdpo_path):
"""打印权重键结构与前缀,帮助实现 GDPO->diffusers UNet 映射。"""
if gdpo_path.endswith(".safetensors"):
from safetensors.torch import load_file
sd = load_file(gdpo_path)
else:
sd = torch.load(gdpo_path, map_location="cpu", weights_only=False)
if isinstance(sd, dict) and "state_dict" in sd:
sd = sd["state_dict"]
keys = list(sd.keys())
print("num keys:", len(keys))
for k in keys[:40]:
print(k, tuple(sd[k].shape) if hasattr(sd[k], "shape") else type(sd[k]))
# 判断是否为完整 UNet(含 down_blocks)或 LoRA 或 Pipeline
has_unet = any("down_blocks" in k for k in keys)
has_lora = any("lora" in k.lower() for k in keys)
print("has_unet_blocks:", has_unet, "| has_lora:", has_lora)
# ---------------------------------------------------------------------------
# 手工 LoRA(对任意 Conv2d/Linear 注入,规避 peft 在剪枝/删模块后的解析问题)
# ---------------------------------------------------------------------------
class LoRAConv2d(nn.Module):
def __init__(self, conv: nn.Conv2d, r: int, alpha: float = 1.0):
super().__init__()
self.conv = conv
self.r = max(1, r)
self.alpha = alpha
self.cin = conv.in_channels
self.cout = conv.out_channels
self.lora_a = nn.Parameter(torch.zeros(self.cin, self.r))
self.lora_b = nn.Parameter(torch.zeros(self.r, self.cout))
nn.init.kaiming_uniform_(self.lora_a, a=5 ** 0.5)
nn.init.zeros_(self.lora_b)
self.requires_grad_(False)
self.lora_a.requires_grad_(True)
self.lora_b.requires_grad_(True)
def forward(self, x):
y = self.conv(x)
if self.training or True:
# 1x1 conv low-rank: 输入cin->r->cout, 保持空间尺寸
z = F.conv2d(x, self.lora_a.t().view(self.r, self.cin, 1, 1))
z = F.conv2d(z, self.lora_b.t().view(self.cout, self.r, 1, 1))
return y + self.alpha * z
return y
class LoRALinear(nn.Module):
def __init__(self, lin: nn.Linear, r: int, alpha: float = 1.0):
super().__init__()
self.lin = lin
self.r = max(1, r)
self.alpha = alpha
cin, cout = lin.in_features, lin.out_features
self.lora_a = nn.Parameter(torch.zeros(cin, self.r))
self.lora_b = nn.Parameter(torch.zeros(self.r, cout))
nn.init.kaiming_uniform_(self.lora_a, a=5 ** 0.5)
nn.init.zeros_(self.lora_b)
self.requires_grad_(False)
self.lora_a.requires_grad_(True)
self.lora_b.requires_grad_(True)
def forward(self, x):
y = self.lin(x)
z = F.linear(x, self.lora_a.t())
z = F.linear(z, self.lora_b.t())
return y + self.alpha * z
def _names(model):
for n, m in model.named_modules():
if isinstance(m, (nn.Conv2d, nn.Linear)):
yield n, m
def inject_lora(model, rank=64, alpha=1.0, skip_bias_norm=True, include=("conv", "to_q", "to_k", "to_v", "proj", "ff", "linear")):
"""替换模型内所有 Conv2d/Linear 为 LoRA 包装(原始权重冻结,仅训练 lora_a/b)。
include: 子串过滤,None=全部。"""
for n, m in list(_names(model)):
if include is not None and not any(s in n for s in include):
continue
parent, attr = _find_parent(model, n)
if isinstance(m, nn.Conv2d) and m.kernel_size == (1, 1):
setattr(parent, attr, LoRAConv2d(m, rank, alpha))
elif isinstance(m, nn.Conv2d):
continue # 3x3/stride>1 conv: 1x1 ???????, ??
elif isinstance(m, nn.Linear):
setattr(parent, attr, LoRALinear(m, rank, alpha))
# ????, ??? LoRA A/B ???(?????, ??"? LoRA ??")
model.requires_grad_(False)
for m in model.modules():
if isinstance(m, (LoRAConv2d, LoRALinear)):
m.lora_a.requires_grad_(True)
m.lora_b.requires_grad_(True)
return model
def _find_parent(model, name):
parts = name.split(".")
node = model
for p in parts[:-1]:
node = getattr(node, p)
return node, parts[-1]
def lora_params(model):
for p in model.parameters():
if p.requires_grad:
yield p
# ---------------------------------------------------------------------------
# ???????(????/????/EMA/???) 2026-09-06
# ---------------------------------------------------------------------------
def is_finite(x):
"""??/???????(? NaN/Inf)?"""
try:
if torch.is_tensor(x):
return bool(torch.isfinite(x.float()).all().item())
return bool(math.isfinite(float(x)))
except Exception:
return False
def check_tensor(x, name, log=None):
"""??/??????: ?? True=???"""
if x is None:
return False
if torch.is_tensor(x) and not is_finite(x):
msg = f"[anomaly] {name} contains NaN/Inf"
print(msg, flush=True)
if log is not None:
log(msg)
return True
return False
def clip_and_check_grads(params, max_norm, log=None):
"""???? + NaN/Inf ??; ?? True=????(??? step)?"""
grads = [p.grad for p in params if p.grad is not None]
bad = False
for g in grads:
if not is_finite(g):
bad = True
msg = "[anomaly] grad contains NaN/Inf; skip this optimizer step"
print(msg, flush=True)
if log is not None:
log(msg)
break
if bad:
return True
if max_norm and max_norm > 0 and grads:
total = torch.nn.utils.clip_grad_norm_(params, max_norm=max_norm)
if not is_finite(total):
msg = "[anomaly] grad total norm NaN; skip step"
print(msg, flush=True)
if log is not None:
log(msg)
return True
return False
class EMA:
"""??????(?? trainable/lora ??)?"""
def __init__(self, params, decay=0.999):
self.decay = decay
self.shadow = {id(p): p.detach().clone().float() for p in params if p.requires_grad}
@torch.no_grad()
def update(self, params):
d = self.decay
for p in params:
if not p.requires_grad or id(p) not in self.shadow:
continue
self.shadow[id(p)].mul_(d).add_(p.detach().float(), alpha=1 - d)
def state_dict(self, params):
return {id(p): self.shadow[id(p)] for p in params if id(p) in self.shadow}
def preview_grid(tensors, path, vmin=-1.0, vmax=1.0):
"""? [B,C,H,W] ??([-1,1]) ?????? PNG, ????????/?????"""
import numpy as np
from PIL import Image
ims = []
for t in tensors:
t = t.detach().float().clamp(vmin, vmax)
t = (t - vmin) / (vmax - vmin)
b = t[0].clamp(0, 1).permute(1, 2, 0).cpu().numpy()
ims.append(Image.fromarray((b * 255).astype(np.uint8)))
w = sum(im.width for im in ims); h = max(im.height for im in ims)
canvas = Image.new("RGB", (w, h), (0, 0, 0))
x = 0
for im in ims:
canvas.paste(im, (x, 0)); x += im.width
canvas.save(path, quality=92)
def count_params(model, only_trainable=False):
if only_trainable:
return sum(p.numel() for p in model.parameters() if p.requires_grad)
return sum(p.numel() for p in model.parameters())
if __name__ == "__main__":
print("common module OK; official dir:", OFFICIAL)