|
Download RUNBOOK_A100.md from XenderYang/CSIGv3_train_script: direct link, hf CLI and curl.
- Browser
- Download file 3.72 kB
-
https://huggingface.co/XenderYang/CSIGv3_train_script/resolve/main/RUNBOOK_A100.md
- Command line
-
hf download hf://XenderYang/CSIGv3_train_script/RUNBOOK_A100.md
-
curl -L -o RUNBOOK_A100.md https://huggingface.co/XenderYang/CSIGv3_train_script/resolve/main/RUNBOOK_A100.md
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 冒烟 | |