CSIGv3_train_script / RUNBOOK_A100.md
XenderYang's picture
fix smoke bugs + anomaly guards; runbook update
ca2409a verified
|
Raw History Blame Contribute Delete
3.72 kB
# CSIGv3 AdcSR 训练脚本包 (A100 runbook)
## 0) 数据
```bash
hf download XenderYang/CSIGv2_SR-Train --repo-type dataset --local-dir ./data
cd data && for f in *.tar; do tar -xf "$f"; done && cd ..
# 重新生成清单(manifest 内为 Windows 绝对路径, 不可直接用)
python tools/build_manifest.py --hr_dirs data/patches \
--real_lr data/real_pairs/RealSR/train_lr --real_hr data/real_pairs/RealSR/train_hr \
--out data/manifest_train.json
python tools/build_manifest.py --hr_dirs data/patches_val \
--real_lr data/real_pairs/RealSR/test_lr --real_hr data/real_pairs/RealSR/test_hr \
--out data/manifest_val.json
```
## 1) 环境
```bash
conda create -n AdcSR python=3.10 -y && conda activate AdcSR
pip install -r requirements.txt
# 数据仍以本仓库 scripts/setup_env.sh 指引; 权重:
bash scripts/download_weights.sh # 官方 AdcSR 5件 + SD2.1(fp16) + GDPO(可选)
python scripts/env_check.py # GPU/驱动/关键 import/单迭代自检
```
## 2) S0 冒烟(每个配置 2 迭代, 必须先过)
```bash
bash scripts/run_smoke.sh
# 或手动:
python src/train_lora.py --config configs/config_smoke.yml --manifest data/manifest_train.json --real_prob 0.0 --steps 2 --batch_size 1 --grad_accum 1 --init_net weight/net_params_200.pkl --out weight/smoke --log_dir logs/smoke_lora
python src/train_stage2.py --config configs/config_smoke.yml --steps 2 --batch_size 1 --grad_accum 1 --teacher osediff --skip_ram --init_net weight/net_params_200.pkl --out weight/smoke --log_dir logs/smoke_s2
# GDPO 教师冒烟(先 probe)
python scripts/probe_gdpo.py weight/gdpo
python src/train_stage2.py --config configs/config_smoke.yml --steps 2 --batch_size 1 --teacher gdpo --gdpo_dir weight/gdpo --init_net weight/net_params_200.pkl --out weight/smoke --log_dir logs/smoke_s2_gdpo
```
## 3) 正式训练
```bash
bash scripts/run_stage1.sh # S1 LoRA 域适配 (20k step)
# eval 选 best -> weight/s1/net_params_BEST.pkl
bash scripts/run_stage2.sh gdpo # S2 GDPO 蒸馏(回退 osediff) (30k step)
bash scripts/run_stage3.sh # S3 域微调 (20k step)
```
## 4) 验证/导出/提交
见 src/eval_val.py / export_jit.py / inference_4k.py / check_submission.py
(赛题验证 3 对: data/val_official; 打分测试集: 见工程 README)
## 注意
- 训练脚本依赖 torchvision(官方 bsr/degradations.py import)。A100(Linux+torchvision)正常;
本机 Windows 沙箱若 torchvision 报 `operator torchvision::nms does not exist`, 请在自己终端(非沙箱)运行。
- 数据清单中的真实退化对仅 RealSR 400/99; 其余母本在 data/clean_all, data/clean。
## 本地冒烟结果(2026-09-06, conda tvtest CPU)
- conda 环境 torch 2.5.1 + torchvision 0.20.1 可用(需 PYTHONNOUSERSITE=1 屏蔽用户目录污染)
- 修复清单(已同步到本仓库):
* common.py: LoRAConv2d cin/cout 未定义; LoRA 仅注入 Linear/1x1Conv(3x3 分解不匹配);
整体冻结仅训 LoRA; assemble_full_student 展开 up_blocks(ModuleList 不能进 Sequential);
load_diffusers_sd 自动 variant=fp16; 新增 is_finite/check_tensor/clip_and_check_grads/EMA/preview_grid
* train_lora.py/train_stage2.py: NaN/Inf 检测(输入/输出/损失/梯度)、梯度裁剪(clip_grad 1.0)、
skip-step 连续异常>20 中止、EMA(默认0.999)、周期预览图(vis_every)、num_workers 冒烟用0
* official/bsr/utils/img_process_util.py: filter2D .view->.reshape(非连续张量)
* configs/config_smoke.yml: gt_size 128->512(Net 输入需 LR128=HR512/4)
- 冒烟结果: 退化管线 OK(HR512->LR128, 数值正常, 预览 logs/smoke_degrade_preview.png);
S1 LoRA 2步 OK(checkpoint/lora_ema/预览); S2 本地 CPU 内存不足被系统终止 -> 在 A100 跑 S2/GDPO 冒烟