XenderYang commited on
Commit
ca2409a
·
verified ·
1 Parent(s): 4811c23

fix smoke bugs + anomaly guards; runbook update

Browse files
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: 128
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.view(b * c, 1, ph, pw)
26
- kernel = kernel.view(1, 1, k, k)
27
- return F.conv2d(img, kernel, padding=0).view(b, c, h, w)
28
  else:
29
- img = img.view(1, b * c, ph, pw)
30
- kernel = kernel.view(b, 1, k, k).repeat(1, c, 1, 1).view(b * c, 1, k, k)
31
- return F.conv2d(img, kernel, groups=b * c).view(b, c, h, w)
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
- pipe = StableDiffusionPipeline.from_pretrained(model_id, torch_dtype=dtype).to(device)
 
 
 
 
 
 
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, cout = conv.in_channels, conv.out_channels
139
- self.lora_a = nn.Parameter(torch.zeros(cin, self.r))
140
- self.lora_b = nn.Parameter(torch.zeros(self.r, cout))
 
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 scaler is not None:
206
- scaler.step(optimizer); scaler.update()
207
  else:
208
- optimizer.step()
209
- optimizer.zero_grad(set_to_none=True)
210
- sched.step()
 
 
 
 
 
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
- opt_s.step(); opt_s.zero_grad(set_to_none=True); sched.step()
 
 
 
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
- opt_D.step(); opt_D.zero_grad(set_to_none=True)
 
 
 
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"))