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