fix smoke bugs + anomaly guards; runbook update
Browse files- RUNBOOK_A100.md +13 -0
- configs/config_smoke.yml +36 -36
- official/bsr/utils/img_process_util.py +6 -6
- src/common.py +105 -10
- src/train_lora.py +57 -8
- src/train_stage2.py +40 -4
RUNBOOK_A100.md
CHANGED
|
@@ -49,3 +49,16 @@ bash scripts/run_stage3.sh # S3 域微调 (20k step)
|
|
| 49 |
- 训练脚本依赖 torchvision(官方 bsr/degradations.py import)。A100(Linux+torchvision)正常;
|
| 50 |
本机 Windows 沙箱若 torchvision 报 `operator torchvision::nms does not exist`, 请在自己终端(非沙箱)运行。
|
| 51 |
- 数据清单中的真实退化对仅 RealSR 400/99; 其余母本在 data/clean_all, data/clean。
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 49 |
- 训练脚本依赖 torchvision(官方 bsr/degradations.py import)。A100(Linux+torchvision)正常;
|
| 50 |
本机 Windows 沙箱若 torchvision 报 `operator torchvision::nms does not exist`, 请在自己终端(非沙箱)运行。
|
| 51 |
- 数据清单中的真实退化对仅 RealSR 400/99; 其余母本在 data/clean_all, data/clean。
|
| 52 |
+
|
| 53 |
+
## 本地冒烟结果(2026-09-06, conda tvtest CPU)
|
| 54 |
+
- conda 环境 torch 2.5.1 + torchvision 0.20.1 可用(需 PYTHONNOUSERSITE=1 屏蔽用户目录污染)
|
| 55 |
+
- 修复清单(已同步到本仓库):
|
| 56 |
+
* common.py: LoRAConv2d cin/cout 未定义; LoRA 仅注入 Linear/1x1Conv(3x3 分解不匹配);
|
| 57 |
+
整体冻结仅训 LoRA; assemble_full_student 展开 up_blocks(ModuleList 不能进 Sequential);
|
| 58 |
+
load_diffusers_sd 自动 variant=fp16; 新增 is_finite/check_tensor/clip_and_check_grads/EMA/preview_grid
|
| 59 |
+
* train_lora.py/train_stage2.py: NaN/Inf 检测(输入/输出/损失/梯度)、梯度裁剪(clip_grad 1.0)、
|
| 60 |
+
skip-step 连续异常>20 中止、EMA(默认0.999)、周期预览图(vis_every)、num_workers 冒烟用0
|
| 61 |
+
* official/bsr/utils/img_process_util.py: filter2D .view->.reshape(非连续张量)
|
| 62 |
+
* configs/config_smoke.yml: gt_size 128->512(Net 输入需 LR128=HR512/4)
|
| 63 |
+
- 冒烟结果: 退化管线 OK(HR512->LR128, 数值正常, 预览 logs/smoke_degrade_preview.png);
|
| 64 |
+
S1 LoRA 2步 OK(checkpoint/lora_ema/预览); S2 本地 CPU 内存不足被系统终止 -> 在 A100 跑 S2/GDPO 冒烟
|
configs/config_smoke.yml
CHANGED
|
@@ -1,37 +1,37 @@
|
|
| 1 |
-
dataroot_gt: data/patches
|
| 2 |
-
gt_size:
|
| 3 |
-
iter_num: 2
|
| 4 |
-
# Real-ESRGAN style degradation (shared by all stages)
|
| 5 |
-
scale: 4
|
| 6 |
-
resize_prob: [0.2, 0.7, 0.1]
|
| 7 |
-
resize_range: [0.3, 1.5]
|
| 8 |
-
gaussian_noise_prob: 0.5
|
| 9 |
-
noise_range: [1, 15]
|
| 10 |
-
poisson_scale_range: [0.05, 2.0]
|
| 11 |
-
gray_noise_prob: 0.4
|
| 12 |
-
jpeg_range: [60, 95]
|
| 13 |
-
second_blur_prob: 0.5
|
| 14 |
-
resize_prob2: [0.3, 0.4, 0.3]
|
| 15 |
-
resize_range2: [0.6, 1.2]
|
| 16 |
-
gaussian_noise_prob2: 0.5
|
| 17 |
-
noise_range2: [1, 12]
|
| 18 |
-
poisson_scale_range2: [0.05, 1.0]
|
| 19 |
-
gray_noise_prob2: 0.4
|
| 20 |
-
jpeg_range2: [60, 100]
|
| 21 |
-
blur_kernel_size: 21
|
| 22 |
-
kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
|
| 23 |
-
kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
|
| 24 |
-
sinc_prob: 0.1
|
| 25 |
-
blur_sigma: [0.2, 1.5]
|
| 26 |
-
betag_range: [0.5, 2.0]
|
| 27 |
-
betap_range: [1, 1.5]
|
| 28 |
-
blur_kernel_size2: 11
|
| 29 |
-
kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
|
| 30 |
-
kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
|
| 31 |
-
sinc_prob2: 0.1
|
| 32 |
-
blur_sigma2: [0.2, 1.0]
|
| 33 |
-
betag_range2: [0.5, 2.0]
|
| 34 |
-
betap_range2: [1, 1.5]
|
| 35 |
-
final_sinc_prob: 0.8
|
| 36 |
-
use_hflip: True
|
| 37 |
use_rot: False
|
|
|
|
| 1 |
+
dataroot_gt: data/patches
|
| 2 |
+
gt_size: 512
|
| 3 |
+
iter_num: 2
|
| 4 |
+
# Real-ESRGAN style degradation (shared by all stages)
|
| 5 |
+
scale: 4
|
| 6 |
+
resize_prob: [0.2, 0.7, 0.1]
|
| 7 |
+
resize_range: [0.3, 1.5]
|
| 8 |
+
gaussian_noise_prob: 0.5
|
| 9 |
+
noise_range: [1, 15]
|
| 10 |
+
poisson_scale_range: [0.05, 2.0]
|
| 11 |
+
gray_noise_prob: 0.4
|
| 12 |
+
jpeg_range: [60, 95]
|
| 13 |
+
second_blur_prob: 0.5
|
| 14 |
+
resize_prob2: [0.3, 0.4, 0.3]
|
| 15 |
+
resize_range2: [0.6, 1.2]
|
| 16 |
+
gaussian_noise_prob2: 0.5
|
| 17 |
+
noise_range2: [1, 12]
|
| 18 |
+
poisson_scale_range2: [0.05, 1.0]
|
| 19 |
+
gray_noise_prob2: 0.4
|
| 20 |
+
jpeg_range2: [60, 100]
|
| 21 |
+
blur_kernel_size: 21
|
| 22 |
+
kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
|
| 23 |
+
kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
|
| 24 |
+
sinc_prob: 0.1
|
| 25 |
+
blur_sigma: [0.2, 1.5]
|
| 26 |
+
betag_range: [0.5, 2.0]
|
| 27 |
+
betap_range: [1, 1.5]
|
| 28 |
+
blur_kernel_size2: 11
|
| 29 |
+
kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
|
| 30 |
+
kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
|
| 31 |
+
sinc_prob2: 0.1
|
| 32 |
+
blur_sigma2: [0.2, 1.0]
|
| 33 |
+
betag_range2: [0.5, 2.0]
|
| 34 |
+
betap_range2: [1, 1.5]
|
| 35 |
+
final_sinc_prob: 0.8
|
| 36 |
+
use_hflip: True
|
| 37 |
use_rot: False
|
official/bsr/utils/img_process_util.py
CHANGED
|
@@ -22,13 +22,13 @@ def filter2D(img, kernel):
|
|
| 22 |
|
| 23 |
if kernel.size(0) == 1:
|
| 24 |
# apply the same kernel to all batch images
|
| 25 |
-
img = img.
|
| 26 |
-
kernel = kernel.
|
| 27 |
-
return F.conv2d(img, kernel, padding=0).
|
| 28 |
else:
|
| 29 |
-
img = img.
|
| 30 |
-
kernel = kernel.
|
| 31 |
-
return F.conv2d(img, kernel, groups=b * c).
|
| 32 |
|
| 33 |
|
| 34 |
def usm_sharp(img, weight=0.5, radius=50, threshold=10):
|
|
|
|
| 22 |
|
| 23 |
if kernel.size(0) == 1:
|
| 24 |
# apply the same kernel to all batch images
|
| 25 |
+
img = img.reshape(b * c, 1, ph, pw)
|
| 26 |
+
kernel = kernel.reshape(1, 1, k, k)
|
| 27 |
+
return F.conv2d(img, kernel, padding=0).reshape(b, c, h, w)
|
| 28 |
else:
|
| 29 |
+
img = img.reshape(1, b * c, ph, pw)
|
| 30 |
+
kernel = kernel.reshape(b, 1, k, k).repeat(1, c, 1, 1).reshape(b * c, 1, k, k)
|
| 31 |
+
return F.conv2d(img, kernel, groups=b * c).reshape(b, c, h, w)
|
| 32 |
|
| 33 |
|
| 34 |
def usm_sharp(img, weight=0.5, radius=50, threshold=10):
|
src/common.py
CHANGED
|
@@ -2,7 +2,7 @@
|
|
| 2 |
"""AdcSR 工程公共工具:路径、模型装配、手工 LoRA 注入、教师加载、GDPO probe。
|
| 3 |
训练脚本统一从这里 import,禁止各自重复实现装配逻辑。
|
| 4 |
"""
|
| 5 |
-
import os, sys, copy, json, types
|
| 6 |
from pathlib import Path
|
| 7 |
|
| 8 |
REPO = Path(__file__).resolve().parents[1]
|
|
@@ -21,9 +21,15 @@ import torch.nn.functional as F
|
|
| 21 |
# ---------------------------------------------------------------------------
|
| 22 |
# 模型装配(与 official/test.py 全链一致)
|
| 23 |
# ---------------------------------------------------------------------------
|
| 24 |
-
def load_diffusers_sd(model_id, dtype=torch.float32, device="cpu"):
|
| 25 |
from diffusers import StableDiffusionPipeline
|
| 26 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 27 |
return pipe.vae, pipe.unet, pipe.text_encoder, pipe.tokenizer
|
| 28 |
|
| 29 |
def load_pruned_decoder(half_decoder_ckpt, device="cpu", dtype=torch.float32):
|
|
@@ -51,7 +57,7 @@ def assemble_full_student(unet, decoder, net_weights=None, device="cuda", dtype=
|
|
| 51 |
sd = {k.replace("module.", "", 1): v for k, v in sd.items()}
|
| 52 |
net.load_state_dict(sd, strict=True)
|
| 53 |
net.to(device=device, dtype=dtype)
|
| 54 |
-
tail_mods = [decoder.up_blocks, decoder.conv_norm_out, decoder.conv_act, decoder.conv_out]
|
| 55 |
for m in tail_mods:
|
| 56 |
m.to(device=device, dtype=dtype)
|
| 57 |
full = nn.Sequential(net, *tail_mods)
|
|
@@ -135,9 +141,10 @@ class LoRAConv2d(nn.Module):
|
|
| 135 |
self.conv = conv
|
| 136 |
self.r = max(1, r)
|
| 137 |
self.alpha = alpha
|
| 138 |
-
cin
|
| 139 |
-
self.
|
| 140 |
-
self.
|
|
|
|
| 141 |
nn.init.kaiming_uniform_(self.lora_a, a=5 ** 0.5)
|
| 142 |
nn.init.zeros_(self.lora_b)
|
| 143 |
self.requires_grad_(False)
|
|
@@ -148,8 +155,8 @@ class LoRAConv2d(nn.Module):
|
|
| 148 |
y = self.conv(x)
|
| 149 |
if self.training or True:
|
| 150 |
# 1x1 conv low-rank: 输入cin->r->cout, 保持空间尺寸
|
| 151 |
-
z = F.conv2d(x, self.lora_a.t().view(self.r, cin, 1, 1))
|
| 152 |
-
z = F.conv2d(z, self.lora_b.t().view(cout, self.r, 1, 1))
|
| 153 |
return y + self.alpha * z
|
| 154 |
return y
|
| 155 |
|
|
@@ -186,10 +193,18 @@ def inject_lora(model, rank=64, alpha=1.0, skip_bias_norm=True, include=("conv",
|
|
| 186 |
if include is not None and not any(s in n for s in include):
|
| 187 |
continue
|
| 188 |
parent, attr = _find_parent(model, n)
|
| 189 |
-
if isinstance(m, nn.Conv2d):
|
| 190 |
setattr(parent, attr, LoRAConv2d(m, rank, alpha))
|
|
|
|
|
|
|
| 191 |
elif isinstance(m, nn.Linear):
|
| 192 |
setattr(parent, attr, LoRALinear(m, rank, alpha))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 193 |
return model
|
| 194 |
|
| 195 |
def _find_parent(model, name):
|
|
@@ -204,6 +219,86 @@ def lora_params(model):
|
|
| 204 |
if p.requires_grad:
|
| 205 |
yield p
|
| 206 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 207 |
def count_params(model, only_trainable=False):
|
| 208 |
if only_trainable:
|
| 209 |
return sum(p.numel() for p in model.parameters() if p.requires_grad)
|
|
|
|
| 2 |
"""AdcSR 工程公共工具:路径、模型装配、手工 LoRA 注入、教师加载、GDPO probe。
|
| 3 |
训练脚本统一从这里 import,禁止各自重复实现装配逻辑。
|
| 4 |
"""
|
| 5 |
+
import os, sys, copy, json, types, math
|
| 6 |
from pathlib import Path
|
| 7 |
|
| 8 |
REPO = Path(__file__).resolve().parents[1]
|
|
|
|
| 21 |
# ---------------------------------------------------------------------------
|
| 22 |
# 模型装配(与 official/test.py 全链一致)
|
| 23 |
# ---------------------------------------------------------------------------
|
| 24 |
+
def load_diffusers_sd(model_id, dtype=torch.float32, device="cpu", variant=None):
|
| 25 |
from diffusers import StableDiffusionPipeline
|
| 26 |
+
if variant is None:
|
| 27 |
+
# ???? fp16 ???????; ??? variant="" ???
|
| 28 |
+
import os as _os
|
| 29 |
+
if _os.path.isdir(model_id) and _os.path.exists(_os.path.join(model_id, "unet", "diffusion_pytorch_model.fp16.safetensors")):
|
| 30 |
+
variant = "fp16"
|
| 31 |
+
pipe = StableDiffusionPipeline.from_pretrained(model_id, torch_dtype=dtype,
|
| 32 |
+
variant=variant).to(device)
|
| 33 |
return pipe.vae, pipe.unet, pipe.text_encoder, pipe.tokenizer
|
| 34 |
|
| 35 |
def load_pruned_decoder(half_decoder_ckpt, device="cpu", dtype=torch.float32):
|
|
|
|
| 57 |
sd = {k.replace("module.", "", 1): v for k, v in sd.items()}
|
| 58 |
net.load_state_dict(sd, strict=True)
|
| 59 |
net.to(device=device, dtype=dtype)
|
| 60 |
+
tail_mods = [*decoder.up_blocks, decoder.conv_norm_out, decoder.conv_act, decoder.conv_out]
|
| 61 |
for m in tail_mods:
|
| 62 |
m.to(device=device, dtype=dtype)
|
| 63 |
full = nn.Sequential(net, *tail_mods)
|
|
|
|
| 141 |
self.conv = conv
|
| 142 |
self.r = max(1, r)
|
| 143 |
self.alpha = alpha
|
| 144 |
+
self.cin = conv.in_channels
|
| 145 |
+
self.cout = conv.out_channels
|
| 146 |
+
self.lora_a = nn.Parameter(torch.zeros(self.cin, self.r))
|
| 147 |
+
self.lora_b = nn.Parameter(torch.zeros(self.r, self.cout))
|
| 148 |
nn.init.kaiming_uniform_(self.lora_a, a=5 ** 0.5)
|
| 149 |
nn.init.zeros_(self.lora_b)
|
| 150 |
self.requires_grad_(False)
|
|
|
|
| 155 |
y = self.conv(x)
|
| 156 |
if self.training or True:
|
| 157 |
# 1x1 conv low-rank: 输入cin->r->cout, 保持空间尺寸
|
| 158 |
+
z = F.conv2d(x, self.lora_a.t().view(self.r, self.cin, 1, 1))
|
| 159 |
+
z = F.conv2d(z, self.lora_b.t().view(self.cout, self.r, 1, 1))
|
| 160 |
return y + self.alpha * z
|
| 161 |
return y
|
| 162 |
|
|
|
|
| 193 |
if include is not None and not any(s in n for s in include):
|
| 194 |
continue
|
| 195 |
parent, attr = _find_parent(model, n)
|
| 196 |
+
if isinstance(m, nn.Conv2d) and m.kernel_size == (1, 1):
|
| 197 |
setattr(parent, attr, LoRAConv2d(m, rank, alpha))
|
| 198 |
+
elif isinstance(m, nn.Conv2d):
|
| 199 |
+
continue # 3x3/stride>1 conv: 1x1 ???????, ??
|
| 200 |
elif isinstance(m, nn.Linear):
|
| 201 |
setattr(parent, attr, LoRALinear(m, rank, alpha))
|
| 202 |
+
# ????, ??? LoRA A/B ???(?????, ??"? LoRA ??")
|
| 203 |
+
model.requires_grad_(False)
|
| 204 |
+
for m in model.modules():
|
| 205 |
+
if isinstance(m, (LoRAConv2d, LoRALinear)):
|
| 206 |
+
m.lora_a.requires_grad_(True)
|
| 207 |
+
m.lora_b.requires_grad_(True)
|
| 208 |
return model
|
| 209 |
|
| 210 |
def _find_parent(model, name):
|
|
|
|
| 219 |
if p.requires_grad:
|
| 220 |
yield p
|
| 221 |
|
| 222 |
+
# ---------------------------------------------------------------------------
|
| 223 |
+
# ???????(????/????/EMA/???) 2026-09-06
|
| 224 |
+
# ---------------------------------------------------------------------------
|
| 225 |
+
def is_finite(x):
|
| 226 |
+
"""??/???????(? NaN/Inf)?"""
|
| 227 |
+
try:
|
| 228 |
+
if torch.is_tensor(x):
|
| 229 |
+
return bool(torch.isfinite(x.float()).all().item())
|
| 230 |
+
return bool(math.isfinite(float(x)))
|
| 231 |
+
except Exception:
|
| 232 |
+
return False
|
| 233 |
+
|
| 234 |
+
def check_tensor(x, name, log=None):
|
| 235 |
+
"""??/??????: ?? True=???"""
|
| 236 |
+
if x is None:
|
| 237 |
+
return False
|
| 238 |
+
if torch.is_tensor(x) and not is_finite(x):
|
| 239 |
+
msg = f"[anomaly] {name} contains NaN/Inf"
|
| 240 |
+
print(msg, flush=True)
|
| 241 |
+
if log is not None:
|
| 242 |
+
log(msg)
|
| 243 |
+
return True
|
| 244 |
+
return False
|
| 245 |
+
|
| 246 |
+
def clip_and_check_grads(params, max_norm, log=None):
|
| 247 |
+
"""???? + NaN/Inf ??; ?? True=????(??? step)?"""
|
| 248 |
+
grads = [p.grad for p in params if p.grad is not None]
|
| 249 |
+
bad = False
|
| 250 |
+
for g in grads:
|
| 251 |
+
if not is_finite(g):
|
| 252 |
+
bad = True
|
| 253 |
+
msg = "[anomaly] grad contains NaN/Inf; skip this optimizer step"
|
| 254 |
+
print(msg, flush=True)
|
| 255 |
+
if log is not None:
|
| 256 |
+
log(msg)
|
| 257 |
+
break
|
| 258 |
+
if bad:
|
| 259 |
+
return True
|
| 260 |
+
if max_norm and max_norm > 0 and grads:
|
| 261 |
+
total = torch.nn.utils.clip_grad_norm_(params, max_norm=max_norm)
|
| 262 |
+
if not is_finite(total):
|
| 263 |
+
msg = "[anomaly] grad total norm NaN; skip step"
|
| 264 |
+
print(msg, flush=True)
|
| 265 |
+
if log is not None:
|
| 266 |
+
log(msg)
|
| 267 |
+
return True
|
| 268 |
+
return False
|
| 269 |
+
|
| 270 |
+
class EMA:
|
| 271 |
+
"""??????(?? trainable/lora ??)?"""
|
| 272 |
+
def __init__(self, params, decay=0.999):
|
| 273 |
+
self.decay = decay
|
| 274 |
+
self.shadow = {id(p): p.detach().clone().float() for p in params if p.requires_grad}
|
| 275 |
+
@torch.no_grad()
|
| 276 |
+
def update(self, params):
|
| 277 |
+
d = self.decay
|
| 278 |
+
for p in params:
|
| 279 |
+
if not p.requires_grad or id(p) not in self.shadow:
|
| 280 |
+
continue
|
| 281 |
+
self.shadow[id(p)].mul_(d).add_(p.detach().float(), alpha=1 - d)
|
| 282 |
+
def state_dict(self, params):
|
| 283 |
+
return {id(p): self.shadow[id(p)] for p in params if id(p) in self.shadow}
|
| 284 |
+
|
| 285 |
+
def preview_grid(tensors, path, vmin=-1.0, vmax=1.0):
|
| 286 |
+
"""? [B,C,H,W] ??([-1,1]) ?????? PNG, ????????/?????"""
|
| 287 |
+
import numpy as np
|
| 288 |
+
from PIL import Image
|
| 289 |
+
ims = []
|
| 290 |
+
for t in tensors:
|
| 291 |
+
t = t.detach().float().clamp(vmin, vmax)
|
| 292 |
+
t = (t - vmin) / (vmax - vmin)
|
| 293 |
+
b = t[0].clamp(0, 1).permute(1, 2, 0).cpu().numpy()
|
| 294 |
+
ims.append(Image.fromarray((b * 255).astype(np.uint8)))
|
| 295 |
+
w = sum(im.width for im in ims); h = max(im.height for im in ims)
|
| 296 |
+
canvas = Image.new("RGB", (w, h), (0, 0, 0))
|
| 297 |
+
x = 0
|
| 298 |
+
for im in ims:
|
| 299 |
+
canvas.paste(im, (x, 0)); x += im.width
|
| 300 |
+
canvas.save(path, quality=92)
|
| 301 |
+
|
| 302 |
def count_params(model, only_trainable=False):
|
| 303 |
if only_trainable:
|
| 304 |
return sum(p.numel() for p in model.parameters() if p.requires_grad)
|
src/train_lora.py
CHANGED
|
@@ -21,7 +21,8 @@ from omegaconf import OmegaConf
|
|
| 21 |
|
| 22 |
from common import (ensure_official, load_diffusers_sd, load_pruned_decoder,
|
| 23 |
assemble_full_student, inject_lora, lora_params,
|
| 24 |
-
count_params, build_net
|
|
|
|
| 25 |
ensure_official()
|
| 26 |
from dataset import RealESRGANDataset, RealESRGANDegrader # official
|
| 27 |
|
|
@@ -127,6 +128,9 @@ def main():
|
|
| 127 |
ap.add_argument("--no_bf16", dest="bf16", action="store_false")
|
| 128 |
ap.add_argument("--seed", type=int, default=123)
|
| 129 |
ap.add_argument("--num_workers", type=int, default=8)
|
|
|
|
|
|
|
|
|
|
| 130 |
args = ap.parse_args()
|
| 131 |
random.seed(args.seed); torch.manual_seed(args.seed)
|
| 132 |
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
@@ -175,9 +179,11 @@ def main():
|
|
| 175 |
except Exception as e:
|
| 176 |
log(f"[warn] pyiqa dists 不可用: {e}; 该损失置 0")
|
| 177 |
|
| 178 |
-
# ---- train loop ----
|
| 179 |
syn_iter = iter(syn_dl); real_iter = iter(real_dl) if real_dl else None
|
| 180 |
-
step = 0
|
|
|
|
|
|
|
| 181 |
full.train()
|
| 182 |
optimizer.zero_grad(set_to_none=True)
|
| 183 |
while step < args.steps:
|
|
@@ -193,6 +199,11 @@ def main():
|
|
| 193 |
real_iter = iter(real_dl) if real_dl else None
|
| 194 |
continue
|
| 195 |
lr_t, hr_t = lr_t.to(device), hr_t.to(device)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 196 |
if dtype == torch.bfloat16:
|
| 197 |
with torch.autocast("cuda", dtype=torch.bfloat16):
|
| 198 |
out = full(lr_t)
|
|
@@ -200,21 +211,50 @@ def main():
|
|
| 200 |
else:
|
| 201 |
out = full(lr_t)
|
| 202 |
loss, items = _losses(out, hr_t, full, args, lpips_fn, dists_fn)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 203 |
(loss / args.grad_accum).backward()
|
| 204 |
if (step + 1) % args.grad_accum == 0:
|
| 205 |
-
if
|
| 206 |
-
|
| 207 |
else:
|
| 208 |
-
|
| 209 |
-
|
| 210 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 211 |
if (step + 1) % 50 == 0:
|
| 212 |
log(f"step {step+1}/{args.steps} loss {loss.item():.4f} " +
|
| 213 |
" ".join(f"{k}:{v:.4f}" for k, v in items.items()))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 214 |
if (step + 1) % args.save_every == 0:
|
| 215 |
_save(full, args.out, step + 1)
|
|
|
|
|
|
|
| 216 |
step += 1
|
| 217 |
_save(full, args.out, step)
|
|
|
|
|
|
|
| 218 |
log("S1 done")
|
| 219 |
|
| 220 |
def _losses(out, hr, full, args, lpips_fn, dists_fn):
|
|
@@ -250,5 +290,14 @@ def _save(full, out_dir, step):
|
|
| 250 |
torch.save(net.state_dict(), os.path.join(out_dir, f"net_params_{step}.pkl"))
|
| 251 |
torch.save(full.state_dict(), os.path.join(out_dir, f"full_params_{step}.pkl"))
|
| 252 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 253 |
if __name__ == "__main__":
|
| 254 |
main()
|
|
|
|
| 21 |
|
| 22 |
from common import (ensure_official, load_diffusers_sd, load_pruned_decoder,
|
| 23 |
assemble_full_student, inject_lora, lora_params,
|
| 24 |
+
count_params, build_net, is_finite, check_tensor,
|
| 25 |
+
clip_and_check_grads, EMA, preview_grid)
|
| 26 |
ensure_official()
|
| 27 |
from dataset import RealESRGANDataset, RealESRGANDegrader # official
|
| 28 |
|
|
|
|
| 128 |
ap.add_argument("--no_bf16", dest="bf16", action="store_false")
|
| 129 |
ap.add_argument("--seed", type=int, default=123)
|
| 130 |
ap.add_argument("--num_workers", type=int, default=8)
|
| 131 |
+
ap.add_argument("--clip_grad", type=float, default=1.0, help="??????; 0=??")
|
| 132 |
+
ap.add_argument("--ema_decay", type=float, default=0.999, help="EMA ??; 0=??")
|
| 133 |
+
ap.add_argument("--vis_every", type=int, default=500, help="? N ?? LR/HR/?????")
|
| 134 |
args = ap.parse_args()
|
| 135 |
random.seed(args.seed); torch.manual_seed(args.seed)
|
| 136 |
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
|
|
| 179 |
except Exception as e:
|
| 180 |
log(f"[warn] pyiqa dists 不可用: {e}; 该损失置 0")
|
| 181 |
|
| 182 |
+
# ---- train loop (????/????/EMA/???) ----
|
| 183 |
syn_iter = iter(syn_dl); real_iter = iter(real_dl) if real_dl else None
|
| 184 |
+
step = 0; skip_streak = 0
|
| 185 |
+
ema = EMA(params, args.ema_decay) if args.ema_decay > 0 else None
|
| 186 |
+
name_of = {id(p): n for n, p in full.named_parameters() if p.requires_grad}
|
| 187 |
full.train()
|
| 188 |
optimizer.zero_grad(set_to_none=True)
|
| 189 |
while step < args.steps:
|
|
|
|
| 199 |
real_iter = iter(real_dl) if real_dl else None
|
| 200 |
continue
|
| 201 |
lr_t, hr_t = lr_t.to(device), hr_t.to(device)
|
| 202 |
+
if check_tensor(lr_t, "lr", log) or check_tensor(hr_t, "hr", log):
|
| 203 |
+
skip_streak += 1
|
| 204 |
+
if skip_streak > 20:
|
| 205 |
+
log("[anomaly] too many bad batches, abort"); break
|
| 206 |
+
continue
|
| 207 |
if dtype == torch.bfloat16:
|
| 208 |
with torch.autocast("cuda", dtype=torch.bfloat16):
|
| 209 |
out = full(lr_t)
|
|
|
|
| 211 |
else:
|
| 212 |
out = full(lr_t)
|
| 213 |
loss, items = _losses(out, hr_t, full, args, lpips_fn, dists_fn)
|
| 214 |
+
if check_tensor(out, "output", log):
|
| 215 |
+
skip_streak += 1
|
| 216 |
+
if skip_streak > 20:
|
| 217 |
+
log("[anomaly] too many bad outputs, abort"); break
|
| 218 |
+
optimizer.zero_grad(set_to_none=True)
|
| 219 |
+
continue
|
| 220 |
+
if not is_finite(loss):
|
| 221 |
+
log(f"[anomaly] loss NaN/Inf at step {step+1}; skip step")
|
| 222 |
+
skip_streak += 1
|
| 223 |
+
optimizer.zero_grad(set_to_none=True)
|
| 224 |
+
if skip_streak > 20:
|
| 225 |
+
log("[anomaly] too many bad losses, abort"); break
|
| 226 |
+
continue
|
| 227 |
+
skip_streak = 0
|
| 228 |
(loss / args.grad_accum).backward()
|
| 229 |
if (step + 1) % args.grad_accum == 0:
|
| 230 |
+
if clip_and_check_grads(params, args.clip_grad, log):
|
| 231 |
+
optimizer.zero_grad(set_to_none=True)
|
| 232 |
else:
|
| 233 |
+
if scaler is not None:
|
| 234 |
+
scaler.step(optimizer); scaler.update()
|
| 235 |
+
else:
|
| 236 |
+
optimizer.step()
|
| 237 |
+
optimizer.zero_grad(set_to_none=True)
|
| 238 |
+
sched.step()
|
| 239 |
+
if ema is not None:
|
| 240 |
+
ema.update(params)
|
| 241 |
if (step + 1) % 50 == 0:
|
| 242 |
log(f"step {step+1}/{args.steps} loss {loss.item():.4f} " +
|
| 243 |
" ".join(f"{k}:{v:.4f}" for k, v in items.items()))
|
| 244 |
+
if args.vis_every > 0 and (step + 1) % args.vis_every == 0:
|
| 245 |
+
try:
|
| 246 |
+
preview_grid([lr_t.float()[:1], out.float()[:1], hr_t.float()[:1]],
|
| 247 |
+
os.path.join(args.log_dir, f"step_{step+1:06d}.png"))
|
| 248 |
+
except Exception as e:
|
| 249 |
+
log(f"[warn] preview fail: {e}")
|
| 250 |
if (step + 1) % args.save_every == 0:
|
| 251 |
_save(full, args.out, step + 1)
|
| 252 |
+
if ema is not None:
|
| 253 |
+
_save_ema(ema, name_of, args.out, step + 1)
|
| 254 |
step += 1
|
| 255 |
_save(full, args.out, step)
|
| 256 |
+
if ema is not None:
|
| 257 |
+
_save_ema(ema, name_of, args.out, step)
|
| 258 |
log("S1 done")
|
| 259 |
|
| 260 |
def _losses(out, hr, full, args, lpips_fn, dists_fn):
|
|
|
|
| 290 |
torch.save(net.state_dict(), os.path.join(out_dir, f"net_params_{step}.pkl"))
|
| 291 |
torch.save(full.state_dict(), os.path.join(out_dir, f"full_params_{step}.pkl"))
|
| 292 |
|
| 293 |
+
def _save_ema(ema, name_of, out_dir, step):
|
| 294 |
+
sd = {}
|
| 295 |
+
for pid, val in ema.shadow.items():
|
| 296 |
+
nm = name_of.get(pid)
|
| 297 |
+
if nm:
|
| 298 |
+
sd[nm] = val.detach().cpu().clone()
|
| 299 |
+
if sd:
|
| 300 |
+
torch.save(sd, os.path.join(out_dir, f"lora_ema_{step}.pkl"))
|
| 301 |
+
|
| 302 |
if __name__ == "__main__":
|
| 303 |
main()
|
src/train_stage2.py
CHANGED
|
@@ -17,7 +17,8 @@ from omegaconf import OmegaConf
|
|
| 17 |
|
| 18 |
from common import (ensure_official, load_diffusers_sd, load_pruned_decoder,
|
| 19 |
build_net, build_discriminator_unet, load_osediff_teacher,
|
| 20 |
-
load_gdpo_teacher, count_params
|
|
|
|
| 21 |
ensure_official()
|
| 22 |
from dataset import RealESRGANDataset, RealESRGANDegrader # official
|
| 23 |
from train_lora import highfreq_mask # 复用小波高频 mask
|
|
@@ -44,6 +45,8 @@ def main():
|
|
| 44 |
ap.add_argument("--skip_ram", action="store_true")
|
| 45 |
ap.add_argument("--bf16", action="store_true", default=True); ap.add_argument("--no_bf16", dest="bf16", action="store_false")
|
| 46 |
ap.add_argument("--seed", type=int, default=123); ap.add_argument("--num_workers", type=int, default=8)
|
|
|
|
|
|
|
| 47 |
args = ap.parse_args()
|
| 48 |
random.seed(args.seed); torch.manual_seed(args.seed)
|
| 49 |
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
@@ -147,7 +150,9 @@ def main():
|
|
| 147 |
for p in dists_fn.parameters(): p.requires_grad_(False)
|
| 148 |
except Exception as e: log(f"[warn] pyiqa dists 不可用 {e}")
|
| 149 |
|
| 150 |
-
dl_iter = iter(dataloader); step = 0
|
|
|
|
|
|
|
| 151 |
opt_s.zero_grad(set_to_none=True); opt_D.zero_grad(set_to_none=True)
|
| 152 |
loss_D = torch.tensor(0.0, device=device)
|
| 153 |
while step < args.steps:
|
|
@@ -157,6 +162,11 @@ def main():
|
|
| 157 |
dl_iter = iter(dataloader); continue
|
| 158 |
LR, HR = degrader.degrade(batch) # [0,1] range
|
| 159 |
LR, HR = LR.to(device), HR.to(device)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 160 |
with torch.no_grad(), amp:
|
| 161 |
tag = DAPE.generate_tag(ram_tf(LR))[0]
|
| 162 |
text_input = tokenizer(tag, max_length=tokenizer.model_max_length,
|
|
@@ -174,6 +184,11 @@ def main():
|
|
| 174 |
z0_t = decoder.conv_in(z0_t); z0_t = decoder.mid_block(z0_t)
|
| 175 |
z0_gt = vae.post_quant_conv(HR_lat)
|
| 176 |
z0_gt = decoder.conv_in(z0_gt); z0_gt = decoder.mid_block(z0_gt)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 177 |
with amp:
|
| 178 |
z0_s = net(LRn)
|
| 179 |
tB = torch.full((z0_s.shape[0],), 999, dtype=torch.long, device=device)
|
|
@@ -185,6 +200,7 @@ def main():
|
|
| 185 |
loss_distil = (m * (z0_s - z0_t).abs()).mean()
|
| 186 |
loss_adv = (m * F.softplus(-d_s)).mean()
|
| 187 |
loss = args.w_distil * loss_distil + args.w_adv * loss_adv
|
|
|
|
| 188 |
if args.w_lpips + args.w_dists + args.w_color > 0:
|
| 189 |
img_s = tail(z0_s)
|
| 190 |
if lpips_fn is not None and args.w_lpips > 0:
|
|
@@ -196,19 +212,39 @@ def main():
|
|
| 196 |
loss = loss + args.w_color * (
|
| 197 |
(img_s.float().mean(dim=(2, 3)) - HRn.float().mean(dim=(2, 3))).abs().mean()
|
| 198 |
+ (img_s.float().std(dim=(2, 3)) - HRn.float().std(dim=(2, 3))).abs().mean())
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 199 |
(loss / args.grad_accum).backward()
|
| 200 |
if (step + 1) % args.grad_accum == 0:
|
| 201 |
-
|
|
|
|
|
|
|
|
|
|
| 202 |
with amp:
|
| 203 |
tB = torch.full((z0_s.shape[0],), 999, dtype=torch.long, device=device)
|
| 204 |
pred_real = unet_D(z0_gt.detach(), tB, encoder_hidden_states=enc, return_dict=False)[0]
|
| 205 |
pred_fake = unet_D(z0_s.detach(), tB, encoder_hidden_states=enc, return_dict=False)[0]
|
| 206 |
loss_D = F.softplus(pred_fake).mean() + F.softplus(-pred_real).mean()
|
| 207 |
(loss_D / args.grad_accum).backward()
|
| 208 |
-
|
|
|
|
|
|
|
|
|
|
| 209 |
if (step + 1) % 50 == 0:
|
| 210 |
log(f"step {step+1}/{args.steps} g {loss.item():.4f} distil {loss_distil.item():.4f} "
|
| 211 |
f"adv {loss_adv.item():.4f} D {loss_D.item():.4f}")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 212 |
if (step + 1) % args.save_every == 0:
|
| 213 |
torch.save(net.state_dict(), os.path.join(args.out, f"net_params_{step+1}.pkl"))
|
| 214 |
torch.save(tail.state_dict(), os.path.join(args.out, f"tail_params_{step+1}.pkl"))
|
|
|
|
| 17 |
|
| 18 |
from common import (ensure_official, load_diffusers_sd, load_pruned_decoder,
|
| 19 |
build_net, build_discriminator_unet, load_osediff_teacher,
|
| 20 |
+
load_gdpo_teacher, count_params, is_finite, check_tensor,
|
| 21 |
+
clip_and_check_grads, preview_grid)
|
| 22 |
ensure_official()
|
| 23 |
from dataset import RealESRGANDataset, RealESRGANDegrader # official
|
| 24 |
from train_lora import highfreq_mask # 复用小波高频 mask
|
|
|
|
| 45 |
ap.add_argument("--skip_ram", action="store_true")
|
| 46 |
ap.add_argument("--bf16", action="store_true", default=True); ap.add_argument("--no_bf16", dest="bf16", action="store_false")
|
| 47 |
ap.add_argument("--seed", type=int, default=123); ap.add_argument("--num_workers", type=int, default=8)
|
| 48 |
+
ap.add_argument("--clip_grad", type=float, default=1.0, help="??/?????????; 0=??")
|
| 49 |
+
ap.add_argument("--vis_every", type=int, default=500, help="? N ?? ??RGB/HR ??")
|
| 50 |
args = ap.parse_args()
|
| 51 |
random.seed(args.seed); torch.manual_seed(args.seed)
|
| 52 |
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
|
|
| 150 |
for p in dists_fn.parameters(): p.requires_grad_(False)
|
| 151 |
except Exception as e: log(f"[warn] pyiqa dists 不可用 {e}")
|
| 152 |
|
| 153 |
+
dl_iter = iter(dataloader); step = 0; skip_streak = 0
|
| 154 |
+
params_s = list(net.parameters()) + list(tail.parameters())
|
| 155 |
+
params_D = [p for p in unet_D.parameters() if p.requires_grad]
|
| 156 |
opt_s.zero_grad(set_to_none=True); opt_D.zero_grad(set_to_none=True)
|
| 157 |
loss_D = torch.tensor(0.0, device=device)
|
| 158 |
while step < args.steps:
|
|
|
|
| 162 |
dl_iter = iter(dataloader); continue
|
| 163 |
LR, HR = degrader.degrade(batch) # [0,1] range
|
| 164 |
LR, HR = LR.to(device), HR.to(device)
|
| 165 |
+
if check_tensor(LR, "lr", log) or check_tensor(HR, "hr", log):
|
| 166 |
+
skip_streak += 1
|
| 167 |
+
if skip_streak > 20:
|
| 168 |
+
log("[anomaly] too many bad batches, abort"); break
|
| 169 |
+
continue
|
| 170 |
with torch.no_grad(), amp:
|
| 171 |
tag = DAPE.generate_tag(ram_tf(LR))[0]
|
| 172 |
text_input = tokenizer(tag, max_length=tokenizer.model_max_length,
|
|
|
|
| 184 |
z0_t = decoder.conv_in(z0_t); z0_t = decoder.mid_block(z0_t)
|
| 185 |
z0_gt = vae.post_quant_conv(HR_lat)
|
| 186 |
z0_gt = decoder.conv_in(z0_gt); z0_gt = decoder.mid_block(z0_gt)
|
| 187 |
+
if check_tensor(z0_t, "z0_teacher", log) or check_tensor(z0_gt, "z0_gt", log):
|
| 188 |
+
skip_streak += 1
|
| 189 |
+
if skip_streak > 20:
|
| 190 |
+
log("[anomaly] too many bad teachers, abort"); break
|
| 191 |
+
continue
|
| 192 |
with amp:
|
| 193 |
z0_s = net(LRn)
|
| 194 |
tB = torch.full((z0_s.shape[0],), 999, dtype=torch.long, device=device)
|
|
|
|
| 200 |
loss_distil = (m * (z0_s - z0_t).abs()).mean()
|
| 201 |
loss_adv = (m * F.softplus(-d_s)).mean()
|
| 202 |
loss = args.w_distil * loss_distil + args.w_adv * loss_adv
|
| 203 |
+
img_s = None
|
| 204 |
if args.w_lpips + args.w_dists + args.w_color > 0:
|
| 205 |
img_s = tail(z0_s)
|
| 206 |
if lpips_fn is not None and args.w_lpips > 0:
|
|
|
|
| 212 |
loss = loss + args.w_color * (
|
| 213 |
(img_s.float().mean(dim=(2, 3)) - HRn.float().mean(dim=(2, 3))).abs().mean()
|
| 214 |
+ (img_s.float().std(dim=(2, 3)) - HRn.float().std(dim=(2, 3))).abs().mean())
|
| 215 |
+
if check_tensor(z0_s, "z0_student", log) or not is_finite(loss):
|
| 216 |
+
log(f"[anomaly] student/loss NaN at step {step+1}; skip step")
|
| 217 |
+
skip_streak += 1
|
| 218 |
+
opt_s.zero_grad(set_to_none=True); opt_D.zero_grad(set_to_none=True)
|
| 219 |
+
if skip_streak > 20:
|
| 220 |
+
log("[anomaly] too many bad steps, abort"); break
|
| 221 |
+
continue
|
| 222 |
+
skip_streak = 0
|
| 223 |
(loss / args.grad_accum).backward()
|
| 224 |
if (step + 1) % args.grad_accum == 0:
|
| 225 |
+
if clip_and_check_grads(params_s, args.clip_grad, log):
|
| 226 |
+
opt_s.zero_grad(set_to_none=True)
|
| 227 |
+
else:
|
| 228 |
+
opt_s.step(); opt_s.zero_grad(set_to_none=True); sched.step()
|
| 229 |
with amp:
|
| 230 |
tB = torch.full((z0_s.shape[0],), 999, dtype=torch.long, device=device)
|
| 231 |
pred_real = unet_D(z0_gt.detach(), tB, encoder_hidden_states=enc, return_dict=False)[0]
|
| 232 |
pred_fake = unet_D(z0_s.detach(), tB, encoder_hidden_states=enc, return_dict=False)[0]
|
| 233 |
loss_D = F.softplus(pred_fake).mean() + F.softplus(-pred_real).mean()
|
| 234 |
(loss_D / args.grad_accum).backward()
|
| 235 |
+
if clip_and_check_grads(params_D, args.clip_grad, log):
|
| 236 |
+
opt_D.zero_grad(set_to_none=True)
|
| 237 |
+
else:
|
| 238 |
+
opt_D.step(); opt_D.zero_grad(set_to_none=True)
|
| 239 |
if (step + 1) % 50 == 0:
|
| 240 |
log(f"step {step+1}/{args.steps} g {loss.item():.4f} distil {loss_distil.item():.4f} "
|
| 241 |
f"adv {loss_adv.item():.4f} D {loss_D.item():.4f}")
|
| 242 |
+
if args.vis_every > 0 and (step + 1) % args.vis_every == 0 and img_s is not None:
|
| 243 |
+
try:
|
| 244 |
+
preview_grid([LRn.float()[:1], img_s.float()[:1], HRn.float()[:1]],
|
| 245 |
+
os.path.join(args.log_dir, f"step_{step+1:06d}.png"))
|
| 246 |
+
except Exception as e:
|
| 247 |
+
log(f"[warn] preview fail: {e}")
|
| 248 |
if (step + 1) % args.save_every == 0:
|
| 249 |
torch.save(net.state_dict(), os.path.join(args.out, f"net_params_{step+1}.pkl"))
|
| 250 |
torch.save(tail.state_dict(), os.path.join(args.out, f"tail_params_{step+1}.pkl"))
|