diff --git a/NOTES.md b/NOTES.md new file mode 100644 index 0000000000000000000000000000000000000000..4e357e4f3cf9c8c64281e85f3dc6a9b14414b3f1 --- /dev/null +++ b/NOTES.md @@ -0,0 +1,107 @@ +# 实施与 S0 冒烟要点(给下一执行者/服务器环境) + +## 已实现(本地已验证) +- 工程脚手架 + 官方 clone(official/) + git 已提交(5a1afb7) +- tools/*.py 全部 CPU 可跑: clean_images(阈值/去重/可复制到 clean) / tag_categories(7类启发式) / + build_manifest / split_val(70组隔离) / build_patches(>=20k patch) / download_batch(编排) + 已用合成小图冒烟: 模糊/全黑/过曝/纯色正确剔除, 打标与 patch 构建正常。 +- src/*.py 与 scripts/*.py 语法编译通过(py_compile)。 + +## 必须上 A100 用真实权重验证的点(S0, 1 迭代/配置) +1. scripts/setup_env.sh + scripts/download_weights.sh(Plan A) 后运行 scripts/env_check.py。 +2. scripts/run_smoke.sh: train_lora 与 train_stage2 各 1 迭代; 重点核对: + - Net+halfDecoder 全链装配 forward 形状: LR[1,3,128,128] -> RGB[1,3,512,512] + (若 shape 报错, 以 official/test.py 官方装配为唯一基准修正 src/common.build_net/assemble_full_student) + - teacher z0 前向与判别器 dtype/shape 匹配(bf16 + autocast) + - loss 有限值且可保存 checkpoint、可续训 +3. 基线复现: 用官方 net_params_200.pkl + official/test.py 在 RealSR 复现论文量级指标; + 并在 A100 记录 OSEDiff 512 fp16 时延(基准分母)。 +4. GDPO 教师: 先 python scripts/probe_gdpo.py <权重路径> 确认键格式; 若格式不兼容, + 按 probe 输出在 src/common.load_gdpo_teacher 补映射, 否则回退 --teacher osediff(默认已支持)。 + +## 执行顺序(每阶段 gate 见方案 §6) +1. tools 数据构建 B1-B4 -> clean -> tag -> patches>=20k -> manifest -> split_val(70) +2. pack_upload.ps1 上传代码+数据; 服务器下载权重 +3. S0 smoke -> S1(LoRA) -> eval_val 选 best -> S2(GDPO蒸馏) -> eval_val 选 best -> S3 +4. 每次选点: 用 src/eval_val.py 输出 proxy, 把最优权重软链/复制为 + weight/s1/net_params_BEST.pkl (run_stage2 引用) / weight/s2/net_params_BEST.pkl (run_stage3 引用) +5. 速度: src/latency_test.py; 出图: src/inference_4k.py; 导出: src/export_jit.py; 提交包检查: src/check_submission.py + +## 已知待办/说明 +- 训练脚本仅在服务器有 torch/权重时可运行; 本地 3050 4G 只做代码与数据开发。 +- run_stage1/2/3.sh 里的 net_params_BEST.pkl 为占位名, 按实际 eval 结果更新。 +- inference_4k/export_jit 采用“512窗->128->官方4x链->512”的 x1 包装语义; + 若官方评测机 runner 口径不同, 以官方样例 runner 实测为准调整 export_jit.SR512。 +- 类别打标为启发式(文件名关键词+亮度/绿色占比), 抽样人工复核后再上传训练。 + +## 试点批次进度 (B3 NKUSR8K 首 100 张, 2026-09-05) +- raw 100 张(~1.9GB) -> clean 90 张(10 张 aHash 重复剔除; 8K 图质量均达标) -> 已删 raw 释放空间 +- 分类: city_view 41 / night_building 26 / plant 23 (skyscraper/shop_street/residential 为 0, 需 4KLSDB/RealSR 等源补充) +- patches: 2998 张 512x512 (city_view 1366 / night_building 866 / plant 766), manifest: data/manifest_train.json +- 人工复核图: data/review/NKUSR8K/*.jpg (每类 20 张拼图) +- 注意: synthetic patch 的 val 划分须按"源图"分组(文件名前缀)而非单 patch, 防同源泄露; 最终 val 主要来自 real pairs。 + +## 数据集构建进度 (2026-09-05 晚 更新, 本批次完成) +- NKUSR8K 全量 1000 张: 下载 900(跳过试点100) -> 清洗 kept 639 (dup 259/过曝 2) -> 已删 raw(17.9GB 释放) +- 合并 clean_all/NKUSR8K = 729 张(试点90 + 全量639, 4096 长边 jpg, 文件名前缀 NKUSR8K_) +- 启发式分类: night_building 231 / plant 255 / city_view 243 (skyscraper/shop/residential 启发式分不出, 无碍训练) +- 源级 val 划分: 70 源留作 val (split_val_sources.json), train 源 659 +- patches: data/patches 23,999 (city 7866/night 7684/plant 8449) + data/patches_val 560 (val 源) +- manifest_train.json: synthetic 23,999 + RealSR 训练对 400; manifest_val.json: synthetic 560 + RealSR 测试对 99 +- RealSR 对清单: clean_pairs 配对命名 Canon_001_LR4.png <-> Canon_001_HR.png (build_manifest 已按 scene key 匹配) +- 人工复核拼图: data/review/NKUSR8K_all/{city_view,night_building,plant}.jpg + +## 待办 +- [ ] DRealSR 训练集: 需人工(百度网盘 osiy / Google Drive 文件夹), 见 data/SOURCES.md; 到手后走 clean_pairs 入 real_pairs +- [ ] 4KLSDB: 默认跳过(稀疏类抽取不划算, 结论见 SOURCES.md); 若需补量可用 tools 的 RangeReader 思路实现 extract_klsdb_rg.py +- [ ] 打包上传: scripts/pack_upload.ps1 (需密码), 服务器 setup_env.sh + download_weights.sh + S0 smoke -> S1/S2/S3 +- [ ] 本地清理: 确认后删除 data/clean/NKUSR8K* / data/clean_all(已并入 patches, 仅留存备查) + +## 2026-09-05 深夜: CLIP/4KLSDB/稀疏类抓取 进展 +- 修正: NKUSR8K 无夜景(CLIP 确认 night_building=0); tag_categories 夜景仅允许文件名关键词, 不再按暗度推断 +- CLIP(openai/clip-vit-base-patch32, CPU, transformers5.x 需 .pooler_output) 打标 NKUSR8K 729: + skyscraper 204 / plant 354 / residential 63 / other_building 41 / storefront 22 / mall 21 / city_view 5, reject 19 +- 分类升级为 8 类: night_building / skyscraper / storefront(店铺门牌) / mall(商场外拍) / plant / + residential / city_view / other_building (tools: clip_filter.py, tag_categories.py) +- 4KLSDB 抽取管线落地(全部本地可跑): + klsdb_index.py(206 分片 id->shard/rg 索引, ~27min) + klsdb_extract.py(rank 按 caption 评分排序 + extract 按 row-group 拉 hr) + 第一批: 18 row-group, 11.6GB 流量 -> 1060 张 4K 原始 -> clean 802 张(data/clean/4KLSDB_b1) +- 稀疏类原图 URL 抓取(fetch_klsdb_urls.py, 成功率 ~96%): night_building 575 / mall 355 / + storefront 501 / residential 495 / skyscraper 486 -> 后台逐类清洗中(data/clean/4KLSDB_url/) +- DRealSR: 用户给的 GDrive 文件夹(1tP5m4k1...)只含 Test_x2/x3/x4.zip(测试集); 训练集在百度网盘 osiy(需登录) +- SSDLite: 本地 Windows torchvision CPU wheel ABI 不兼容(2.7/2.6+cpu 均报 torchvision::nms), 已放弃本地; + 方案: 本地用 CLIP; SSDLite person/物过滤代码留到 A100(Linux+torchvision 正常) 或后续换环境再启用 +- 待办: URL 清洗完成 -> CLIP 复核(去 people/indoor/other) -> 全池 pHash 去重 -> 合并 clean_all(带源前缀) + -> 8 类平衡报表 -> sample_review 拼图人工复核 -> build_patches(>=? 目标 5000-7000 HR 源后 50k+ patch) + -> 再跑 4KLSDB b2/b3 补量到总池 5000-7000 + +## 2026-09-06 凌晨: LIU4K v2 引入 + 策略更新 +- 4KLSDB 政策: 建筑/风景/花草均可入池(增强鲁棒), 夜景适量即可(不堆量) +- LIU4K v2 (GDrive 1FtVQtY2t_ecuy_gzJqZ-CatqrJBAdq_d): 4 组多段zip(Animal/Building/Street/Mountain), + 只取 Building/Street/Mountain; Building = 400 张高分辨率 PNG(3992x2242~6208x4139, Pexels 建筑图, 10.2GB) + 解压用 NVIDIA App 7-Zip 22.01(自带 7z.exe); 旧 7za 9.20 不支持 spanned zip +- DRealSR: 放弃训练(用户确认); SSDLite: 取消(用户确认); 过滤以 CLIP 为主 +- URL 稀疏类清洗后: night_building 409 / storefront 446 / mall 248 / residential 339 / skyscraper 331 (共 1773) +- 4KLSDB b1 CLIP 后 accept 604 (plant172/other_building171/residential124/city_view66/storefront35/skyscraper29/mall6/night1), reject 198 +- NKUSR8K CLIP accept 710 (skyscraper204/plant354/residential63/other_building41/storefront22/mall21/city_view5), reject 19 +- 新工具: merge_reports.py(多源合并/剔 reject/去重) + download_gdrive_parts.py(分卷断点下载) + +## 2026-09-06 凌晨2: 四源合并池(CLIP accept + URL文件夹标签 + reject剔除) +- 池: logs/cat_pool_all.json, 共 3072 张清洗后 HR 原图 + night_building 332 / skyscraper 575 / storefront 445 / mall 252 / plant 540 / residential 526 / + city_view 102 / other_building 300 +- 人工复核拼图: data/review/pool_all/{8类}.jpg (每类 20 张) +- LIU4K Building: 400 PNG(10.2GB) -> clean 320 -> jpg(clean_all/LIU4K) -> CLIP accept 235(reject 85 室内/人像) +- 剩余待补: city_view(102, 目标>=400), mall(252); 总池 3072 -> 目标 5000-7000 +- 下一批: 4KLSDB b2/b3(城景/商场/其他建筑加量) + LIU4K Street/Mountain(城市街道/山地风景) + +## 2026-09-06 凌晨: 收尾流水线启动(finalize_all.ps1 后台自动链) +- 补量URL抓取: plant 790+383 / city_view 398+240 / other_building 554+259 / storefront 129 / mall 160 (raw) + (skyscraper/residential 候选枯竭未新增; fetch 去重由 jsonl 断点保证) +- 4KLSDB b2 事故: klsdb_extract 未跳过已用 row-group -> b2 与 b1 完全重复(802/802), 已删 b2 目录; + 已修复 extract 记录 used_rgs.json, 后续 b3 不再重复 +- clean_url_vol.ps1 正清洗 5 类 raw(plant/city_view/other_building/storefront/mall) +- finalize_all.ps1(等清洗完): clip_url3(全URL目录) -> merge_reports(cat_pool_final, id去重) + -> split_pool(70源val) -> build_patches2(train 45k, val 560) -> build_manifest(train/val + RealSR对) +- 产出: logs/cat_pool_final.json / cat_train.json / cat_val.json / data/patches / data/patches_val / + data/manifest_train.json / data/manifest_val.json / data/README.md diff --git a/README.md b/README.md new file mode 100644 index 0000000000000000000000000000000000000000..d7f69dfc6571ed9eb31cbf2fbd067f641146734d --- /dev/null +++ b/README.md @@ -0,0 +1,36 @@ +# AdcSR 竞赛冲分工程(CSIG 2026 Camera学术之星 赛道二) + +以官方 Guaishou74851/AdcSR (CVPR 2025) 为基座的分阶段训练工程: +S1 LoRA/QLoRA 像素域适配 -> S2 GDPO-SR 教师升级蒸馏 -> S3 域微调 -> 速度优化 -> 赛题提交包。 + +## 目录 +- `official/` 官方仓库原样 clone(只读参照;训练脚本通过 sys.path 引用其 model/dataset/utils/ram) +- `src/` 训练/蒸馏/导出/推理/评测 +- `tools/` 数据下载、清洗、打标、建 patch、清单、val 划分(纯 CPU) +- `configs/` 各阶段 yaml +- `scripts/` 服务器环境、权重、运行、打包上传、归档 +- `weight/ models/ logs/ data/` 运行产物(不入库) + +## 本地 vs 服务器 +- 本地(RTX 3050 4G / F 盘 ~44G):代码开发、CPU 数据清洗、gt<=128 的 1 迭代冒烟。 +- aspire2a(单卡 A100 40G):S1/S2/S3 全量训练、评测、导出、100 张 4K 出图。 + +## 快速开始(服务器) +```bash +bash scripts/setup_env.sh # conda + 依赖 +bash scripts/download_weights.sh # Plan A: 服务器拉权重(AdcSR 5件套+SD2.1+GDPO) +python scripts/env_check.py +bash scripts/run_smoke.sh # S0 冒烟(每个训练配置 1 迭代) +bash scripts/run_stage1.sh # S1 LoRA +bash scripts/run_stage2.sh # S2 GDPO 教师蒸馏 +bash scripts/run_stage3.sh # S3 域微调 +python src/eval_val.py --weights # val 评测 + 代理分 +python src/export_jit.py --weights --out model_dir +python src/inference_4k.py --lr_dir --weights --out output_dir +python src/check_submission.py --zip +``` + +## 训练/评测口径(与官方一致) +- 输入 LR 128x128 (scale=4, gt 512) -> 输出 512x512 RGB;官方 runner 用 512x512 包装/分块。 +- val 代理分 proxy = 0.25*CLIPIQA + 0.2*(MUSIQ/100) + 0.15*MANIQA + 0.15*(1-NIQE/10) + 0.25*(1-LPIPS) +- 保真护栏:PSNR 不低于官方基线 0.5dB,SSIM 不低于 0.01。 diff --git a/RUNBOOK_A100.md b/RUNBOOK_A100.md new file mode 100644 index 0000000000000000000000000000000000000000..9e29d61bdfc5f3b905a0848eee99e7c19d35bb9c --- /dev/null +++ b/RUNBOOK_A100.md @@ -0,0 +1,51 @@ +# 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。 diff --git a/configs/config_base.yml b/configs/config_base.yml new file mode 100644 index 0000000000000000000000000000000000000000..2287bc8dd2d667ecd89369e87add18c4fa676a9e --- /dev/null +++ b/configs/config_base.yml @@ -0,0 +1,35 @@ +# base degradation config (override dataroot_gt/gt_size/iter_num per stage) +# Real-ESRGAN style degradation (shared by all stages) +scale: 4 +resize_prob: [0.2, 0.7, 0.1] +resize_range: [0.3, 1.5] +gaussian_noise_prob: 0.5 +noise_range: [1, 15] +poisson_scale_range: [0.05, 2.0] +gray_noise_prob: 0.4 +jpeg_range: [60, 95] +second_blur_prob: 0.5 +resize_prob2: [0.3, 0.4, 0.3] +resize_range2: [0.6, 1.2] +gaussian_noise_prob2: 0.5 +noise_range2: [1, 12] +poisson_scale_range2: [0.05, 1.0] +gray_noise_prob2: 0.4 +jpeg_range2: [60, 100] +blur_kernel_size: 21 +kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] +kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] +sinc_prob: 0.1 +blur_sigma: [0.2, 1.5] +betag_range: [0.5, 2.0] +betap_range: [1, 1.5] +blur_kernel_size2: 11 +kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] +kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] +sinc_prob2: 0.1 +blur_sigma2: [0.2, 1.0] +betag_range2: [0.5, 2.0] +betap_range2: [1, 1.5] +final_sinc_prob: 0.8 +use_hflip: True +use_rot: False diff --git a/configs/config_s1_lora.yml b/configs/config_s1_lora.yml new file mode 100644 index 0000000000000000000000000000000000000000..b169ff93a6d2ead328122c3b529a100aceb026e2 --- /dev/null +++ b/configs/config_s1_lora.yml @@ -0,0 +1,37 @@ +dataroot_gt: data/patches +gt_size: 512 +iter_num: 1000 +# Real-ESRGAN style degradation (shared by all stages) +scale: 4 +resize_prob: [0.2, 0.7, 0.1] +resize_range: [0.3, 1.5] +gaussian_noise_prob: 0.5 +noise_range: [1, 15] +poisson_scale_range: [0.05, 2.0] +gray_noise_prob: 0.4 +jpeg_range: [60, 95] +second_blur_prob: 0.5 +resize_prob2: [0.3, 0.4, 0.3] +resize_range2: [0.6, 1.2] +gaussian_noise_prob2: 0.5 +noise_range2: [1, 12] +poisson_scale_range2: [0.05, 1.0] +gray_noise_prob2: 0.4 +jpeg_range2: [60, 100] +blur_kernel_size: 21 +kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] +kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] +sinc_prob: 0.1 +blur_sigma: [0.2, 1.5] +betag_range: [0.5, 2.0] +betap_range: [1, 1.5] +blur_kernel_size2: 11 +kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] +kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] +sinc_prob2: 0.1 +blur_sigma2: [0.2, 1.0] +betag_range2: [0.5, 2.0] +betap_range2: [1, 1.5] +final_sinc_prob: 0.8 +use_hflip: True +use_rot: False diff --git a/configs/config_s2_distill.yml b/configs/config_s2_distill.yml new file mode 100644 index 0000000000000000000000000000000000000000..b169ff93a6d2ead328122c3b529a100aceb026e2 --- /dev/null +++ b/configs/config_s2_distill.yml @@ -0,0 +1,37 @@ +dataroot_gt: data/patches +gt_size: 512 +iter_num: 1000 +# Real-ESRGAN style degradation (shared by all stages) +scale: 4 +resize_prob: [0.2, 0.7, 0.1] +resize_range: [0.3, 1.5] +gaussian_noise_prob: 0.5 +noise_range: [1, 15] +poisson_scale_range: [0.05, 2.0] +gray_noise_prob: 0.4 +jpeg_range: [60, 95] +second_blur_prob: 0.5 +resize_prob2: [0.3, 0.4, 0.3] +resize_range2: [0.6, 1.2] +gaussian_noise_prob2: 0.5 +noise_range2: [1, 12] +poisson_scale_range2: [0.05, 1.0] +gray_noise_prob2: 0.4 +jpeg_range2: [60, 100] +blur_kernel_size: 21 +kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] +kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] +sinc_prob: 0.1 +blur_sigma: [0.2, 1.5] +betag_range: [0.5, 2.0] +betap_range: [1, 1.5] +blur_kernel_size2: 11 +kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] +kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] +sinc_prob2: 0.1 +blur_sigma2: [0.2, 1.0] +betag_range2: [0.5, 2.0] +betap_range2: [1, 1.5] +final_sinc_prob: 0.8 +use_hflip: True +use_rot: False diff --git a/configs/config_s3_ft.yml b/configs/config_s3_ft.yml new file mode 100644 index 0000000000000000000000000000000000000000..b169ff93a6d2ead328122c3b529a100aceb026e2 --- /dev/null +++ b/configs/config_s3_ft.yml @@ -0,0 +1,37 @@ +dataroot_gt: data/patches +gt_size: 512 +iter_num: 1000 +# Real-ESRGAN style degradation (shared by all stages) +scale: 4 +resize_prob: [0.2, 0.7, 0.1] +resize_range: [0.3, 1.5] +gaussian_noise_prob: 0.5 +noise_range: [1, 15] +poisson_scale_range: [0.05, 2.0] +gray_noise_prob: 0.4 +jpeg_range: [60, 95] +second_blur_prob: 0.5 +resize_prob2: [0.3, 0.4, 0.3] +resize_range2: [0.6, 1.2] +gaussian_noise_prob2: 0.5 +noise_range2: [1, 12] +poisson_scale_range2: [0.05, 1.0] +gray_noise_prob2: 0.4 +jpeg_range2: [60, 100] +blur_kernel_size: 21 +kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] +kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] +sinc_prob: 0.1 +blur_sigma: [0.2, 1.5] +betag_range: [0.5, 2.0] +betap_range: [1, 1.5] +blur_kernel_size2: 11 +kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] +kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] +sinc_prob2: 0.1 +blur_sigma2: [0.2, 1.0] +betag_range2: [0.5, 2.0] +betap_range2: [1, 1.5] +final_sinc_prob: 0.8 +use_hflip: True +use_rot: False diff --git a/configs/config_smoke.yml b/configs/config_smoke.yml new file mode 100644 index 0000000000000000000000000000000000000000..57413da4079a93d63ef5228bf533f3dbbe260ea5 --- /dev/null +++ b/configs/config_smoke.yml @@ -0,0 +1,37 @@ +dataroot_gt: data/patches +gt_size: 128 +iter_num: 2 +# Real-ESRGAN style degradation (shared by all stages) +scale: 4 +resize_prob: [0.2, 0.7, 0.1] +resize_range: [0.3, 1.5] +gaussian_noise_prob: 0.5 +noise_range: [1, 15] +poisson_scale_range: [0.05, 2.0] +gray_noise_prob: 0.4 +jpeg_range: [60, 95] +second_blur_prob: 0.5 +resize_prob2: [0.3, 0.4, 0.3] +resize_range2: [0.6, 1.2] +gaussian_noise_prob2: 0.5 +noise_range2: [1, 12] +poisson_scale_range2: [0.05, 1.0] +gray_noise_prob2: 0.4 +jpeg_range2: [60, 100] +blur_kernel_size: 21 +kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] +kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] +sinc_prob: 0.1 +blur_sigma: [0.2, 1.5] +betag_range: [0.5, 2.0] +betap_range: [1, 1.5] +blur_kernel_size2: 11 +kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] +kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] +sinc_prob2: 0.1 +blur_sigma2: [0.2, 1.0] +betag_range2: [0.5, 2.0] +betap_range2: [1, 1.5] +final_sinc_prob: 0.8 +use_hflip: True +use_rot: False diff --git a/official/bsr/__pycache__/degradations.cpython-312.pyc b/official/bsr/__pycache__/degradations.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..2659824cdcd3f23b00cfbb73ebdd1a2cfb755eb5 Binary files /dev/null and b/official/bsr/__pycache__/degradations.cpython-312.pyc differ diff --git a/official/bsr/degradations.py b/official/bsr/degradations.py new file mode 100644 index 0000000000000000000000000000000000000000..f40d0fd77a7128202070a18d4c5195ebff25056a --- /dev/null +++ b/official/bsr/degradations.py @@ -0,0 +1,764 @@ +import cv2 +import math +import numpy as np +import random +import torch +from scipy import special +from scipy.stats import multivariate_normal +from torchvision.transforms._functional_tensor import rgb_to_grayscale + +# -------------------------------------------------------------------- # +# --------------------------- blur kernels --------------------------- # +# -------------------------------------------------------------------- # + + +# --------------------------- util functions --------------------------- # +def sigma_matrix2(sig_x, sig_y, theta): + """Calculate the rotated sigma matrix (two dimensional matrix). + + Args: + sig_x (float): + sig_y (float): + theta (float): Radian measurement. + + Returns: + ndarray: Rotated sigma matrix. + """ + d_matrix = np.array([[sig_x**2, 0], [0, sig_y**2]]) + u_matrix = np.array([[np.cos(theta), -np.sin(theta)], [np.sin(theta), np.cos(theta)]]) + return np.dot(u_matrix, np.dot(d_matrix, u_matrix.T)) + + +def mesh_grid(kernel_size): + """Generate the mesh grid, centering at zero. + + Args: + kernel_size (int): + + Returns: + xy (ndarray): with the shape (kernel_size, kernel_size, 2) + xx (ndarray): with the shape (kernel_size, kernel_size) + yy (ndarray): with the shape (kernel_size, kernel_size) + """ + ax = np.arange(-kernel_size // 2 + 1., kernel_size // 2 + 1.) + xx, yy = np.meshgrid(ax, ax) + xy = np.hstack((xx.reshape((kernel_size * kernel_size, 1)), yy.reshape(kernel_size * kernel_size, + 1))).reshape(kernel_size, kernel_size, 2) + return xy, xx, yy + + +def pdf2(sigma_matrix, grid): + """Calculate PDF of the bivariate Gaussian distribution. + + Args: + sigma_matrix (ndarray): with the shape (2, 2) + grid (ndarray): generated by :func:`mesh_grid`, + with the shape (K, K, 2), K is the kernel size. + + Returns: + kernel (ndarrray): un-normalized kernel. + """ + inverse_sigma = np.linalg.inv(sigma_matrix) + kernel = np.exp(-0.5 * np.sum(np.dot(grid, inverse_sigma) * grid, 2)) + return kernel + + +def cdf2(d_matrix, grid): + """Calculate the CDF of the standard bivariate Gaussian distribution. + Used in skewed Gaussian distribution. + + Args: + d_matrix (ndarrasy): skew matrix. + grid (ndarray): generated by :func:`mesh_grid`, + with the shape (K, K, 2), K is the kernel size. + + Returns: + cdf (ndarray): skewed cdf. + """ + rv = multivariate_normal([0, 0], [[1, 0], [0, 1]]) + grid = np.dot(grid, d_matrix) + cdf = rv.cdf(grid) + return cdf + + +def bivariate_Gaussian(kernel_size, sig_x, sig_y, theta, grid=None, isotropic=True): + """Generate a bivariate isotropic or anisotropic Gaussian kernel. + + In the isotropic mode, only `sig_x` is used. `sig_y` and `theta` is ignored. + + Args: + kernel_size (int): + sig_x (float): + sig_y (float): + theta (float): Radian measurement. + grid (ndarray, optional): generated by :func:`mesh_grid`, + with the shape (K, K, 2), K is the kernel size. Default: None + isotropic (bool): + + Returns: + kernel (ndarray): normalized kernel. + """ + if grid is None: + grid, _, _ = mesh_grid(kernel_size) + if isotropic: + sigma_matrix = np.array([[sig_x**2, 0], [0, sig_x**2]]) + else: + sigma_matrix = sigma_matrix2(sig_x, sig_y, theta) + kernel = pdf2(sigma_matrix, grid) + kernel = kernel / np.sum(kernel) + return kernel + + +def bivariate_generalized_Gaussian(kernel_size, sig_x, sig_y, theta, beta, grid=None, isotropic=True): + """Generate a bivariate generalized Gaussian kernel. + + ``Paper: Parameter Estimation For Multivariate Generalized Gaussian Distributions`` + + In the isotropic mode, only `sig_x` is used. `sig_y` and `theta` is ignored. + + Args: + kernel_size (int): + sig_x (float): + sig_y (float): + theta (float): Radian measurement. + beta (float): shape parameter, beta = 1 is the normal distribution. + grid (ndarray, optional): generated by :func:`mesh_grid`, + with the shape (K, K, 2), K is the kernel size. Default: None + + Returns: + kernel (ndarray): normalized kernel. + """ + if grid is None: + grid, _, _ = mesh_grid(kernel_size) + if isotropic: + sigma_matrix = np.array([[sig_x**2, 0], [0, sig_x**2]]) + else: + sigma_matrix = sigma_matrix2(sig_x, sig_y, theta) + inverse_sigma = np.linalg.inv(sigma_matrix) + kernel = np.exp(-0.5 * np.power(np.sum(np.dot(grid, inverse_sigma) * grid, 2), beta)) + kernel = kernel / np.sum(kernel) + return kernel + + +def bivariate_plateau(kernel_size, sig_x, sig_y, theta, beta, grid=None, isotropic=True): + """Generate a plateau-like anisotropic kernel. + + 1 / (1+x^(beta)) + + Reference: https://stats.stackexchange.com/questions/203629/is-there-a-plateau-shaped-distribution + + In the isotropic mode, only `sig_x` is used. `sig_y` and `theta` is ignored. + + Args: + kernel_size (int): + sig_x (float): + sig_y (float): + theta (float): Radian measurement. + beta (float): shape parameter, beta = 1 is the normal distribution. + grid (ndarray, optional): generated by :func:`mesh_grid`, + with the shape (K, K, 2), K is the kernel size. Default: None + + Returns: + kernel (ndarray): normalized kernel. + """ + if grid is None: + grid, _, _ = mesh_grid(kernel_size) + if isotropic: + sigma_matrix = np.array([[sig_x**2, 0], [0, sig_x**2]]) + else: + sigma_matrix = sigma_matrix2(sig_x, sig_y, theta) + inverse_sigma = np.linalg.inv(sigma_matrix) + kernel = np.reciprocal(np.power(np.sum(np.dot(grid, inverse_sigma) * grid, 2), beta) + 1) + kernel = kernel / np.sum(kernel) + return kernel + + +def random_bivariate_Gaussian(kernel_size, + sigma_x_range, + sigma_y_range, + rotation_range, + noise_range=None, + isotropic=True): + """Randomly generate bivariate isotropic or anisotropic Gaussian kernels. + + In the isotropic mode, only `sigma_x_range` is used. `sigma_y_range` and `rotation_range` is ignored. + + Args: + kernel_size (int): + sigma_x_range (tuple): [0.6, 5] + sigma_y_range (tuple): [0.6, 5] + rotation range (tuple): [-math.pi, math.pi] + noise_range(tuple, optional): multiplicative kernel noise, + [0.75, 1.25]. Default: None + + Returns: + kernel (ndarray): + """ + assert kernel_size % 2 == 1, 'Kernel size must be an odd number.' + assert sigma_x_range[0] < sigma_x_range[1], 'Wrong sigma_x_range.' + sigma_x = np.random.uniform(sigma_x_range[0], sigma_x_range[1]) + if isotropic is False: + assert sigma_y_range[0] < sigma_y_range[1], 'Wrong sigma_y_range.' + assert rotation_range[0] < rotation_range[1], 'Wrong rotation_range.' + sigma_y = np.random.uniform(sigma_y_range[0], sigma_y_range[1]) + rotation = np.random.uniform(rotation_range[0], rotation_range[1]) + else: + sigma_y = sigma_x + rotation = 0 + + kernel = bivariate_Gaussian(kernel_size, sigma_x, sigma_y, rotation, isotropic=isotropic) + + # add multiplicative noise + if noise_range is not None: + assert noise_range[0] < noise_range[1], 'Wrong noise range.' + noise = np.random.uniform(noise_range[0], noise_range[1], size=kernel.shape) + kernel = kernel * noise + kernel = kernel / np.sum(kernel) + return kernel + + +def random_bivariate_generalized_Gaussian(kernel_size, + sigma_x_range, + sigma_y_range, + rotation_range, + beta_range, + noise_range=None, + isotropic=True): + """Randomly generate bivariate generalized Gaussian kernels. + + In the isotropic mode, only `sigma_x_range` is used. `sigma_y_range` and `rotation_range` is ignored. + + Args: + kernel_size (int): + sigma_x_range (tuple): [0.6, 5] + sigma_y_range (tuple): [0.6, 5] + rotation range (tuple): [-math.pi, math.pi] + beta_range (tuple): [0.5, 8] + noise_range(tuple, optional): multiplicative kernel noise, + [0.75, 1.25]. Default: None + + Returns: + kernel (ndarray): + """ + assert kernel_size % 2 == 1, 'Kernel size must be an odd number.' + assert sigma_x_range[0] < sigma_x_range[1], 'Wrong sigma_x_range.' + sigma_x = np.random.uniform(sigma_x_range[0], sigma_x_range[1]) + if isotropic is False: + assert sigma_y_range[0] < sigma_y_range[1], 'Wrong sigma_y_range.' + assert rotation_range[0] < rotation_range[1], 'Wrong rotation_range.' + sigma_y = np.random.uniform(sigma_y_range[0], sigma_y_range[1]) + rotation = np.random.uniform(rotation_range[0], rotation_range[1]) + else: + sigma_y = sigma_x + rotation = 0 + + # assume beta_range[0] < 1 < beta_range[1] + if np.random.uniform() < 0.5: + beta = np.random.uniform(beta_range[0], 1) + else: + beta = np.random.uniform(1, beta_range[1]) + + kernel = bivariate_generalized_Gaussian(kernel_size, sigma_x, sigma_y, rotation, beta, isotropic=isotropic) + + # add multiplicative noise + if noise_range is not None: + assert noise_range[0] < noise_range[1], 'Wrong noise range.' + noise = np.random.uniform(noise_range[0], noise_range[1], size=kernel.shape) + kernel = kernel * noise + kernel = kernel / np.sum(kernel) + return kernel + + +def random_bivariate_plateau(kernel_size, + sigma_x_range, + sigma_y_range, + rotation_range, + beta_range, + noise_range=None, + isotropic=True): + """Randomly generate bivariate plateau kernels. + + In the isotropic mode, only `sigma_x_range` is used. `sigma_y_range` and `rotation_range` is ignored. + + Args: + kernel_size (int): + sigma_x_range (tuple): [0.6, 5] + sigma_y_range (tuple): [0.6, 5] + rotation range (tuple): [-math.pi/2, math.pi/2] + beta_range (tuple): [1, 4] + noise_range(tuple, optional): multiplicative kernel noise, + [0.75, 1.25]. Default: None + + Returns: + kernel (ndarray): + """ + assert kernel_size % 2 == 1, 'Kernel size must be an odd number.' + assert sigma_x_range[0] < sigma_x_range[1], 'Wrong sigma_x_range.' + sigma_x = np.random.uniform(sigma_x_range[0], sigma_x_range[1]) + if isotropic is False: + assert sigma_y_range[0] < sigma_y_range[1], 'Wrong sigma_y_range.' + assert rotation_range[0] < rotation_range[1], 'Wrong rotation_range.' + sigma_y = np.random.uniform(sigma_y_range[0], sigma_y_range[1]) + rotation = np.random.uniform(rotation_range[0], rotation_range[1]) + else: + sigma_y = sigma_x + rotation = 0 + + # TODO: this may be not proper + if np.random.uniform() < 0.5: + beta = np.random.uniform(beta_range[0], 1) + else: + beta = np.random.uniform(1, beta_range[1]) + + kernel = bivariate_plateau(kernel_size, sigma_x, sigma_y, rotation, beta, isotropic=isotropic) + # add multiplicative noise + if noise_range is not None: + assert noise_range[0] < noise_range[1], 'Wrong noise range.' + noise = np.random.uniform(noise_range[0], noise_range[1], size=kernel.shape) + kernel = kernel * noise + kernel = kernel / np.sum(kernel) + + return kernel + + +def random_mixed_kernels(kernel_list, + kernel_prob, + kernel_size=21, + sigma_x_range=(0.6, 5), + sigma_y_range=(0.6, 5), + rotation_range=(-math.pi, math.pi), + betag_range=(0.5, 8), + betap_range=(0.5, 8), + noise_range=None): + """Randomly generate mixed kernels. + + Args: + kernel_list (tuple): a list name of kernel types, + support ['iso', 'aniso', 'skew', 'generalized', 'plateau_iso', + 'plateau_aniso'] + kernel_prob (tuple): corresponding kernel probability for each + kernel type + kernel_size (int): + sigma_x_range (tuple): [0.6, 5] + sigma_y_range (tuple): [0.6, 5] + rotation range (tuple): [-math.pi, math.pi] + beta_range (tuple): [0.5, 8] + noise_range(tuple, optional): multiplicative kernel noise, + [0.75, 1.25]. Default: None + + Returns: + kernel (ndarray): + """ + kernel_type = random.choices(kernel_list, kernel_prob)[0] + if kernel_type == 'iso': + kernel = random_bivariate_Gaussian( + kernel_size, sigma_x_range, sigma_y_range, rotation_range, noise_range=noise_range, isotropic=True) + elif kernel_type == 'aniso': + kernel = random_bivariate_Gaussian( + kernel_size, sigma_x_range, sigma_y_range, rotation_range, noise_range=noise_range, isotropic=False) + elif kernel_type == 'generalized_iso': + kernel = random_bivariate_generalized_Gaussian( + kernel_size, + sigma_x_range, + sigma_y_range, + rotation_range, + betag_range, + noise_range=noise_range, + isotropic=True) + elif kernel_type == 'generalized_aniso': + kernel = random_bivariate_generalized_Gaussian( + kernel_size, + sigma_x_range, + sigma_y_range, + rotation_range, + betag_range, + noise_range=noise_range, + isotropic=False) + elif kernel_type == 'plateau_iso': + kernel = random_bivariate_plateau( + kernel_size, sigma_x_range, sigma_y_range, rotation_range, betap_range, noise_range=None, isotropic=True) + elif kernel_type == 'plateau_aniso': + kernel = random_bivariate_plateau( + kernel_size, sigma_x_range, sigma_y_range, rotation_range, betap_range, noise_range=None, isotropic=False) + return kernel + + +np.seterr(divide='ignore', invalid='ignore') + + +def circular_lowpass_kernel(cutoff, kernel_size, pad_to=0): + """2D sinc filter + + Reference: https://dsp.stackexchange.com/questions/58301/2-d-circularly-symmetric-low-pass-filter + + Args: + cutoff (float): cutoff frequency in radians (pi is max) + kernel_size (int): horizontal and vertical size, must be odd. + pad_to (int): pad kernel size to desired size, must be odd or zero. + """ + assert kernel_size % 2 == 1, 'Kernel size must be an odd number.' + kernel = np.fromfunction( + lambda x, y: cutoff * special.j1(cutoff * np.sqrt( + (x - (kernel_size - 1) / 2)**2 + (y - (kernel_size - 1) / 2)**2)) / (2 * np.pi * np.sqrt( + (x - (kernel_size - 1) / 2)**2 + (y - (kernel_size - 1) / 2)**2)), [kernel_size, kernel_size]) + kernel[(kernel_size - 1) // 2, (kernel_size - 1) // 2] = cutoff**2 / (4 * np.pi) + kernel = kernel / np.sum(kernel) + if pad_to > kernel_size: + pad_size = (pad_to - kernel_size) // 2 + kernel = np.pad(kernel, ((pad_size, pad_size), (pad_size, pad_size))) + return kernel + + +# ------------------------------------------------------------- # +# --------------------------- noise --------------------------- # +# ------------------------------------------------------------- # + +# ----------------------- Gaussian Noise ----------------------- # + + +def generate_gaussian_noise(img, sigma=10, gray_noise=False): + """Generate Gaussian noise. + + Args: + img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32. + sigma (float): Noise scale (measured in range 255). Default: 10. + + Returns: + (Numpy array): Returned noisy image, shape (h, w, c), range[0, 1], + float32. + """ + if gray_noise: + noise = np.float32(np.random.randn(*(img.shape[0:2]))) * sigma / 255. + noise = np.expand_dims(noise, axis=2).repeat(3, axis=2) + else: + noise = np.float32(np.random.randn(*(img.shape))) * sigma / 255. + return noise + + +def add_gaussian_noise(img, sigma=10, clip=True, rounds=False, gray_noise=False): + """Add Gaussian noise. + + Args: + img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32. + sigma (float): Noise scale (measured in range 255). Default: 10. + + Returns: + (Numpy array): Returned noisy image, shape (h, w, c), range[0, 1], + float32. + """ + noise = generate_gaussian_noise(img, sigma, gray_noise) + out = img + noise + if clip and rounds: + out = np.clip((out * 255.0).round(), 0, 255) / 255. + elif clip: + out = np.clip(out, 0, 1) + elif rounds: + out = (out * 255.0).round() / 255. + return out + + +def generate_gaussian_noise_pt(img, sigma=10, gray_noise=0): + """Add Gaussian noise (PyTorch version). + + Args: + img (Tensor): Shape (b, c, h, w), range[0, 1], float32. + scale (float | Tensor): Noise scale. Default: 1.0. + + Returns: + (Tensor): Returned noisy image, shape (b, c, h, w), range[0, 1], + float32. + """ + b, _, h, w = img.size() + if not isinstance(sigma, (float, int)): + sigma = sigma.view(img.size(0), 1, 1, 1) + if isinstance(gray_noise, (float, int)): + cal_gray_noise = gray_noise > 0 + else: + gray_noise = gray_noise.view(b, 1, 1, 1) + cal_gray_noise = torch.sum(gray_noise) > 0 + + if cal_gray_noise: + noise_gray = torch.randn(*img.size()[2:4], dtype=img.dtype, device=img.device) * sigma / 255. + noise_gray = noise_gray.view(b, 1, h, w) + + # always calculate color noise + noise = torch.randn(*img.size(), dtype=img.dtype, device=img.device) * sigma / 255. + + if cal_gray_noise: + noise = noise * (1 - gray_noise) + noise_gray * gray_noise + return noise + + +def add_gaussian_noise_pt(img, sigma=10, gray_noise=0, clip=True, rounds=False): + """Add Gaussian noise (PyTorch version). + + Args: + img (Tensor): Shape (b, c, h, w), range[0, 1], float32. + scale (float | Tensor): Noise scale. Default: 1.0. + + Returns: + (Tensor): Returned noisy image, shape (b, c, h, w), range[0, 1], + float32. + """ + noise = generate_gaussian_noise_pt(img, sigma, gray_noise) + out = img + noise + if clip and rounds: + out = torch.clamp((out * 255.0).round(), 0, 255) / 255. + elif clip: + out = torch.clamp(out, 0, 1) + elif rounds: + out = (out * 255.0).round() / 255. + return out + + +# ----------------------- Random Gaussian Noise ----------------------- # +def random_generate_gaussian_noise(img, sigma_range=(0, 10), gray_prob=0): + sigma = np.random.uniform(sigma_range[0], sigma_range[1]) + if np.random.uniform() < gray_prob: + gray_noise = True + else: + gray_noise = False + return generate_gaussian_noise(img, sigma, gray_noise) + + +def random_add_gaussian_noise(img, sigma_range=(0, 1.0), gray_prob=0, clip=True, rounds=False): + noise = random_generate_gaussian_noise(img, sigma_range, gray_prob) + out = img + noise + if clip and rounds: + out = np.clip((out * 255.0).round(), 0, 255) / 255. + elif clip: + out = np.clip(out, 0, 1) + elif rounds: + out = (out * 255.0).round() / 255. + return out + + +def random_generate_gaussian_noise_pt(img, sigma_range=(0, 10), gray_prob=0): + sigma = torch.rand( + img.size(0), dtype=img.dtype, device=img.device) * (sigma_range[1] - sigma_range[0]) + sigma_range[0] + gray_noise = torch.rand(img.size(0), dtype=img.dtype, device=img.device) + gray_noise = (gray_noise < gray_prob).float() + return generate_gaussian_noise_pt(img, sigma, gray_noise) + + +def random_add_gaussian_noise_pt(img, sigma_range=(0, 1.0), gray_prob=0, clip=True, rounds=False): + noise = random_generate_gaussian_noise_pt(img, sigma_range, gray_prob) + out = img + noise + if clip and rounds: + out = torch.clamp((out * 255.0).round(), 0, 255) / 255. + elif clip: + out = torch.clamp(out, 0, 1) + elif rounds: + out = (out * 255.0).round() / 255. + return out + + +# ----------------------- Poisson (Shot) Noise ----------------------- # + + +def generate_poisson_noise(img, scale=1.0, gray_noise=False): + """Generate poisson noise. + + Reference: https://github.com/scikit-image/scikit-image/blob/main/skimage/util/noise.py#L37-L219 + + Args: + img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32. + scale (float): Noise scale. Default: 1.0. + gray_noise (bool): Whether generate gray noise. Default: False. + + Returns: + (Numpy array): Returned noisy image, shape (h, w, c), range[0, 1], + float32. + """ + if gray_noise: + img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY) + # round and clip image for counting vals correctly + img = np.clip((img * 255.0).round(), 0, 255) / 255. + vals = len(np.unique(img)) + vals = 2**np.ceil(np.log2(vals)) + out = np.float32(np.random.poisson(img * vals) / float(vals)) + noise = out - img + if gray_noise: + noise = np.repeat(noise[:, :, np.newaxis], 3, axis=2) + return noise * scale + + +def add_poisson_noise(img, scale=1.0, clip=True, rounds=False, gray_noise=False): + """Add poisson noise. + + Args: + img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32. + scale (float): Noise scale. Default: 1.0. + gray_noise (bool): Whether generate gray noise. Default: False. + + Returns: + (Numpy array): Returned noisy image, shape (h, w, c), range[0, 1], + float32. + """ + noise = generate_poisson_noise(img, scale, gray_noise) + out = img + noise + if clip and rounds: + out = np.clip((out * 255.0).round(), 0, 255) / 255. + elif clip: + out = np.clip(out, 0, 1) + elif rounds: + out = (out * 255.0).round() / 255. + return out + + +def generate_poisson_noise_pt(img, scale=1.0, gray_noise=0): + """Generate a batch of poisson noise (PyTorch version) + + Args: + img (Tensor): Input image, shape (b, c, h, w), range [0, 1], float32. + scale (float | Tensor): Noise scale. Number or Tensor with shape (b). + Default: 1.0. + gray_noise (float | Tensor): 0-1 number or Tensor with shape (b). + 0 for False, 1 for True. Default: 0. + + Returns: + (Tensor): Returned noisy image, shape (b, c, h, w), range[0, 1], + float32. + """ + b, _, h, w = img.size() + if isinstance(gray_noise, (float, int)): + cal_gray_noise = gray_noise > 0 + else: + gray_noise = gray_noise.view(b, 1, 1, 1) + cal_gray_noise = torch.sum(gray_noise) > 0 + if cal_gray_noise: + img_gray = rgb_to_grayscale(img, num_output_channels=1) + # round and clip image for counting vals correctly + img_gray = torch.clamp((img_gray * 255.0).round(), 0, 255) / 255. + # use for-loop to get the unique values for each sample + vals_list = [len(torch.unique(img_gray[i, :, :, :])) for i in range(b)] + vals_list = [2**np.ceil(np.log2(vals)) for vals in vals_list] + vals = img_gray.new_tensor(vals_list).view(b, 1, 1, 1) + out = torch.poisson(img_gray * vals) / vals + noise_gray = out - img_gray + noise_gray = noise_gray.expand(b, 3, h, w) + + # always calculate color noise + # round and clip image for counting vals correctly + img = torch.clamp((img * 255.0).round(), 0, 255) / 255. + # use for-loop to get the unique values for each sample + vals_list = [len(torch.unique(img[i, :, :, :])) for i in range(b)] + vals_list = [2**np.ceil(np.log2(vals)) for vals in vals_list] + vals = img.new_tensor(vals_list).view(b, 1, 1, 1) + out = torch.poisson(img * vals) / vals + noise = out - img + if cal_gray_noise: + noise = noise * (1 - gray_noise) + noise_gray * gray_noise + if not isinstance(scale, (float, int)): + scale = scale.view(b, 1, 1, 1) + return noise * scale + + +def add_poisson_noise_pt(img, scale=1.0, clip=True, rounds=False, gray_noise=0): + """Add poisson noise to a batch of images (PyTorch version). + + Args: + img (Tensor): Input image, shape (b, c, h, w), range [0, 1], float32. + scale (float | Tensor): Noise scale. Number or Tensor with shape (b). + Default: 1.0. + gray_noise (float | Tensor): 0-1 number or Tensor with shape (b). + 0 for False, 1 for True. Default: 0. + + Returns: + (Tensor): Returned noisy image, shape (b, c, h, w), range[0, 1], + float32. + """ + noise = generate_poisson_noise_pt(img, scale, gray_noise) + out = img + noise + if clip and rounds: + out = torch.clamp((out * 255.0).round(), 0, 255) / 255. + elif clip: + out = torch.clamp(out, 0, 1) + elif rounds: + out = (out * 255.0).round() / 255. + return out + + +# ----------------------- Random Poisson (Shot) Noise ----------------------- # + + +def random_generate_poisson_noise(img, scale_range=(0, 1.0), gray_prob=0): + scale = np.random.uniform(scale_range[0], scale_range[1]) + if np.random.uniform() < gray_prob: + gray_noise = True + else: + gray_noise = False + return generate_poisson_noise(img, scale, gray_noise) + + +def random_add_poisson_noise(img, scale_range=(0, 1.0), gray_prob=0, clip=True, rounds=False): + noise = random_generate_poisson_noise(img, scale_range, gray_prob) + out = img + noise + if clip and rounds: + out = np.clip((out * 255.0).round(), 0, 255) / 255. + elif clip: + out = np.clip(out, 0, 1) + elif rounds: + out = (out * 255.0).round() / 255. + return out + + +def random_generate_poisson_noise_pt(img, scale_range=(0, 1.0), gray_prob=0): + scale = torch.rand( + img.size(0), dtype=img.dtype, device=img.device) * (scale_range[1] - scale_range[0]) + scale_range[0] + gray_noise = torch.rand(img.size(0), dtype=img.dtype, device=img.device) + gray_noise = (gray_noise < gray_prob).float() + return generate_poisson_noise_pt(img, scale, gray_noise) + + +def random_add_poisson_noise_pt(img, scale_range=(0, 1.0), gray_prob=0, clip=True, rounds=False): + noise = random_generate_poisson_noise_pt(img, scale_range, gray_prob) + out = img + noise + if clip and rounds: + out = torch.clamp((out * 255.0).round(), 0, 255) / 255. + elif clip: + out = torch.clamp(out, 0, 1) + elif rounds: + out = (out * 255.0).round() / 255. + return out + + +# ------------------------------------------------------------------------ # +# --------------------------- JPEG compression --------------------------- # +# ------------------------------------------------------------------------ # + + +def add_jpg_compression(img, quality=90): + """Add JPG compression artifacts. + + Args: + img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32. + quality (float): JPG compression quality. 0 for lowest quality, 100 for + best quality. Default: 90. + + Returns: + (Numpy array): Returned image after JPG, shape (h, w, c), range[0, 1], + float32. + """ + img = np.clip(img, 0, 1) + encode_param = [int(cv2.IMWRITE_JPEG_QUALITY), quality] + _, encimg = cv2.imencode('.jpg', img * 255., encode_param) + img = np.float32(cv2.imdecode(encimg, 1)) / 255. + return img + + +def random_add_jpg_compression(img, quality_range=(90, 100)): + """Randomly add JPG compression artifacts. + + Args: + img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32. + quality_range (tuple[float] | list[float]): JPG compression quality + range. 0 for lowest quality, 100 for best quality. + Default: (90, 100). + + Returns: + (Numpy array): Returned image after JPG, shape (h, w, c), range[0, 1], + float32. + """ + quality = np.random.uniform(quality_range[0], quality_range[1]) + return add_jpg_compression(img, quality) diff --git a/official/bsr/transforms.py b/official/bsr/transforms.py new file mode 100644 index 0000000000000000000000000000000000000000..85d1bc2b3587995f9d87d242bd266c50846f95fd --- /dev/null +++ b/official/bsr/transforms.py @@ -0,0 +1,179 @@ +import cv2 +import random +import torch + + +def mod_crop(img, scale): + """Mod crop images, used during testing. + + Args: + img (ndarray): Input image. + scale (int): Scale factor. + + Returns: + ndarray: Result image. + """ + img = img.copy() + if img.ndim in (2, 3): + h, w = img.shape[0], img.shape[1] + h_remainder, w_remainder = h % scale, w % scale + img = img[:h - h_remainder, :w - w_remainder, ...] + else: + raise ValueError(f'Wrong img ndim: {img.ndim}.') + return img + + +def paired_random_crop(img_gts, img_lqs, gt_patch_size, scale, gt_path=None): + """Paired random crop. Support Numpy array and Tensor inputs. + + It crops lists of lq and gt images with corresponding locations. + + Args: + img_gts (list[ndarray] | ndarray | list[Tensor] | Tensor): GT images. Note that all images + should have the same shape. If the input is an ndarray, it will + be transformed to a list containing itself. + img_lqs (list[ndarray] | ndarray): LQ images. Note that all images + should have the same shape. If the input is an ndarray, it will + be transformed to a list containing itself. + gt_patch_size (int): GT patch size. + scale (int): Scale factor. + gt_path (str): Path to ground-truth. Default: None. + + Returns: + list[ndarray] | ndarray: GT images and LQ images. If returned results + only have one element, just return ndarray. + """ + + if not isinstance(img_gts, list): + img_gts = [img_gts] + if not isinstance(img_lqs, list): + img_lqs = [img_lqs] + + # determine input type: Numpy array or Tensor + input_type = 'Tensor' if torch.is_tensor(img_gts[0]) else 'Numpy' + + if input_type == 'Tensor': + h_lq, w_lq = img_lqs[0].size()[-2:] + h_gt, w_gt = img_gts[0].size()[-2:] + else: + h_lq, w_lq = img_lqs[0].shape[0:2] + h_gt, w_gt = img_gts[0].shape[0:2] + lq_patch_size = gt_patch_size // scale + + if h_gt != h_lq * scale or w_gt != w_lq * scale: + raise ValueError(f'Scale mismatches. GT ({h_gt}, {w_gt}) is not {scale}x ', + f'multiplication of LQ ({h_lq}, {w_lq}).') + if h_lq < lq_patch_size or w_lq < lq_patch_size: + raise ValueError(f'LQ ({h_lq}, {w_lq}) is smaller than patch size ' + f'({lq_patch_size}, {lq_patch_size}). ' + f'Please remove {gt_path}.') + + # randomly choose top and left coordinates for lq patch + top = random.randint(0, h_lq - lq_patch_size) + left = random.randint(0, w_lq - lq_patch_size) + + # crop lq patch + if input_type == 'Tensor': + img_lqs = [v[:, :, top:top + lq_patch_size, left:left + lq_patch_size] for v in img_lqs] + else: + img_lqs = [v[top:top + lq_patch_size, left:left + lq_patch_size, ...] for v in img_lqs] + + # crop corresponding gt patch + top_gt, left_gt = int(top * scale), int(left * scale) + if input_type == 'Tensor': + img_gts = [v[:, :, top_gt:top_gt + gt_patch_size, left_gt:left_gt + gt_patch_size] for v in img_gts] + else: + img_gts = [v[top_gt:top_gt + gt_patch_size, left_gt:left_gt + gt_patch_size, ...] for v in img_gts] + if len(img_gts) == 1: + img_gts = img_gts[0] + if len(img_lqs) == 1: + img_lqs = img_lqs[0] + return img_gts, img_lqs + + +def augment(imgs, hflip=True, rotation=True, flows=None, return_status=False): + """Augment: horizontal flips OR rotate (0, 90, 180, 270 degrees). + + We use vertical flip and transpose for rotation implementation. + All the images in the list use the same augmentation. + + Args: + imgs (list[ndarray] | ndarray): Images to be augmented. If the input + is an ndarray, it will be transformed to a list. + hflip (bool): Horizontal flip. Default: True. + rotation (bool): Ratotation. Default: True. + flows (list[ndarray]: Flows to be augmented. If the input is an + ndarray, it will be transformed to a list. + Dimension is (h, w, 2). Default: None. + return_status (bool): Return the status of flip and rotation. + Default: False. + + Returns: + list[ndarray] | ndarray: Augmented images and flows. If returned + results only have one element, just return ndarray. + + """ + hflip = hflip and random.random() < 0.5 + vflip = rotation and random.random() < 0.5 + rot90 = rotation and random.random() < 0.5 + + def _augment(img): + if hflip: # horizontal + cv2.flip(img, 1, img) + if vflip: # vertical + cv2.flip(img, 0, img) + if rot90: + img = img.transpose(1, 0, 2) + return img + + def _augment_flow(flow): + if hflip: # horizontal + cv2.flip(flow, 1, flow) + flow[:, :, 0] *= -1 + if vflip: # vertical + cv2.flip(flow, 0, flow) + flow[:, :, 1] *= -1 + if rot90: + flow = flow.transpose(1, 0, 2) + flow = flow[:, :, [1, 0]] + return flow + + if not isinstance(imgs, list): + imgs = [imgs] + imgs = [_augment(img) for img in imgs] + if len(imgs) == 1: + imgs = imgs[0] + + if flows is not None: + if not isinstance(flows, list): + flows = [flows] + flows = [_augment_flow(flow) for flow in flows] + if len(flows) == 1: + flows = flows[0] + return imgs, flows + else: + if return_status: + return imgs, (hflip, vflip, rot90) + else: + return imgs + + +def img_rotate(img, angle, center=None, scale=1.0): + """Rotate image. + + Args: + img (ndarray): Image to be rotated. + angle (float): Rotation angle in degrees. Positive values mean + counter-clockwise rotation. + center (tuple[int]): Rotation center. If the center is None, + initialize it as the center of the image. Default: None. + scale (float): Isotropic scale factor. Default: 1.0. + """ + (h, w) = img.shape[:2] + + if center is None: + center = (w // 2, h // 2) + + matrix = cv2.getRotationMatrix2D(center, angle, scale) + rotated_img = cv2.warpAffine(img, matrix, (w, h)) + return rotated_img diff --git a/official/bsr/utils/__init__.py b/official/bsr/utils/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..85670f1660b4a4f50c3ecd29d933cf0afcf17357 --- /dev/null +++ b/official/bsr/utils/__init__.py @@ -0,0 +1,47 @@ +from .color_util import bgr2ycbcr, rgb2ycbcr, rgb2ycbcr_pt, ycbcr2bgr, ycbcr2rgb +from .diffjpeg import DiffJPEG +from .file_client import FileClient +from .img_process_util import USMSharp, usm_sharp +from .img_util import crop_border, imfrombytes, img2tensor, imwrite, tensor2img +from .logger import AvgTimer, MessageLogger, get_env_info, get_root_logger, init_tb_logger, init_wandb_logger +from .misc import check_resume, get_time_str, make_exp_dirs, mkdir_and_rename, scandir, set_random_seed, sizeof_fmt +from .options import yaml_load + +__all__ = [ + # color_util.py + 'bgr2ycbcr', + 'rgb2ycbcr', + 'rgb2ycbcr_pt', + 'ycbcr2bgr', + 'ycbcr2rgb', + # file_client.py + 'FileClient', + # img_util.py + 'img2tensor', + 'tensor2img', + 'imfrombytes', + 'imwrite', + 'crop_border', + # logger.py + 'MessageLogger', + 'AvgTimer', + 'init_tb_logger', + 'init_wandb_logger', + 'get_root_logger', + 'get_env_info', + # misc.py + 'set_random_seed', + 'get_time_str', + 'mkdir_and_rename', + 'make_exp_dirs', + 'scandir', + 'check_resume', + 'sizeof_fmt', + # diffjpeg + 'DiffJPEG', + # img_process_util + 'USMSharp', + 'usm_sharp', + # options + 'yaml_load' +] diff --git a/official/bsr/utils/color_util.py b/official/bsr/utils/color_util.py new file mode 100644 index 0000000000000000000000000000000000000000..8b7676fd78e007300d54950e553f8255b7a86a82 --- /dev/null +++ b/official/bsr/utils/color_util.py @@ -0,0 +1,208 @@ +import numpy as np +import torch + + +def rgb2ycbcr(img, y_only=False): + """Convert a RGB image to YCbCr image. + + This function produces the same results as Matlab's `rgb2ycbcr` function. + It implements the ITU-R BT.601 conversion for standard-definition + television. See more details in + https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion. + + It differs from a similar function in cv2.cvtColor: `RGB <-> YCrCb`. + In OpenCV, it implements a JPEG conversion. See more details in + https://en.wikipedia.org/wiki/YCbCr#JPEG_conversion. + + Args: + img (ndarray): The input image. It accepts: + 1. np.uint8 type with range [0, 255]; + 2. np.float32 type with range [0, 1]. + y_only (bool): Whether to only return Y channel. Default: False. + + Returns: + ndarray: The converted YCbCr image. The output image has the same type + and range as input image. + """ + img_type = img.dtype + img = _convert_input_type_range(img) + if y_only: + out_img = np.dot(img, [65.481, 128.553, 24.966]) + 16.0 + else: + out_img = np.matmul( + img, [[65.481, -37.797, 112.0], [128.553, -74.203, -93.786], [24.966, 112.0, -18.214]]) + [16, 128, 128] + out_img = _convert_output_type_range(out_img, img_type) + return out_img + + +def bgr2ycbcr(img, y_only=False): + """Convert a BGR image to YCbCr image. + + The bgr version of rgb2ycbcr. + It implements the ITU-R BT.601 conversion for standard-definition + television. See more details in + https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion. + + It differs from a similar function in cv2.cvtColor: `BGR <-> YCrCb`. + In OpenCV, it implements a JPEG conversion. See more details in + https://en.wikipedia.org/wiki/YCbCr#JPEG_conversion. + + Args: + img (ndarray): The input image. It accepts: + 1. np.uint8 type with range [0, 255]; + 2. np.float32 type with range [0, 1]. + y_only (bool): Whether to only return Y channel. Default: False. + + Returns: + ndarray: The converted YCbCr image. The output image has the same type + and range as input image. + """ + img_type = img.dtype + img = _convert_input_type_range(img) + if y_only: + out_img = np.dot(img, [24.966, 128.553, 65.481]) + 16.0 + else: + out_img = np.matmul( + img, [[24.966, 112.0, -18.214], [128.553, -74.203, -93.786], [65.481, -37.797, 112.0]]) + [16, 128, 128] + out_img = _convert_output_type_range(out_img, img_type) + return out_img + + +def ycbcr2rgb(img): + """Convert a YCbCr image to RGB image. + + This function produces the same results as Matlab's ycbcr2rgb function. + It implements the ITU-R BT.601 conversion for standard-definition + television. See more details in + https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion. + + It differs from a similar function in cv2.cvtColor: `YCrCb <-> RGB`. + In OpenCV, it implements a JPEG conversion. See more details in + https://en.wikipedia.org/wiki/YCbCr#JPEG_conversion. + + Args: + img (ndarray): The input image. It accepts: + 1. np.uint8 type with range [0, 255]; + 2. np.float32 type with range [0, 1]. + + Returns: + ndarray: The converted RGB image. The output image has the same type + and range as input image. + """ + img_type = img.dtype + img = _convert_input_type_range(img) * 255 + out_img = np.matmul(img, [[0.00456621, 0.00456621, 0.00456621], [0, -0.00153632, 0.00791071], + [0.00625893, -0.00318811, 0]]) * 255.0 + [-222.921, 135.576, -276.836] # noqa: E126 + out_img = _convert_output_type_range(out_img, img_type) + return out_img + + +def ycbcr2bgr(img): + """Convert a YCbCr image to BGR image. + + The bgr version of ycbcr2rgb. + It implements the ITU-R BT.601 conversion for standard-definition + television. See more details in + https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion. + + It differs from a similar function in cv2.cvtColor: `YCrCb <-> BGR`. + In OpenCV, it implements a JPEG conversion. See more details in + https://en.wikipedia.org/wiki/YCbCr#JPEG_conversion. + + Args: + img (ndarray): The input image. It accepts: + 1. np.uint8 type with range [0, 255]; + 2. np.float32 type with range [0, 1]. + + Returns: + ndarray: The converted BGR image. The output image has the same type + and range as input image. + """ + img_type = img.dtype + img = _convert_input_type_range(img) * 255 + out_img = np.matmul(img, [[0.00456621, 0.00456621, 0.00456621], [0.00791071, -0.00153632, 0], + [0, -0.00318811, 0.00625893]]) * 255.0 + [-276.836, 135.576, -222.921] # noqa: E126 + out_img = _convert_output_type_range(out_img, img_type) + return out_img + + +def _convert_input_type_range(img): + """Convert the type and range of the input image. + + It converts the input image to np.float32 type and range of [0, 1]. + It is mainly used for pre-processing the input image in colorspace + conversion functions such as rgb2ycbcr and ycbcr2rgb. + + Args: + img (ndarray): The input image. It accepts: + 1. np.uint8 type with range [0, 255]; + 2. np.float32 type with range [0, 1]. + + Returns: + (ndarray): The converted image with type of np.float32 and range of + [0, 1]. + """ + img_type = img.dtype + img = img.astype(np.float32) + if img_type == np.float32: + pass + elif img_type == np.uint8: + img /= 255. + else: + raise TypeError(f'The img type should be np.float32 or np.uint8, but got {img_type}') + return img + + +def _convert_output_type_range(img, dst_type): + """Convert the type and range of the image according to dst_type. + + It converts the image to desired type and range. If `dst_type` is np.uint8, + images will be converted to np.uint8 type with range [0, 255]. If + `dst_type` is np.float32, it converts the image to np.float32 type with + range [0, 1]. + It is mainly used for post-processing images in colorspace conversion + functions such as rgb2ycbcr and ycbcr2rgb. + + Args: + img (ndarray): The image to be converted with np.float32 type and + range [0, 255]. + dst_type (np.uint8 | np.float32): If dst_type is np.uint8, it + converts the image to np.uint8 type with range [0, 255]. If + dst_type is np.float32, it converts the image to np.float32 type + with range [0, 1]. + + Returns: + (ndarray): The converted image with desired type and range. + """ + if dst_type not in (np.uint8, np.float32): + raise TypeError(f'The dst_type should be np.float32 or np.uint8, but got {dst_type}') + if dst_type == np.uint8: + img = img.round() + else: + img /= 255. + return img.astype(dst_type) + + +def rgb2ycbcr_pt(img, y_only=False): + """Convert RGB images to YCbCr images (PyTorch version). + + It implements the ITU-R BT.601 conversion for standard-definition television. See more details in + https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion. + + Args: + img (Tensor): Images with shape (n, 3, h, w), the range [0, 1], float, RGB format. + y_only (bool): Whether to only return Y channel. Default: False. + + Returns: + (Tensor): converted images with the shape (n, 3/1, h, w), the range [0, 1], float. + """ + if y_only: + weight = torch.tensor([[65.481], [128.553], [24.966]]).to(img) + out_img = torch.matmul(img.permute(0, 2, 3, 1), weight).permute(0, 3, 1, 2) + 16.0 + else: + weight = torch.tensor([[65.481, -37.797, 112.0], [128.553, -74.203, -93.786], [24.966, 112.0, -18.214]]).to(img) + bias = torch.tensor([16, 128, 128]).view(1, 3, 1, 1).to(img) + out_img = torch.matmul(img.permute(0, 2, 3, 1), weight).permute(0, 3, 1, 2) + bias + + out_img = out_img / 255. + return out_img diff --git a/official/bsr/utils/diffjpeg.py b/official/bsr/utils/diffjpeg.py new file mode 100644 index 0000000000000000000000000000000000000000..4273c9663f8b9a0a618676751d89901e636ea241 --- /dev/null +++ b/official/bsr/utils/diffjpeg.py @@ -0,0 +1,515 @@ +""" +Modified from https://github.com/mlomnitz/DiffJPEG + +For images not divisible by 8 +https://dsp.stackexchange.com/questions/35339/jpeg-dct-padding/35343#35343 +""" +import itertools +import numpy as np +import torch +import torch.nn as nn +from torch.nn import functional as F + +# ------------------------ utils ------------------------# +y_table = np.array( + [[16, 11, 10, 16, 24, 40, 51, 61], [12, 12, 14, 19, 26, 58, 60, 55], [14, 13, 16, 24, 40, 57, 69, 56], + [14, 17, 22, 29, 51, 87, 80, 62], [18, 22, 37, 56, 68, 109, 103, 77], [24, 35, 55, 64, 81, 104, 113, 92], + [49, 64, 78, 87, 103, 121, 120, 101], [72, 92, 95, 98, 112, 100, 103, 99]], + dtype=np.float32).T +y_table = nn.Parameter(torch.from_numpy(y_table)) +c_table = np.empty((8, 8), dtype=np.float32) +c_table.fill(99) +c_table[:4, :4] = np.array([[17, 18, 24, 47], [18, 21, 26, 66], [24, 26, 56, 99], [47, 66, 99, 99]]).T +c_table = nn.Parameter(torch.from_numpy(c_table)) + + +def diff_round(x): + """ Differentiable rounding function + """ + return torch.round(x) + (x - torch.round(x))**3 + + +def quality_to_factor(quality): + """ Calculate factor corresponding to quality + + Args: + quality(float): Quality for jpeg compression. + + Returns: + float: Compression factor. + """ + if quality < 50: + quality = 5000. / quality + else: + quality = 200. - quality * 2 + return quality / 100. + + +# ------------------------ compression ------------------------# +class RGB2YCbCrJpeg(nn.Module): + """ Converts RGB image to YCbCr + """ + + def __init__(self): + super(RGB2YCbCrJpeg, self).__init__() + matrix = np.array([[0.299, 0.587, 0.114], [-0.168736, -0.331264, 0.5], [0.5, -0.418688, -0.081312]], + dtype=np.float32).T + self.shift = nn.Parameter(torch.tensor([0., 128., 128.])) + self.matrix = nn.Parameter(torch.from_numpy(matrix)) + + def forward(self, image): + """ + Args: + image(Tensor): batch x 3 x height x width + + Returns: + Tensor: batch x height x width x 3 + """ + image = image.permute(0, 2, 3, 1) + result = torch.tensordot(image, self.matrix, dims=1) + self.shift + return result.view(image.shape) + + +class ChromaSubsampling(nn.Module): + """ Chroma subsampling on CbCr channels + """ + + def __init__(self): + super(ChromaSubsampling, self).__init__() + + def forward(self, image): + """ + Args: + image(tensor): batch x height x width x 3 + + Returns: + y(tensor): batch x height x width + cb(tensor): batch x height/2 x width/2 + cr(tensor): batch x height/2 x width/2 + """ + image_2 = image.permute(0, 3, 1, 2).clone() + cb = F.avg_pool2d(image_2[:, 1, :, :].unsqueeze(1), kernel_size=2, stride=(2, 2), count_include_pad=False) + cr = F.avg_pool2d(image_2[:, 2, :, :].unsqueeze(1), kernel_size=2, stride=(2, 2), count_include_pad=False) + cb = cb.permute(0, 2, 3, 1) + cr = cr.permute(0, 2, 3, 1) + return image[:, :, :, 0], cb.squeeze(3), cr.squeeze(3) + + +class BlockSplitting(nn.Module): + """ Splitting image into patches + """ + + def __init__(self): + super(BlockSplitting, self).__init__() + self.k = 8 + + def forward(self, image): + """ + Args: + image(tensor): batch x height x width + + Returns: + Tensor: batch x h*w/64 x h x w + """ + height, _ = image.shape[1:3] + batch_size = image.shape[0] + image_reshaped = image.view(batch_size, height // self.k, self.k, -1, self.k) + image_transposed = image_reshaped.permute(0, 1, 3, 2, 4) + return image_transposed.contiguous().view(batch_size, -1, self.k, self.k) + + +class DCT8x8(nn.Module): + """ Discrete Cosine Transformation + """ + + def __init__(self): + super(DCT8x8, self).__init__() + tensor = np.zeros((8, 8, 8, 8), dtype=np.float32) + for x, y, u, v in itertools.product(range(8), repeat=4): + tensor[x, y, u, v] = np.cos((2 * x + 1) * u * np.pi / 16) * np.cos((2 * y + 1) * v * np.pi / 16) + alpha = np.array([1. / np.sqrt(2)] + [1] * 7) + self.tensor = nn.Parameter(torch.from_numpy(tensor).float()) + self.scale = nn.Parameter(torch.from_numpy(np.outer(alpha, alpha) * 0.25).float()) + + def forward(self, image): + """ + Args: + image(tensor): batch x height x width + + Returns: + Tensor: batch x height x width + """ + image = image - 128 + result = self.scale * torch.tensordot(image, self.tensor, dims=2) + result.view(image.shape) + return result + + +class YQuantize(nn.Module): + """ JPEG Quantization for Y channel + + Args: + rounding(function): rounding function to use + """ + + def __init__(self, rounding): + super(YQuantize, self).__init__() + self.rounding = rounding + self.y_table = y_table + + def forward(self, image, factor=1): + """ + Args: + image(tensor): batch x height x width + + Returns: + Tensor: batch x height x width + """ + if isinstance(factor, (int, float)): + image = image.float() / (self.y_table * factor) + else: + b = factor.size(0) + table = self.y_table.expand(b, 1, 8, 8) * factor.view(b, 1, 1, 1) + image = image.float() / table + image = self.rounding(image) + return image + + +class CQuantize(nn.Module): + """ JPEG Quantization for CbCr channels + + Args: + rounding(function): rounding function to use + """ + + def __init__(self, rounding): + super(CQuantize, self).__init__() + self.rounding = rounding + self.c_table = c_table + + def forward(self, image, factor=1): + """ + Args: + image(tensor): batch x height x width + + Returns: + Tensor: batch x height x width + """ + if isinstance(factor, (int, float)): + image = image.float() / (self.c_table * factor) + else: + b = factor.size(0) + table = self.c_table.expand(b, 1, 8, 8) * factor.view(b, 1, 1, 1) + image = image.float() / table + image = self.rounding(image) + return image + + +class CompressJpeg(nn.Module): + """Full JPEG compression algorithm + + Args: + rounding(function): rounding function to use + """ + + def __init__(self, rounding=torch.round): + super(CompressJpeg, self).__init__() + self.l1 = nn.Sequential(RGB2YCbCrJpeg(), ChromaSubsampling()) + self.l2 = nn.Sequential(BlockSplitting(), DCT8x8()) + self.c_quantize = CQuantize(rounding=rounding) + self.y_quantize = YQuantize(rounding=rounding) + + def forward(self, image, factor=1): + """ + Args: + image(tensor): batch x 3 x height x width + + Returns: + dict(tensor): Compressed tensor with batch x h*w/64 x 8 x 8. + """ + y, cb, cr = self.l1(image * 255) + components = {'y': y, 'cb': cb, 'cr': cr} + for k in components.keys(): + comp = self.l2(components[k]) + if k in ('cb', 'cr'): + comp = self.c_quantize(comp, factor=factor) + else: + comp = self.y_quantize(comp, factor=factor) + + components[k] = comp + + return components['y'], components['cb'], components['cr'] + + +# ------------------------ decompression ------------------------# + + +class YDequantize(nn.Module): + """Dequantize Y channel + """ + + def __init__(self): + super(YDequantize, self).__init__() + self.y_table = y_table + + def forward(self, image, factor=1): + """ + Args: + image(tensor): batch x height x width + + Returns: + Tensor: batch x height x width + """ + if isinstance(factor, (int, float)): + out = image * (self.y_table * factor) + else: + b = factor.size(0) + table = self.y_table.expand(b, 1, 8, 8) * factor.view(b, 1, 1, 1) + out = image * table + return out + + +class CDequantize(nn.Module): + """Dequantize CbCr channel + """ + + def __init__(self): + super(CDequantize, self).__init__() + self.c_table = c_table + + def forward(self, image, factor=1): + """ + Args: + image(tensor): batch x height x width + + Returns: + Tensor: batch x height x width + """ + if isinstance(factor, (int, float)): + out = image * (self.c_table * factor) + else: + b = factor.size(0) + table = self.c_table.expand(b, 1, 8, 8) * factor.view(b, 1, 1, 1) + out = image * table + return out + + +class iDCT8x8(nn.Module): + """Inverse discrete Cosine Transformation + """ + + def __init__(self): + super(iDCT8x8, self).__init__() + alpha = np.array([1. / np.sqrt(2)] + [1] * 7) + self.alpha = nn.Parameter(torch.from_numpy(np.outer(alpha, alpha)).float()) + tensor = np.zeros((8, 8, 8, 8), dtype=np.float32) + for x, y, u, v in itertools.product(range(8), repeat=4): + tensor[x, y, u, v] = np.cos((2 * u + 1) * x * np.pi / 16) * np.cos((2 * v + 1) * y * np.pi / 16) + self.tensor = nn.Parameter(torch.from_numpy(tensor).float()) + + def forward(self, image): + """ + Args: + image(tensor): batch x height x width + + Returns: + Tensor: batch x height x width + """ + image = image * self.alpha + result = 0.25 * torch.tensordot(image, self.tensor, dims=2) + 128 + result.view(image.shape) + return result + + +class BlockMerging(nn.Module): + """Merge patches into image + """ + + def __init__(self): + super(BlockMerging, self).__init__() + + def forward(self, patches, height, width): + """ + Args: + patches(tensor) batch x height*width/64, height x width + height(int) + width(int) + + Returns: + Tensor: batch x height x width + """ + k = 8 + batch_size = patches.shape[0] + image_reshaped = patches.view(batch_size, height // k, width // k, k, k) + image_transposed = image_reshaped.permute(0, 1, 3, 2, 4) + return image_transposed.contiguous().view(batch_size, height, width) + + +class ChromaUpsampling(nn.Module): + """Upsample chroma layers + """ + + def __init__(self): + super(ChromaUpsampling, self).__init__() + + def forward(self, y, cb, cr): + """ + Args: + y(tensor): y channel image + cb(tensor): cb channel + cr(tensor): cr channel + + Returns: + Tensor: batch x height x width x 3 + """ + + def repeat(x, k=2): + height, width = x.shape[1:3] + x = x.unsqueeze(-1) + x = x.repeat(1, 1, k, k) + x = x.view(-1, height * k, width * k) + return x + + cb = repeat(cb) + cr = repeat(cr) + return torch.cat([y.unsqueeze(3), cb.unsqueeze(3), cr.unsqueeze(3)], dim=3) + + +class YCbCr2RGBJpeg(nn.Module): + """Converts YCbCr image to RGB JPEG + """ + + def __init__(self): + super(YCbCr2RGBJpeg, self).__init__() + + matrix = np.array([[1., 0., 1.402], [1, -0.344136, -0.714136], [1, 1.772, 0]], dtype=np.float32).T + self.shift = nn.Parameter(torch.tensor([0, -128., -128.])) + self.matrix = nn.Parameter(torch.from_numpy(matrix)) + + def forward(self, image): + """ + Args: + image(tensor): batch x height x width x 3 + + Returns: + Tensor: batch x 3 x height x width + """ + result = torch.tensordot(image + self.shift, self.matrix, dims=1) + return result.view(image.shape).permute(0, 3, 1, 2) + + +class DeCompressJpeg(nn.Module): + """Full JPEG decompression algorithm + + Args: + rounding(function): rounding function to use + """ + + def __init__(self, rounding=torch.round): + super(DeCompressJpeg, self).__init__() + self.c_dequantize = CDequantize() + self.y_dequantize = YDequantize() + self.idct = iDCT8x8() + self.merging = BlockMerging() + self.chroma = ChromaUpsampling() + self.colors = YCbCr2RGBJpeg() + + def forward(self, y, cb, cr, imgh, imgw, factor=1): + """ + Args: + compressed(dict(tensor)): batch x h*w/64 x 8 x 8 + imgh(int) + imgw(int) + factor(float) + + Returns: + Tensor: batch x 3 x height x width + """ + components = {'y': y, 'cb': cb, 'cr': cr} + for k in components.keys(): + if k in ('cb', 'cr'): + comp = self.c_dequantize(components[k], factor=factor) + height, width = int(imgh / 2), int(imgw / 2) + else: + comp = self.y_dequantize(components[k], factor=factor) + height, width = imgh, imgw + comp = self.idct(comp) + components[k] = self.merging(comp, height, width) + # + image = self.chroma(components['y'], components['cb'], components['cr']) + image = self.colors(image) + + image = torch.min(255 * torch.ones_like(image), torch.max(torch.zeros_like(image), image)) + return image / 255 + + +# ------------------------ main DiffJPEG ------------------------ # + + +class DiffJPEG(nn.Module): + """This JPEG algorithm result is slightly different from cv2. + DiffJPEG supports batch processing. + + Args: + differentiable(bool): If True, uses custom differentiable rounding function, if False, uses standard torch.round + """ + + def __init__(self, differentiable=True): + super(DiffJPEG, self).__init__() + if differentiable: + rounding = diff_round + else: + rounding = torch.round + + self.compress = CompressJpeg(rounding=rounding) + self.decompress = DeCompressJpeg(rounding=rounding) + + def forward(self, x, quality): + """ + Args: + x (Tensor): Input image, bchw, rgb, [0, 1] + quality(float): Quality factor for jpeg compression scheme. + """ + factor = quality + if isinstance(factor, (int, float)): + factor = quality_to_factor(factor) + else: + for i in range(factor.size(0)): + factor[i] = quality_to_factor(factor[i]) + h, w = x.size()[-2:] + h_pad, w_pad = 0, 0 + # why should use 16 + if h % 16 != 0: + h_pad = 16 - h % 16 + if w % 16 != 0: + w_pad = 16 - w % 16 + x = F.pad(x, (0, w_pad, 0, h_pad), mode='constant', value=0) + + y, cb, cr = self.compress(x, factor=factor) + recovered = self.decompress(y, cb, cr, (h + h_pad), (w + w_pad), factor=factor) + recovered = recovered[:, :, 0:h, 0:w] + return recovered + + +if __name__ == '__main__': + import cv2 + + from bsr.utils import img2tensor, tensor2img + + img_gt = cv2.imread('test.png') / 255. + + # -------------- cv2 -------------- # + encode_param = [int(cv2.IMWRITE_JPEG_QUALITY), 20] + _, encimg = cv2.imencode('.jpg', img_gt * 255., encode_param) + img_lq = np.float32(cv2.imdecode(encimg, 1)) + cv2.imwrite('cv2_JPEG_20.png', img_lq) + + # -------------- DiffJPEG -------------- # + jpeger = DiffJPEG(differentiable=False).cuda() + img_gt = img2tensor(img_gt) + img_gt = torch.stack([img_gt, img_gt]).cuda() + quality = img_gt.new_tensor([20, 40]) + out = jpeger(img_gt, quality=quality) + + cv2.imwrite('pt_JPEG_20.png', tensor2img(out[0])) + cv2.imwrite('pt_JPEG_40.png', tensor2img(out[1])) diff --git a/official/bsr/utils/dist_util.py b/official/bsr/utils/dist_util.py new file mode 100644 index 0000000000000000000000000000000000000000..380f155bc18cc5788d8b14fd18c0c0d748859de2 --- /dev/null +++ b/official/bsr/utils/dist_util.py @@ -0,0 +1,82 @@ +# Modified from https://github.com/open-mmlab/mmcv/blob/master/mmcv/runner/dist_utils.py # noqa: E501 +import functools +import os +import subprocess +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + + +def init_dist(launcher, backend='nccl', **kwargs): + if mp.get_start_method(allow_none=True) is None: + mp.set_start_method('spawn') + if launcher == 'pytorch': + _init_dist_pytorch(backend, **kwargs) + elif launcher == 'slurm': + _init_dist_slurm(backend, **kwargs) + else: + raise ValueError(f'Invalid launcher type: {launcher}') + + +def _init_dist_pytorch(backend, **kwargs): + rank = int(os.environ['RANK']) + num_gpus = torch.cuda.device_count() + torch.cuda.set_device(rank % num_gpus) + dist.init_process_group(backend=backend, **kwargs) + + +def _init_dist_slurm(backend, port=None): + """Initialize slurm distributed training environment. + + If argument ``port`` is not specified, then the master port will be system + environment variable ``MASTER_PORT``. If ``MASTER_PORT`` is not in system + environment variable, then a default port ``29500`` will be used. + + Args: + backend (str): Backend of torch.distributed. + port (int, optional): Master port. Defaults to None. + """ + proc_id = int(os.environ['SLURM_PROCID']) + ntasks = int(os.environ['SLURM_NTASKS']) + node_list = os.environ['SLURM_NODELIST'] + num_gpus = torch.cuda.device_count() + torch.cuda.set_device(proc_id % num_gpus) + addr = subprocess.getoutput(f'scontrol show hostname {node_list} | head -n1') + # specify master port + if port is not None: + os.environ['MASTER_PORT'] = str(port) + elif 'MASTER_PORT' in os.environ: + pass # use MASTER_PORT in the environment variable + else: + # 29500 is torch.distributed default port + os.environ['MASTER_PORT'] = '29500' + os.environ['MASTER_ADDR'] = addr + os.environ['WORLD_SIZE'] = str(ntasks) + os.environ['LOCAL_RANK'] = str(proc_id % num_gpus) + os.environ['RANK'] = str(proc_id) + dist.init_process_group(backend=backend) + + +def get_dist_info(): + if dist.is_available(): + initialized = dist.is_initialized() + else: + initialized = False + if initialized: + rank = dist.get_rank() + world_size = dist.get_world_size() + else: + rank = 0 + world_size = 1 + return rank, world_size + + +def master_only(func): + + @functools.wraps(func) + def wrapper(*args, **kwargs): + rank, _ = get_dist_info() + if rank == 0: + return func(*args, **kwargs) + + return wrapper diff --git a/official/bsr/utils/download_util.py b/official/bsr/utils/download_util.py new file mode 100644 index 0000000000000000000000000000000000000000..43fe80f79e0ad8354002ccd45b1ea4c3c125e983 --- /dev/null +++ b/official/bsr/utils/download_util.py @@ -0,0 +1,98 @@ +import math +import os +import requests +from torch.hub import download_url_to_file, get_dir +from tqdm import tqdm +from urllib.parse import urlparse + +from .misc import sizeof_fmt + + +def download_file_from_google_drive(file_id, save_path): + """Download files from google drive. + + Reference: https://stackoverflow.com/questions/25010369/wget-curl-large-file-from-google-drive + + Args: + file_id (str): File id. + save_path (str): Save path. + """ + + session = requests.Session() + URL = 'https://docs.google.com/uc?export=download' + params = {'id': file_id} + + response = session.get(URL, params=params, stream=True) + token = get_confirm_token(response) + if token: + params['confirm'] = token + response = session.get(URL, params=params, stream=True) + + # get file size + response_file_size = session.get(URL, params=params, stream=True, headers={'Range': 'bytes=0-2'}) + if 'Content-Range' in response_file_size.headers: + file_size = int(response_file_size.headers['Content-Range'].split('/')[1]) + else: + file_size = None + + save_response_content(response, save_path, file_size) + + +def get_confirm_token(response): + for key, value in response.cookies.items(): + if key.startswith('download_warning'): + return value + return None + + +def save_response_content(response, destination, file_size=None, chunk_size=32768): + if file_size is not None: + pbar = tqdm(total=math.ceil(file_size / chunk_size), unit='chunk') + + readable_file_size = sizeof_fmt(file_size) + else: + pbar = None + + with open(destination, 'wb') as f: + downloaded_size = 0 + for chunk in response.iter_content(chunk_size): + downloaded_size += chunk_size + if pbar is not None: + pbar.update(1) + pbar.set_description(f'Download {sizeof_fmt(downloaded_size)} / {readable_file_size}') + if chunk: # filter out keep-alive new chunks + f.write(chunk) + if pbar is not None: + pbar.close() + + +def load_file_from_url(url, model_dir=None, progress=True, file_name=None): + """Load file form http url, will download models if necessary. + + Reference: https://github.com/1adrianb/face-alignment/blob/master/face_alignment/utils.py + + Args: + url (str): URL to be downloaded. + model_dir (str): The path to save the downloaded model. Should be a full path. If None, use pytorch hub_dir. + Default: None. + progress (bool): Whether to show the download progress. Default: True. + file_name (str): The downloaded file name. If None, use the file name in the url. Default: None. + + Returns: + str: The path to the downloaded file. + """ + if model_dir is None: # use the pytorch hub_dir + hub_dir = get_dir() + model_dir = os.path.join(hub_dir, 'checkpoints') + + os.makedirs(model_dir, exist_ok=True) + + parts = urlparse(url) + filename = os.path.basename(parts.path) + if file_name is not None: + filename = file_name + cached_file = os.path.abspath(os.path.join(model_dir, filename)) + if not os.path.exists(cached_file): + print(f'Downloading: "{url}" to {cached_file}\n') + download_url_to_file(url, cached_file, hash_prefix=None, progress=progress) + return cached_file diff --git a/official/bsr/utils/file_client.py b/official/bsr/utils/file_client.py new file mode 100644 index 0000000000000000000000000000000000000000..8f6340e429dcf87f2f48059c292a427f1be97354 --- /dev/null +++ b/official/bsr/utils/file_client.py @@ -0,0 +1,167 @@ +# Modified from https://github.com/open-mmlab/mmcv/blob/master/mmcv/fileio/file_client.py # noqa: E501 +from abc import ABCMeta, abstractmethod + + +class BaseStorageBackend(metaclass=ABCMeta): + """Abstract class of storage backends. + + All backends need to implement two apis: ``get()`` and ``get_text()``. + ``get()`` reads the file as a byte stream and ``get_text()`` reads the file + as texts. + """ + + @abstractmethod + def get(self, filepath): + pass + + @abstractmethod + def get_text(self, filepath): + pass + + +class MemcachedBackend(BaseStorageBackend): + """Memcached storage backend. + + Attributes: + server_list_cfg (str): Config file for memcached server list. + client_cfg (str): Config file for memcached client. + sys_path (str | None): Additional path to be appended to `sys.path`. + Default: None. + """ + + def __init__(self, server_list_cfg, client_cfg, sys_path=None): + if sys_path is not None: + import sys + sys.path.append(sys_path) + try: + import mc + except ImportError: + raise ImportError('Please install memcached to enable MemcachedBackend.') + + self.server_list_cfg = server_list_cfg + self.client_cfg = client_cfg + self._client = mc.MemcachedClient.GetInstance(self.server_list_cfg, self.client_cfg) + # mc.pyvector servers as a point which points to a memory cache + self._mc_buffer = mc.pyvector() + + def get(self, filepath): + filepath = str(filepath) + import mc + self._client.Get(filepath, self._mc_buffer) + value_buf = mc.ConvertBuffer(self._mc_buffer) + return value_buf + + def get_text(self, filepath): + raise NotImplementedError + + +class HardDiskBackend(BaseStorageBackend): + """Raw hard disks storage backend.""" + + def get(self, filepath): + filepath = str(filepath) + with open(filepath, 'rb') as f: + value_buf = f.read() + return value_buf + + def get_text(self, filepath): + filepath = str(filepath) + with open(filepath, 'r') as f: + value_buf = f.read() + return value_buf + + +class LmdbBackend(BaseStorageBackend): + """Lmdb storage backend. + + Args: + db_paths (str | list[str]): Lmdb database paths. + client_keys (str | list[str]): Lmdb client keys. Default: 'default'. + readonly (bool, optional): Lmdb environment parameter. If True, + disallow any write operations. Default: True. + lock (bool, optional): Lmdb environment parameter. If False, when + concurrent access occurs, do not lock the database. Default: False. + readahead (bool, optional): Lmdb environment parameter. If False, + disable the OS filesystem readahead mechanism, which may improve + random read performance when a database is larger than RAM. + Default: False. + + Attributes: + db_paths (list): Lmdb database path. + _client (list): A list of several lmdb envs. + """ + + def __init__(self, db_paths, client_keys='default', readonly=True, lock=False, readahead=False, **kwargs): + try: + import lmdb + except ImportError: + raise ImportError('Please install lmdb to enable LmdbBackend.') + + if isinstance(client_keys, str): + client_keys = [client_keys] + + if isinstance(db_paths, list): + self.db_paths = [str(v) for v in db_paths] + elif isinstance(db_paths, str): + self.db_paths = [str(db_paths)] + assert len(client_keys) == len(self.db_paths), ('client_keys and db_paths should have the same length, ' + f'but received {len(client_keys)} and {len(self.db_paths)}.') + + self._client = {} + for client, path in zip(client_keys, self.db_paths): + self._client[client] = lmdb.open(path, readonly=readonly, lock=lock, readahead=readahead, **kwargs) + + def get(self, filepath, client_key): + """Get values according to the filepath from one lmdb named client_key. + + Args: + filepath (str | obj:`Path`): Here, filepath is the lmdb key. + client_key (str): Used for distinguishing different lmdb envs. + """ + filepath = str(filepath) + assert client_key in self._client, (f'client_key {client_key} is not in lmdb clients.') + client = self._client[client_key] + with client.begin(write=False) as txn: + value_buf = txn.get(filepath.encode('ascii')) + return value_buf + + def get_text(self, filepath): + raise NotImplementedError + + +class FileClient(object): + """A general file client to access files in different backend. + + The client loads a file or text in a specified backend from its path + and return it as a binary file. it can also register other backend + accessor with a given name and backend class. + + Attributes: + backend (str): The storage backend type. Options are "disk", + "memcached" and "lmdb". + client (:obj:`BaseStorageBackend`): The backend object. + """ + + _backends = { + 'disk': HardDiskBackend, + 'memcached': MemcachedBackend, + 'lmdb': LmdbBackend, + } + + def __init__(self, backend='disk', **kwargs): + if backend not in self._backends: + raise ValueError(f'Backend {backend} is not supported. Currently supported ones' + f' are {list(self._backends.keys())}') + self.backend = backend + self.client = self._backends[backend](**kwargs) + + def get(self, filepath, client_key='default'): + # client_key is used only for lmdb, where different fileclients have + # different lmdb environments. + if self.backend == 'lmdb': + return self.client.get(filepath, client_key) + else: + return self.client.get(filepath) + + def get_text(self, filepath): + return self.client.get_text(filepath) diff --git a/official/bsr/utils/flow_util.py b/official/bsr/utils/flow_util.py new file mode 100644 index 0000000000000000000000000000000000000000..d133012fddf0dd338ea4764cff4f83a02a36781a --- /dev/null +++ b/official/bsr/utils/flow_util.py @@ -0,0 +1,170 @@ +# Modified from https://github.com/open-mmlab/mmcv/blob/master/mmcv/video/optflow.py # noqa: E501 +import cv2 +import numpy as np +import os + + +def flowread(flow_path, quantize=False, concat_axis=0, *args, **kwargs): + """Read an optical flow map. + + Args: + flow_path (ndarray or str): Flow path. + quantize (bool): whether to read quantized pair, if set to True, + remaining args will be passed to :func:`dequantize_flow`. + concat_axis (int): The axis that dx and dy are concatenated, + can be either 0 or 1. Ignored if quantize is False. + + Returns: + ndarray: Optical flow represented as a (h, w, 2) numpy array + """ + if quantize: + assert concat_axis in [0, 1] + cat_flow = cv2.imread(flow_path, cv2.IMREAD_UNCHANGED) + if cat_flow.ndim != 2: + raise IOError(f'{flow_path} is not a valid quantized flow file, its dimension is {cat_flow.ndim}.') + assert cat_flow.shape[concat_axis] % 2 == 0 + dx, dy = np.split(cat_flow, 2, axis=concat_axis) + flow = dequantize_flow(dx, dy, *args, **kwargs) + else: + with open(flow_path, 'rb') as f: + try: + header = f.read(4).decode('utf-8') + except Exception: + raise IOError(f'Invalid flow file: {flow_path}') + else: + if header != 'PIEH': + raise IOError(f'Invalid flow file: {flow_path}, header does not contain PIEH') + + w = np.fromfile(f, np.int32, 1).squeeze() + h = np.fromfile(f, np.int32, 1).squeeze() + flow = np.fromfile(f, np.float32, w * h * 2).reshape((h, w, 2)) + + return flow.astype(np.float32) + + +def flowwrite(flow, filename, quantize=False, concat_axis=0, *args, **kwargs): + """Write optical flow to file. + + If the flow is not quantized, it will be saved as a .flo file losslessly, + otherwise a jpeg image which is lossy but of much smaller size. (dx and dy + will be concatenated horizontally into a single image if quantize is True.) + + Args: + flow (ndarray): (h, w, 2) array of optical flow. + filename (str): Output filepath. + quantize (bool): Whether to quantize the flow and save it to 2 jpeg + images. If set to True, remaining args will be passed to + :func:`quantize_flow`. + concat_axis (int): The axis that dx and dy are concatenated, + can be either 0 or 1. Ignored if quantize is False. + """ + if not quantize: + with open(filename, 'wb') as f: + f.write('PIEH'.encode('utf-8')) + np.array([flow.shape[1], flow.shape[0]], dtype=np.int32).tofile(f) + flow = flow.astype(np.float32) + flow.tofile(f) + f.flush() + else: + assert concat_axis in [0, 1] + dx, dy = quantize_flow(flow, *args, **kwargs) + dxdy = np.concatenate((dx, dy), axis=concat_axis) + os.makedirs(os.path.dirname(filename), exist_ok=True) + cv2.imwrite(filename, dxdy) + + +def quantize_flow(flow, max_val=0.02, norm=True): + """Quantize flow to [0, 255]. + + After this step, the size of flow will be much smaller, and can be + dumped as jpeg images. + + Args: + flow (ndarray): (h, w, 2) array of optical flow. + max_val (float): Maximum value of flow, values beyond + [-max_val, max_val] will be truncated. + norm (bool): Whether to divide flow values by image width/height. + + Returns: + tuple[ndarray]: Quantized dx and dy. + """ + h, w, _ = flow.shape + dx = flow[..., 0] + dy = flow[..., 1] + if norm: + dx = dx / w # avoid inplace operations + dy = dy / h + # use 255 levels instead of 256 to make sure 0 is 0 after dequantization. + flow_comps = [quantize(d, -max_val, max_val, 255, np.uint8) for d in [dx, dy]] + return tuple(flow_comps) + + +def dequantize_flow(dx, dy, max_val=0.02, denorm=True): + """Recover from quantized flow. + + Args: + dx (ndarray): Quantized dx. + dy (ndarray): Quantized dy. + max_val (float): Maximum value used when quantizing. + denorm (bool): Whether to multiply flow values with width/height. + + Returns: + ndarray: Dequantized flow. + """ + assert dx.shape == dy.shape + assert dx.ndim == 2 or (dx.ndim == 3 and dx.shape[-1] == 1) + + dx, dy = [dequantize(d, -max_val, max_val, 255) for d in [dx, dy]] + + if denorm: + dx *= dx.shape[1] + dy *= dx.shape[0] + flow = np.dstack((dx, dy)) + return flow + + +def quantize(arr, min_val, max_val, levels, dtype=np.int64): + """Quantize an array of (-inf, inf) to [0, levels-1]. + + Args: + arr (ndarray): Input array. + min_val (scalar): Minimum value to be clipped. + max_val (scalar): Maximum value to be clipped. + levels (int): Quantization levels. + dtype (np.type): The type of the quantized array. + + Returns: + tuple: Quantized array. + """ + if not (isinstance(levels, int) and levels > 1): + raise ValueError(f'levels must be a positive integer, but got {levels}') + if min_val >= max_val: + raise ValueError(f'min_val ({min_val}) must be smaller than max_val ({max_val})') + + arr = np.clip(arr, min_val, max_val) - min_val + quantized_arr = np.minimum(np.floor(levels * arr / (max_val - min_val)).astype(dtype), levels - 1) + + return quantized_arr + + +def dequantize(arr, min_val, max_val, levels, dtype=np.float64): + """Dequantize an array. + + Args: + arr (ndarray): Input array. + min_val (scalar): Minimum value to be clipped. + max_val (scalar): Maximum value to be clipped. + levels (int): Quantization levels. + dtype (np.type): The type of the dequantized array. + + Returns: + tuple: Dequantized array. + """ + if not (isinstance(levels, int) and levels > 1): + raise ValueError(f'levels must be a positive integer, but got {levels}') + if min_val >= max_val: + raise ValueError(f'min_val ({min_val}) must be smaller than max_val ({max_val})') + + dequantized_arr = (arr + 0.5).astype(dtype) * (max_val - min_val) / levels + min_val + + return dequantized_arr diff --git a/official/bsr/utils/img_process_util.py b/official/bsr/utils/img_process_util.py new file mode 100644 index 0000000000000000000000000000000000000000..fb5fbc9468ca1861fe7d6eae28128172b9e70001 --- /dev/null +++ b/official/bsr/utils/img_process_util.py @@ -0,0 +1,83 @@ +import cv2 +import numpy as np +import torch +from torch.nn import functional as F + + +def filter2D(img, kernel): + """PyTorch version of cv2.filter2D + + Args: + img (Tensor): (b, c, h, w) + kernel (Tensor): (b, k, k) + """ + k = kernel.size(-1) + b, c, h, w = img.size() + if k % 2 == 1: + img = F.pad(img, (k // 2, k // 2, k // 2, k // 2), mode='reflect') + else: + raise ValueError('Wrong kernel size') + + ph, pw = img.size()[-2:] + + if kernel.size(0) == 1: + # apply the same kernel to all batch images + img = img.view(b * c, 1, ph, pw) + kernel = kernel.view(1, 1, k, k) + return F.conv2d(img, kernel, padding=0).view(b, c, h, w) + else: + img = img.view(1, b * c, ph, pw) + kernel = kernel.view(b, 1, k, k).repeat(1, c, 1, 1).view(b * c, 1, k, k) + return F.conv2d(img, kernel, groups=b * c).view(b, c, h, w) + + +def usm_sharp(img, weight=0.5, radius=50, threshold=10): + """USM sharpening. + + Input image: I; Blurry image: B. + 1. sharp = I + weight * (I - B) + 2. Mask = 1 if abs(I - B) > threshold, else: 0 + 3. Blur mask: + 4. Out = Mask * sharp + (1 - Mask) * I + + + Args: + img (Numpy array): Input image, HWC, BGR; float32, [0, 1]. + weight (float): Sharp weight. Default: 1. + radius (float): Kernel size of Gaussian blur. Default: 50. + threshold (int): + """ + if radius % 2 == 0: + radius += 1 + blur = cv2.GaussianBlur(img, (radius, radius), 0) + residual = img - blur + mask = np.abs(residual) * 255 > threshold + mask = mask.astype('float32') + soft_mask = cv2.GaussianBlur(mask, (radius, radius), 0) + + sharp = img + weight * residual + sharp = np.clip(sharp, 0, 1) + return soft_mask * sharp + (1 - soft_mask) * img + + +class USMSharp(torch.nn.Module): + + def __init__(self, radius=50, sigma=0): + super(USMSharp, self).__init__() + if radius % 2 == 0: + radius += 1 + self.radius = radius + kernel = cv2.getGaussianKernel(radius, sigma) + kernel = torch.FloatTensor(np.dot(kernel, kernel.transpose())).unsqueeze_(0) + self.register_buffer('kernel', kernel) + + def forward(self, img, weight=0.5, threshold=10): + blur = filter2D(img, self.kernel) + residual = img - blur + + mask = torch.abs(residual) * 255 > threshold + mask = mask.float() + soft_mask = filter2D(mask, self.kernel) + sharp = img + weight * residual + sharp = torch.clip(sharp, 0, 1) + return soft_mask * sharp + (1 - soft_mask) * img diff --git a/official/bsr/utils/img_util.py b/official/bsr/utils/img_util.py new file mode 100644 index 0000000000000000000000000000000000000000..3ad2be2c5556ddb6076eb01a44c487447dc7fcf1 --- /dev/null +++ b/official/bsr/utils/img_util.py @@ -0,0 +1,172 @@ +import cv2 +import math +import numpy as np +import os +import torch +from torchvision.utils import make_grid + + +def img2tensor(imgs, bgr2rgb=True, float32=True): + """Numpy array to tensor. + + Args: + imgs (list[ndarray] | ndarray): Input images. + bgr2rgb (bool): Whether to change bgr to rgb. + float32 (bool): Whether to change to float32. + + Returns: + list[tensor] | tensor: Tensor images. If returned results only have + one element, just return tensor. + """ + + def _totensor(img, bgr2rgb, float32): + if img.shape[2] == 3 and bgr2rgb: + if img.dtype == 'float64': + img = img.astype('float32') + img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) + img = torch.from_numpy(img.transpose(2, 0, 1)) + if float32: + img = img.float() + return img + + if isinstance(imgs, list): + return [_totensor(img, bgr2rgb, float32) for img in imgs] + else: + return _totensor(imgs, bgr2rgb, float32) + + +def tensor2img(tensor, rgb2bgr=True, out_type=np.uint8, min_max=(0, 1)): + """Convert torch Tensors into image numpy arrays. + + After clamping to [min, max], values will be normalized to [0, 1]. + + Args: + tensor (Tensor or list[Tensor]): Accept shapes: + 1) 4D mini-batch Tensor of shape (B x 3/1 x H x W); + 2) 3D Tensor of shape (3/1 x H x W); + 3) 2D Tensor of shape (H x W). + Tensor channel should be in RGB order. + rgb2bgr (bool): Whether to change rgb to bgr. + out_type (numpy type): output types. If ``np.uint8``, transform outputs + to uint8 type with range [0, 255]; otherwise, float type with + range [0, 1]. Default: ``np.uint8``. + min_max (tuple[int]): min and max values for clamp. + + Returns: + (Tensor or list): 3D ndarray of shape (H x W x C) OR 2D ndarray of + shape (H x W). The channel order is BGR. + """ + if not (torch.is_tensor(tensor) or (isinstance(tensor, list) and all(torch.is_tensor(t) for t in tensor))): + raise TypeError(f'tensor or list of tensors expected, got {type(tensor)}') + + if torch.is_tensor(tensor): + tensor = [tensor] + result = [] + for _tensor in tensor: + _tensor = _tensor.squeeze(0).float().detach().cpu().clamp_(*min_max) + _tensor = (_tensor - min_max[0]) / (min_max[1] - min_max[0]) + + n_dim = _tensor.dim() + if n_dim == 4: + img_np = make_grid(_tensor, nrow=int(math.sqrt(_tensor.size(0))), normalize=False).numpy() + img_np = img_np.transpose(1, 2, 0) + if rgb2bgr: + img_np = cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR) + elif n_dim == 3: + img_np = _tensor.numpy() + img_np = img_np.transpose(1, 2, 0) + if img_np.shape[2] == 1: # gray image + img_np = np.squeeze(img_np, axis=2) + else: + if rgb2bgr: + img_np = cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR) + elif n_dim == 2: + img_np = _tensor.numpy() + else: + raise TypeError(f'Only support 4D, 3D or 2D tensor. But received with dimension: {n_dim}') + if out_type == np.uint8: + # Unlike MATLAB, numpy.unit8() WILL NOT round by default. + img_np = (img_np * 255.0).round() + img_np = img_np.astype(out_type) + result.append(img_np) + if len(result) == 1: + result = result[0] + return result + + +def tensor2img_fast(tensor, rgb2bgr=True, min_max=(0, 1)): + """This implementation is slightly faster than tensor2img. + It now only supports torch tensor with shape (1, c, h, w). + + Args: + tensor (Tensor): Now only support torch tensor with (1, c, h, w). + rgb2bgr (bool): Whether to change rgb to bgr. Default: True. + min_max (tuple[int]): min and max values for clamp. + """ + output = tensor.squeeze(0).detach().clamp_(*min_max).permute(1, 2, 0) + output = (output - min_max[0]) / (min_max[1] - min_max[0]) * 255 + output = output.type(torch.uint8).cpu().numpy() + if rgb2bgr: + output = cv2.cvtColor(output, cv2.COLOR_RGB2BGR) + return output + + +def imfrombytes(content, flag='color', float32=False): + """Read an image from bytes. + + Args: + content (bytes): Image bytes got from files or other streams. + flag (str): Flags specifying the color type of a loaded image, + candidates are `color`, `grayscale` and `unchanged`. + float32 (bool): Whether to change to float32., If True, will also norm + to [0, 1]. Default: False. + + Returns: + ndarray: Loaded image array. + """ + img_np = np.frombuffer(content, np.uint8) + imread_flags = {'color': cv2.IMREAD_COLOR, 'grayscale': cv2.IMREAD_GRAYSCALE, 'unchanged': cv2.IMREAD_UNCHANGED} + img = cv2.imdecode(img_np, imread_flags[flag]) + if float32: + img = img.astype(np.float32) / 255. + return img + + +def imwrite(img, file_path, params=None, auto_mkdir=True): + """Write image to file. + + Args: + img (ndarray): Image array to be written. + file_path (str): Image file path. + params (None or list): Same as opencv's :func:`imwrite` interface. + auto_mkdir (bool): If the parent folder of `file_path` does not exist, + whether to create it automatically. + + Returns: + bool: Successful or not. + """ + if auto_mkdir: + dir_name = os.path.abspath(os.path.dirname(file_path)) + os.makedirs(dir_name, exist_ok=True) + ok = cv2.imwrite(file_path, img, params) + if not ok: + raise IOError('Failed in writing images.') + + +def crop_border(imgs, crop_border): + """Crop borders of images. + + Args: + imgs (list[ndarray] | ndarray): Images with shape (h, w, c). + crop_border (int): Crop border for each end of height and weight. + + Returns: + list[ndarray]: Cropped images. + """ + if crop_border == 0: + return imgs + else: + if isinstance(imgs, list): + return [v[crop_border:-crop_border, crop_border:-crop_border, ...] for v in imgs] + else: + return imgs[crop_border:-crop_border, crop_border:-crop_border, ...] diff --git a/official/bsr/utils/lmdb_util.py b/official/bsr/utils/lmdb_util.py new file mode 100644 index 0000000000000000000000000000000000000000..591182df8ebb469f23443cd45f540e6010b26fc5 --- /dev/null +++ b/official/bsr/utils/lmdb_util.py @@ -0,0 +1,199 @@ +import cv2 +import lmdb +import sys +from multiprocessing import Pool +from os import path as osp +from tqdm import tqdm + + +def make_lmdb_from_imgs(data_path, + lmdb_path, + img_path_list, + keys, + batch=5000, + compress_level=1, + multiprocessing_read=False, + n_thread=40, + map_size=None): + """Make lmdb from images. + + Contents of lmdb. The file structure is: + + :: + + example.lmdb + ├── data.mdb + ├── lock.mdb + ├── meta_info.txt + + The data.mdb and lock.mdb are standard lmdb files and you can refer to + https://lmdb.readthedocs.io/en/release/ for more details. + + The meta_info.txt is a specified txt file to record the meta information + of our datasets. It will be automatically created when preparing + datasets by our provided dataset tools. + Each line in the txt file records 1)image name (with extension), + 2)image shape, and 3)compression level, separated by a white space. + + For example, the meta information could be: + `000_00000000.png (720,1280,3) 1`, which means: + 1) image name (with extension): 000_00000000.png; + 2) image shape: (720,1280,3); + 3) compression level: 1 + + We use the image name without extension as the lmdb key. + + If `multiprocessing_read` is True, it will read all the images to memory + using multiprocessing. Thus, your server needs to have enough memory. + + Args: + data_path (str): Data path for reading images. + lmdb_path (str): Lmdb save path. + img_path_list (str): Image path list. + keys (str): Used for lmdb keys. + batch (int): After processing batch images, lmdb commits. + Default: 5000. + compress_level (int): Compress level when encoding images. Default: 1. + multiprocessing_read (bool): Whether use multiprocessing to read all + the images to memory. Default: False. + n_thread (int): For multiprocessing. + map_size (int | None): Map size for lmdb env. If None, use the + estimated size from images. Default: None + """ + + assert len(img_path_list) == len(keys), ('img_path_list and keys should have the same length, ' + f'but got {len(img_path_list)} and {len(keys)}') + print(f'Create lmdb for {data_path}, save to {lmdb_path}...') + print(f'Totoal images: {len(img_path_list)}') + if not lmdb_path.endswith('.lmdb'): + raise ValueError("lmdb_path must end with '.lmdb'.") + if osp.exists(lmdb_path): + print(f'Folder {lmdb_path} already exists. Exit.') + sys.exit(1) + + if multiprocessing_read: + # read all the images to memory (multiprocessing) + dataset = {} # use dict to keep the order for multiprocessing + shapes = {} + print(f'Read images with multiprocessing, #thread: {n_thread} ...') + pbar = tqdm(total=len(img_path_list), unit='image') + + def callback(arg): + """get the image data and update pbar.""" + key, dataset[key], shapes[key] = arg + pbar.update(1) + pbar.set_description(f'Read {key}') + + pool = Pool(n_thread) + for path, key in zip(img_path_list, keys): + pool.apply_async(read_img_worker, args=(osp.join(data_path, path), key, compress_level), callback=callback) + pool.close() + pool.join() + pbar.close() + print(f'Finish reading {len(img_path_list)} images.') + + # create lmdb environment + if map_size is None: + # obtain data size for one image + img = cv2.imread(osp.join(data_path, img_path_list[0]), cv2.IMREAD_UNCHANGED) + _, img_byte = cv2.imencode('.png', img, [cv2.IMWRITE_PNG_COMPRESSION, compress_level]) + data_size_per_img = img_byte.nbytes + print('Data size per image is: ', data_size_per_img) + data_size = data_size_per_img * len(img_path_list) + map_size = data_size * 10 + + env = lmdb.open(lmdb_path, map_size=map_size) + + # write data to lmdb + pbar = tqdm(total=len(img_path_list), unit='chunk') + txn = env.begin(write=True) + txt_file = open(osp.join(lmdb_path, 'meta_info.txt'), 'w') + for idx, (path, key) in enumerate(zip(img_path_list, keys)): + pbar.update(1) + pbar.set_description(f'Write {key}') + key_byte = key.encode('ascii') + if multiprocessing_read: + img_byte = dataset[key] + h, w, c = shapes[key] + else: + _, img_byte, img_shape = read_img_worker(osp.join(data_path, path), key, compress_level) + h, w, c = img_shape + + txn.put(key_byte, img_byte) + # write meta information + txt_file.write(f'{key}.png ({h},{w},{c}) {compress_level}\n') + if idx % batch == 0: + txn.commit() + txn = env.begin(write=True) + pbar.close() + txn.commit() + env.close() + txt_file.close() + print('\nFinish writing lmdb.') + + +def read_img_worker(path, key, compress_level): + """Read image worker. + + Args: + path (str): Image path. + key (str): Image key. + compress_level (int): Compress level when encoding images. + + Returns: + str: Image key. + byte: Image byte. + tuple[int]: Image shape. + """ + + img = cv2.imread(path, cv2.IMREAD_UNCHANGED) + if img.ndim == 2: + h, w = img.shape + c = 1 + else: + h, w, c = img.shape + _, img_byte = cv2.imencode('.png', img, [cv2.IMWRITE_PNG_COMPRESSION, compress_level]) + return (key, img_byte, (h, w, c)) + + +class LmdbMaker(): + """LMDB Maker. + + Args: + lmdb_path (str): Lmdb save path. + map_size (int): Map size for lmdb env. Default: 1024 ** 4, 1TB. + batch (int): After processing batch images, lmdb commits. + Default: 5000. + compress_level (int): Compress level when encoding images. Default: 1. + """ + + def __init__(self, lmdb_path, map_size=1024**4, batch=5000, compress_level=1): + if not lmdb_path.endswith('.lmdb'): + raise ValueError("lmdb_path must end with '.lmdb'.") + if osp.exists(lmdb_path): + print(f'Folder {lmdb_path} already exists. Exit.') + sys.exit(1) + + self.lmdb_path = lmdb_path + self.batch = batch + self.compress_level = compress_level + self.env = lmdb.open(lmdb_path, map_size=map_size) + self.txn = self.env.begin(write=True) + self.txt_file = open(osp.join(lmdb_path, 'meta_info.txt'), 'w') + self.counter = 0 + + def put(self, img_byte, key, img_shape): + self.counter += 1 + key_byte = key.encode('ascii') + self.txn.put(key_byte, img_byte) + # write meta information + h, w, c = img_shape + self.txt_file.write(f'{key}.png ({h},{w},{c}) {self.compress_level}\n') + if self.counter % self.batch == 0: + self.txn.commit() + self.txn = self.env.begin(write=True) + + def close(self): + self.txn.commit() + self.env.close() + self.txt_file.close() diff --git a/official/bsr/utils/logger.py b/official/bsr/utils/logger.py new file mode 100644 index 0000000000000000000000000000000000000000..6c0592d2ce50822e8269cbe222cbcd66c04dbb77 --- /dev/null +++ b/official/bsr/utils/logger.py @@ -0,0 +1,213 @@ +import datetime +import logging +import time + +from .dist_util import get_dist_info, master_only + +initialized_logger = {} + + +class AvgTimer(): + + def __init__(self, window=200): + self.window = window # average window + self.current_time = 0 + self.total_time = 0 + self.count = 0 + self.avg_time = 0 + self.start() + + def start(self): + self.start_time = self.tic = time.time() + + def record(self): + self.count += 1 + self.toc = time.time() + self.current_time = self.toc - self.tic + self.total_time += self.current_time + # calculate average time + self.avg_time = self.total_time / self.count + + # reset + if self.count > self.window: + self.count = 0 + self.total_time = 0 + + self.tic = time.time() + + def get_current_time(self): + return self.current_time + + def get_avg_time(self): + return self.avg_time + + +class MessageLogger(): + """Message logger for printing. + + Args: + opt (dict): Config. It contains the following keys: + name (str): Exp name. + logger (dict): Contains 'print_freq' (str) for logger interval. + train (dict): Contains 'total_iter' (int) for total iters. + use_tb_logger (bool): Use tensorboard logger. + start_iter (int): Start iter. Default: 1. + tb_logger (obj:`tb_logger`): Tensorboard logger. Default: None. + """ + + def __init__(self, opt, start_iter=1, tb_logger=None): + self.exp_name = opt['name'] + self.interval = opt['logger']['print_freq'] + self.start_iter = start_iter + self.max_iters = opt['train']['total_iter'] + self.use_tb_logger = opt['logger']['use_tb_logger'] + self.tb_logger = tb_logger + self.start_time = time.time() + self.logger = get_root_logger() + + def reset_start_time(self): + self.start_time = time.time() + + @master_only + def __call__(self, log_vars): + """Format logging message. + + Args: + log_vars (dict): It contains the following keys: + epoch (int): Epoch number. + iter (int): Current iter. + lrs (list): List for learning rates. + + time (float): Iter time. + data_time (float): Data time for each iter. + """ + # epoch, iter, learning rates + epoch = log_vars.pop('epoch') + current_iter = log_vars.pop('iter') + lrs = log_vars.pop('lrs') + + message = (f'[{self.exp_name[:5]}..][epoch:{epoch:3d}, iter:{current_iter:8,d}, lr:(') + for v in lrs: + message += f'{v:.3e},' + message += ')] ' + + # time and estimated time + if 'time' in log_vars.keys(): + iter_time = log_vars.pop('time') + data_time = log_vars.pop('data_time') + + total_time = time.time() - self.start_time + time_sec_avg = total_time / (current_iter - self.start_iter + 1) + eta_sec = time_sec_avg * (self.max_iters - current_iter - 1) + eta_str = str(datetime.timedelta(seconds=int(eta_sec))) + message += f'[eta: {eta_str}, ' + message += f'time (data): {iter_time:.3f} ({data_time:.3f})] ' + + # other items, especially losses + for k, v in log_vars.items(): + message += f'{k}: {v:.4e} ' + # tensorboard logger + if self.use_tb_logger and 'debug' not in self.exp_name: + if k.startswith('l_'): + self.tb_logger.add_scalar(f'losses/{k}', v, current_iter) + else: + self.tb_logger.add_scalar(k, v, current_iter) + self.logger.info(message) + + +@master_only +def init_tb_logger(log_dir): + from torch.utils.tensorboard import SummaryWriter + tb_logger = SummaryWriter(log_dir=log_dir) + return tb_logger + + +@master_only +def init_wandb_logger(opt): + """We now only use wandb to sync tensorboard log.""" + import wandb + logger = get_root_logger() + + project = opt['logger']['wandb']['project'] + resume_id = opt['logger']['wandb'].get('resume_id') + if resume_id: + wandb_id = resume_id + resume = 'allow' + logger.warning(f'Resume wandb logger with id={wandb_id}.') + else: + wandb_id = wandb.util.generate_id() + resume = 'never' + + wandb.init(id=wandb_id, resume=resume, name=opt['name'], config=opt, project=project, sync_tensorboard=True) + + logger.info(f'Use wandb logger with id={wandb_id}; project={project}.') + + +def get_root_logger(logger_name='basicsr', log_level=logging.INFO, log_file=None): + """Get the root logger. + + The logger will be initialized if it has not been initialized. By default a + StreamHandler will be added. If `log_file` is specified, a FileHandler will + also be added. + + Args: + logger_name (str): root logger name. Default: 'basicsr'. + log_file (str | None): The log filename. If specified, a FileHandler + will be added to the root logger. + log_level (int): The root logger level. Note that only the process of + rank 0 is affected, while other processes will set the level to + "Error" and be silent most of the time. + + Returns: + logging.Logger: The root logger. + """ + logger = logging.getLogger(logger_name) + # if the logger has been initialized, just return it + if logger_name in initialized_logger: + return logger + + format_str = '%(asctime)s %(levelname)s: %(message)s' + stream_handler = logging.StreamHandler() + stream_handler.setFormatter(logging.Formatter(format_str)) + logger.addHandler(stream_handler) + logger.propagate = False + rank, _ = get_dist_info() + if rank != 0: + logger.setLevel('ERROR') + elif log_file is not None: + logger.setLevel(log_level) + # add file handler + file_handler = logging.FileHandler(log_file, 'w') + file_handler.setFormatter(logging.Formatter(format_str)) + file_handler.setLevel(log_level) + logger.addHandler(file_handler) + initialized_logger[logger_name] = True + return logger + + +def get_env_info(): + """Get environment information. + + Currently, only log the software version. + """ + import torch + import torchvision + + from basicsr.version import __version__ + msg = r""" + ____ _ _____ ____ + / __ ) ____ _ _____ (_)_____/ ___/ / __ \ + / __ |/ __ `// ___// // ___/\__ \ / /_/ / + / /_/ // /_/ /(__ )/ // /__ ___/ // _, _/ + /_____/ \__,_//____//_/ \___//____//_/ |_| + ______ __ __ __ __ + / ____/____ ____ ____/ / / / __ __ _____ / /__ / / + / / __ / __ \ / __ \ / __ / / / / / / // ___// //_/ / / + / /_/ // /_/ // /_/ // /_/ / / /___/ /_/ // /__ / /< /_/ + \____/ \____/ \____/ \____/ /_____/\____/ \___//_/|_| (_) + """ + msg += ('\nVersion Information: ' + f'\n\tBasicSR: {__version__}' + f'\n\tPyTorch: {torch.__version__}' + f'\n\tTorchVision: {torchvision.__version__}') + return msg diff --git a/official/bsr/utils/matlab_functions.py b/official/bsr/utils/matlab_functions.py new file mode 100644 index 0000000000000000000000000000000000000000..6d0b8cd891338329658e950633745c9a8b2eaad6 --- /dev/null +++ b/official/bsr/utils/matlab_functions.py @@ -0,0 +1,178 @@ +import math +import numpy as np +import torch + + +def cubic(x): + """cubic function used for calculate_weights_indices.""" + absx = torch.abs(x) + absx2 = absx**2 + absx3 = absx**3 + return (1.5 * absx3 - 2.5 * absx2 + 1) * ( + (absx <= 1).type_as(absx)) + (-0.5 * absx3 + 2.5 * absx2 - 4 * absx + 2) * (((absx > 1) * + (absx <= 2)).type_as(absx)) + + +def calculate_weights_indices(in_length, out_length, scale, kernel, kernel_width, antialiasing): + """Calculate weights and indices, used for imresize function. + + Args: + in_length (int): Input length. + out_length (int): Output length. + scale (float): Scale factor. + kernel_width (int): Kernel width. + antialisaing (bool): Whether to apply anti-aliasing when downsampling. + """ + + if (scale < 1) and antialiasing: + # Use a modified kernel (larger kernel width) to simultaneously + # interpolate and antialias + kernel_width = kernel_width / scale + + # Output-space coordinates + x = torch.linspace(1, out_length, out_length) + + # Input-space coordinates. Calculate the inverse mapping such that 0.5 + # in output space maps to 0.5 in input space, and 0.5 + scale in output + # space maps to 1.5 in input space. + u = x / scale + 0.5 * (1 - 1 / scale) + + # What is the left-most pixel that can be involved in the computation? + left = torch.floor(u - kernel_width / 2) + + # What is the maximum number of pixels that can be involved in the + # computation? Note: it's OK to use an extra pixel here; if the + # corresponding weights are all zero, it will be eliminated at the end + # of this function. + p = math.ceil(kernel_width) + 2 + + # The indices of the input pixels involved in computing the k-th output + # pixel are in row k of the indices matrix. + indices = left.view(out_length, 1).expand(out_length, p) + torch.linspace(0, p - 1, p).view(1, p).expand( + out_length, p) + + # The weights used to compute the k-th output pixel are in row k of the + # weights matrix. + distance_to_center = u.view(out_length, 1).expand(out_length, p) - indices + + # apply cubic kernel + if (scale < 1) and antialiasing: + weights = scale * cubic(distance_to_center * scale) + else: + weights = cubic(distance_to_center) + + # Normalize the weights matrix so that each row sums to 1. + weights_sum = torch.sum(weights, 1).view(out_length, 1) + weights = weights / weights_sum.expand(out_length, p) + + # If a column in weights is all zero, get rid of it. only consider the + # first and last column. + weights_zero_tmp = torch.sum((weights == 0), 0) + if not math.isclose(weights_zero_tmp[0], 0, rel_tol=1e-6): + indices = indices.narrow(1, 1, p - 2) + weights = weights.narrow(1, 1, p - 2) + if not math.isclose(weights_zero_tmp[-1], 0, rel_tol=1e-6): + indices = indices.narrow(1, 0, p - 2) + weights = weights.narrow(1, 0, p - 2) + weights = weights.contiguous() + indices = indices.contiguous() + sym_len_s = -indices.min() + 1 + sym_len_e = indices.max() - in_length + indices = indices + sym_len_s - 1 + return weights, indices, int(sym_len_s), int(sym_len_e) + + +@torch.no_grad() +def imresize(img, scale, antialiasing=True): + """imresize function same as MATLAB. + + It now only supports bicubic. + The same scale applies for both height and width. + + Args: + img (Tensor | Numpy array): + Tensor: Input image with shape (c, h, w), [0, 1] range. + Numpy: Input image with shape (h, w, c), [0, 1] range. + scale (float): Scale factor. The same scale applies for both height + and width. + antialisaing (bool): Whether to apply anti-aliasing when downsampling. + Default: True. + + Returns: + Tensor: Output image with shape (c, h, w), [0, 1] range, w/o round. + """ + squeeze_flag = False + if type(img).__module__ == np.__name__: # numpy type + numpy_type = True + if img.ndim == 2: + img = img[:, :, None] + squeeze_flag = True + img = torch.from_numpy(img.transpose(2, 0, 1)).float() + else: + numpy_type = False + if img.ndim == 2: + img = img.unsqueeze(0) + squeeze_flag = True + + in_c, in_h, in_w = img.size() + out_h, out_w = math.ceil(in_h * scale), math.ceil(in_w * scale) + kernel_width = 4 + kernel = 'cubic' + + # get weights and indices + weights_h, indices_h, sym_len_hs, sym_len_he = calculate_weights_indices(in_h, out_h, scale, kernel, kernel_width, + antialiasing) + weights_w, indices_w, sym_len_ws, sym_len_we = calculate_weights_indices(in_w, out_w, scale, kernel, kernel_width, + antialiasing) + # process H dimension + # symmetric copying + img_aug = torch.FloatTensor(in_c, in_h + sym_len_hs + sym_len_he, in_w) + img_aug.narrow(1, sym_len_hs, in_h).copy_(img) + + sym_patch = img[:, :sym_len_hs, :] + inv_idx = torch.arange(sym_patch.size(1) - 1, -1, -1).long() + sym_patch_inv = sym_patch.index_select(1, inv_idx) + img_aug.narrow(1, 0, sym_len_hs).copy_(sym_patch_inv) + + sym_patch = img[:, -sym_len_he:, :] + inv_idx = torch.arange(sym_patch.size(1) - 1, -1, -1).long() + sym_patch_inv = sym_patch.index_select(1, inv_idx) + img_aug.narrow(1, sym_len_hs + in_h, sym_len_he).copy_(sym_patch_inv) + + out_1 = torch.FloatTensor(in_c, out_h, in_w) + kernel_width = weights_h.size(1) + for i in range(out_h): + idx = int(indices_h[i][0]) + for j in range(in_c): + out_1[j, i, :] = img_aug[j, idx:idx + kernel_width, :].transpose(0, 1).mv(weights_h[i]) + + # process W dimension + # symmetric copying + out_1_aug = torch.FloatTensor(in_c, out_h, in_w + sym_len_ws + sym_len_we) + out_1_aug.narrow(2, sym_len_ws, in_w).copy_(out_1) + + sym_patch = out_1[:, :, :sym_len_ws] + inv_idx = torch.arange(sym_patch.size(2) - 1, -1, -1).long() + sym_patch_inv = sym_patch.index_select(2, inv_idx) + out_1_aug.narrow(2, 0, sym_len_ws).copy_(sym_patch_inv) + + sym_patch = out_1[:, :, -sym_len_we:] + inv_idx = torch.arange(sym_patch.size(2) - 1, -1, -1).long() + sym_patch_inv = sym_patch.index_select(2, inv_idx) + out_1_aug.narrow(2, sym_len_ws + in_w, sym_len_we).copy_(sym_patch_inv) + + out_2 = torch.FloatTensor(in_c, out_h, out_w) + kernel_width = weights_w.size(1) + for i in range(out_w): + idx = int(indices_w[i][0]) + for j in range(in_c): + out_2[j, :, i] = out_1_aug[j, :, idx:idx + kernel_width].mv(weights_w[i]) + + if squeeze_flag: + out_2 = out_2.squeeze(0) + if numpy_type: + out_2 = out_2.numpy() + if not squeeze_flag: + out_2 = out_2.transpose(1, 2, 0) + + return out_2 diff --git a/official/bsr/utils/misc.py b/official/bsr/utils/misc.py new file mode 100644 index 0000000000000000000000000000000000000000..a43f878f1d7b61ece665c75a736f3859849b5b42 --- /dev/null +++ b/official/bsr/utils/misc.py @@ -0,0 +1,141 @@ +import numpy as np +import os +import random +import time +import torch +from os import path as osp + +from .dist_util import master_only + + +def set_random_seed(seed): + """Set random seeds.""" + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + + +def get_time_str(): + return time.strftime('%Y%m%d_%H%M%S', time.localtime()) + + +def mkdir_and_rename(path): + """mkdirs. If path exists, rename it with timestamp and create a new one. + + Args: + path (str): Folder path. + """ + if osp.exists(path): + new_name = path + '_archived_' + get_time_str() + print(f'Path already exists. Rename it to {new_name}', flush=True) + os.rename(path, new_name) + os.makedirs(path, exist_ok=True) + + +@master_only +def make_exp_dirs(opt): + """Make dirs for experiments.""" + path_opt = opt['path'].copy() + if opt['is_train']: + mkdir_and_rename(path_opt.pop('experiments_root')) + else: + mkdir_and_rename(path_opt.pop('results_root')) + for key, path in path_opt.items(): + if ('strict_load' in key) or ('pretrain_network' in key) or ('resume' in key) or ('param_key' in key): + continue + else: + os.makedirs(path, exist_ok=True) + + +def scandir(dir_path, suffix=None, recursive=False, full_path=False): + """Scan a directory to find the interested files. + + Args: + dir_path (str): Path of the directory. + suffix (str | tuple(str), optional): File suffix that we are + interested in. Default: None. + recursive (bool, optional): If set to True, recursively scan the + directory. Default: False. + full_path (bool, optional): If set to True, include the dir_path. + Default: False. + + Returns: + A generator for all the interested files with relative paths. + """ + + if (suffix is not None) and not isinstance(suffix, (str, tuple)): + raise TypeError('"suffix" must be a string or tuple of strings') + + root = dir_path + + def _scandir(dir_path, suffix, recursive): + for entry in os.scandir(dir_path): + if not entry.name.startswith('.') and entry.is_file(): + if full_path: + return_path = entry.path + else: + return_path = osp.relpath(entry.path, root) + + if suffix is None: + yield return_path + elif return_path.endswith(suffix): + yield return_path + else: + if recursive: + yield from _scandir(entry.path, suffix=suffix, recursive=recursive) + else: + continue + + return _scandir(dir_path, suffix=suffix, recursive=recursive) + + +def check_resume(opt, resume_iter): + """Check resume states and pretrain_network paths. + + Args: + opt (dict): Options. + resume_iter (int): Resume iteration. + """ + if opt['path']['resume_state']: + # get all the networks + networks = [key for key in opt.keys() if key.startswith('network_')] + flag_pretrain = False + for network in networks: + if opt['path'].get(f'pretrain_{network}') is not None: + flag_pretrain = True + if flag_pretrain: + print('pretrain_network path will be ignored during resuming.') + # set pretrained model paths + for network in networks: + name = f'pretrain_{network}' + basename = network.replace('network_', '') + if opt['path'].get('ignore_resume_networks') is None or (network + not in opt['path']['ignore_resume_networks']): + opt['path'][name] = osp.join(opt['path']['models'], f'net_{basename}_{resume_iter}.pth') + print(f"Set {name} to {opt['path'][name]}") + + # change param_key to params in resume + param_keys = [key for key in opt['path'].keys() if key.startswith('param_key')] + for param_key in param_keys: + if opt['path'][param_key] == 'params_ema': + opt['path'][param_key] = 'params' + print(f'Set {param_key} to params') + + +def sizeof_fmt(size, suffix='B'): + """Get human readable file size. + + Args: + size (int): File size. + suffix (str): Suffix. Default: 'B'. + + Return: + str: Formatted file size. + """ + for unit in ['', 'K', 'M', 'G', 'T', 'P', 'E', 'Z']: + if abs(size) < 1024.0: + return f'{size:3.1f} {unit}{suffix}' + size /= 1024.0 + return f'{size:3.1f} Y{suffix}' diff --git a/official/bsr/utils/options.py b/official/bsr/utils/options.py new file mode 100644 index 0000000000000000000000000000000000000000..d78f19bdf7dec470a6bf678ccef26bc171b2482d --- /dev/null +++ b/official/bsr/utils/options.py @@ -0,0 +1,218 @@ +import argparse +import os +import random +import torch +import yaml +from collections import OrderedDict +from os import path as osp + +from bsr.utils import set_random_seed +from bsr.utils.dist_util import get_dist_info, init_dist, master_only + + +def ordered_yaml(): + """Support OrderedDict for yaml. + + Returns: + tuple: yaml Loader and Dumper. + """ + try: + from yaml import CDumper as Dumper + from yaml import CLoader as Loader + except ImportError: + from yaml import Dumper, Loader + + _mapping_tag = yaml.resolver.BaseResolver.DEFAULT_MAPPING_TAG + + def dict_representer(dumper, data): + return dumper.represent_dict(data.items()) + + def dict_constructor(loader, node): + return OrderedDict(loader.construct_pairs(node)) + + Dumper.add_representer(OrderedDict, dict_representer) + Loader.add_constructor(_mapping_tag, dict_constructor) + return Loader, Dumper + + +def yaml_load(f): + """Load yaml file or string. + + Args: + f (str): File path or a python string. + + Returns: + dict: Loaded dict. + """ + if os.path.isfile(f): + with open(f, 'r') as f: + return yaml.load(f, Loader=ordered_yaml()[0]) + else: + return yaml.load(f, Loader=ordered_yaml()[0]) + + +def dict2str(opt, indent_level=1): + """dict to string for printing options. + + Args: + opt (dict): Option dict. + indent_level (int): Indent level. Default: 1. + + Return: + (str): Option string for printing. + """ + msg = '\n' + for k, v in opt.items(): + if isinstance(v, dict): + msg += ' ' * (indent_level * 2) + k + ':[' + msg += dict2str(v, indent_level + 1) + msg += ' ' * (indent_level * 2) + ']\n' + else: + msg += ' ' * (indent_level * 2) + k + ': ' + str(v) + '\n' + return msg + + +def _postprocess_yml_value(value): + # None + if value == '~' or value.lower() == 'none': + return None + # bool + if value.lower() == 'true': + return True + elif value.lower() == 'false': + return False + # !!float number + if value.startswith('!!float'): + return float(value.replace('!!float', '')) + # number + if value.isdigit(): + return int(value) + elif value.replace('.', '', 1).isdigit() and value.count('.') < 2: + return float(value) + # list + if value.startswith('['): + return eval(value) + # str + return value + + +def parse_options(root_path, is_train=True): + parser = argparse.ArgumentParser() + parser.add_argument('-opt', type=str, required=True, help='Path to option YAML file.') + parser.add_argument('--launcher', choices=['none', 'pytorch', 'slurm'], default='none', help='job launcher') + parser.add_argument('--auto_resume', action='store_true') + parser.add_argument('--debug', action='store_true') + parser.add_argument('--local_rank', type=int, default=0) + parser.add_argument( + '--force_yml', nargs='+', default=None, help='Force to update yml files. Examples: train:ema_decay=0.999') + args = parser.parse_args() + + # parse yml to dict + opt = yaml_load(args.opt) + + # distributed settings + if args.launcher == 'none': + opt['dist'] = False + print('Disable distributed.', flush=True) + else: + opt['dist'] = True + if args.launcher == 'slurm' and 'dist_params' in opt: + init_dist(args.launcher, **opt['dist_params']) + else: + init_dist(args.launcher) + opt['rank'], opt['world_size'] = get_dist_info() + + # random seed + seed = opt.get('manual_seed') + if seed is None: + seed = random.randint(1, 10000) + opt['manual_seed'] = seed + set_random_seed(seed + opt['rank']) + + # force to update yml options + if args.force_yml is not None: + for entry in args.force_yml: + # now do not support creating new keys + keys, value = entry.split('=') + keys, value = keys.strip(), value.strip() + value = _postprocess_yml_value(value) + eval_str = 'opt' + for key in keys.split(':'): + eval_str += f'["{key}"]' + eval_str += '=value' + # using exec function + exec(eval_str) + + opt['auto_resume'] = args.auto_resume + opt['is_train'] = is_train + + # debug setting + if args.debug and not opt['name'].startswith('debug'): + opt['name'] = 'debug_' + opt['name'] + + if opt['num_gpu'] == 'auto': + opt['num_gpu'] = torch.cuda.device_count() + + # datasets + for phase, dataset in opt['datasets'].items(): + # for multiple datasets, e.g., val_1, val_2; test_1, test_2 + phase = phase.split('_')[0] + dataset['phase'] = phase + if 'scale' in opt: + dataset['scale'] = opt['scale'] + if dataset.get('dataroot_gt') is not None: + dataset['dataroot_gt'] = osp.expanduser(dataset['dataroot_gt']) + if dataset.get('dataroot_lq') is not None: + dataset['dataroot_lq'] = osp.expanduser(dataset['dataroot_lq']) + + # paths + for key, val in opt['path'].items(): + if (val is not None) and ('resume_state' in key or 'pretrain_network' in key): + opt['path'][key] = osp.expanduser(val) + + if is_train: + experiments_root = opt['path'].get('experiments_root') + if experiments_root is None: + experiments_root = osp.join(root_path, 'experiments') + experiments_root = osp.join(experiments_root, opt['name']) + + opt['path']['experiments_root'] = experiments_root + opt['path']['models'] = osp.join(experiments_root, 'models') + opt['path']['training_states'] = osp.join(experiments_root, 'training_states') + opt['path']['log'] = experiments_root + opt['path']['visualization'] = osp.join(experiments_root, 'visualization') + + # change some options for debug mode + if 'debug' in opt['name']: + if 'val' in opt: + opt['val']['val_freq'] = 8 + opt['logger']['print_freq'] = 1 + opt['logger']['save_checkpoint_freq'] = 8 + else: # test + results_root = opt['path'].get('results_root') + if results_root is None: + results_root = osp.join(root_path, 'results') + results_root = osp.join(results_root, opt['name']) + + opt['path']['results_root'] = results_root + opt['path']['log'] = results_root + opt['path']['visualization'] = osp.join(results_root, 'visualization') + + return opt, args + + +@master_only +def copy_opt_file(opt_file, experiments_root): + # copy the yml file to the experiment root + import sys + import time + from shutil import copyfile + cmd = ' '.join(sys.argv) + filename = osp.join(experiments_root, osp.basename(opt_file)) + copyfile(opt_file, filename) + + with open(filename, 'r+') as f: + lines = f.readlines() + lines.insert(0, f'# GENERATE TIME: {time.asctime()}\n# CMD:\n# {cmd}\n\n') + f.seek(0) + f.writelines(lines) diff --git a/official/bsr/utils/plot_util.py b/official/bsr/utils/plot_util.py new file mode 100644 index 0000000000000000000000000000000000000000..7094a7f44780e3accbbe985228a0f6f5e0c6b454 --- /dev/null +++ b/official/bsr/utils/plot_util.py @@ -0,0 +1,83 @@ +import re + + +def read_data_from_tensorboard(log_path, tag): + """Get raw data (steps and values) from tensorboard events. + + Args: + log_path (str): Path to the tensorboard log. + tag (str): tag to be read. + """ + from tensorboard.backend.event_processing.event_accumulator import EventAccumulator + + # tensorboard event + event_acc = EventAccumulator(log_path) + event_acc.Reload() + scalar_list = event_acc.Tags()['scalars'] + print('tag list: ', scalar_list) + steps = [int(s.step) for s in event_acc.Scalars(tag)] + values = [s.value for s in event_acc.Scalars(tag)] + return steps, values + + +def read_data_from_txt_2v(path, pattern, step_one=False): + """Read data from txt with 2 returned values (usually [step, value]). + + Args: + path (str): path to the txt file. + pattern (str): re (regular expression) pattern. + step_one (bool): add 1 to steps. Default: False. + """ + with open(path) as f: + lines = f.readlines() + lines = [line.strip() for line in lines] + steps = [] + values = [] + + pattern = re.compile(pattern) + for line in lines: + match = pattern.match(line) + if match: + steps.append(int(match.group(1))) + values.append(float(match.group(2))) + if step_one: + steps = [v + 1 for v in steps] + return steps, values + + +def read_data_from_txt_1v(path, pattern): + """Read data from txt with 1 returned values. + + Args: + path (str): path to the txt file. + pattern (str): re (regular expression) pattern. + """ + with open(path) as f: + lines = f.readlines() + lines = [line.strip() for line in lines] + data = [] + + pattern = re.compile(pattern) + for line in lines: + match = pattern.match(line) + if match: + data.append(float(match.group(1))) + return data + + +def smooth_data(values, smooth_weight): + """ Smooth data using 1st-order IIR low-pass filter (what tensorflow does). + + Reference: https://github.com/tensorflow/tensorboard/blob/f801ebf1f9fbfe2baee1ddd65714d0bccc640fb1/tensorboard/plugins/scalar/vz_line_chart/vz-line-chart.ts#L704 # noqa: E501 + + Args: + values (list): A list of values to be smoothed. + smooth_weight (float): Smooth weight. + """ + values_sm = [] + last_sm_value = values[0] + for value in values: + value_sm = last_sm_value * smooth_weight + (1 - smooth_weight) * value + values_sm.append(value_sm) + last_sm_value = value_sm + return values_sm diff --git a/official/bsr/utils/registry.py b/official/bsr/utils/registry.py new file mode 100644 index 0000000000000000000000000000000000000000..1745e94f2865d8d6cc2a7b6dcd1fdf359232427a --- /dev/null +++ b/official/bsr/utils/registry.py @@ -0,0 +1,88 @@ +# Modified from: https://github.com/facebookresearch/fvcore/blob/master/fvcore/common/registry.py # noqa: E501 + + +class Registry(): + """ + The registry that provides name -> object mapping, to support third-party + users' custom modules. + + To create a registry (e.g. a backbone registry): + + .. code-block:: python + + BACKBONE_REGISTRY = Registry('BACKBONE') + + To register an object: + + .. code-block:: python + + @BACKBONE_REGISTRY.register() + class MyBackbone(): + ... + + Or: + + .. code-block:: python + + BACKBONE_REGISTRY.register(MyBackbone) + """ + + def __init__(self, name): + """ + Args: + name (str): the name of this registry + """ + self._name = name + self._obj_map = {} + + def _do_register(self, name, obj, suffix=None): + if isinstance(suffix, str): + name = name + '_' + suffix + + assert (name not in self._obj_map), (f"An object named '{name}' was already registered " + f"in '{self._name}' registry!") + self._obj_map[name] = obj + + def register(self, obj=None, suffix=None): + """ + Register the given object under the the name `obj.__name__`. + Can be used as either a decorator or not. + See docstring of this class for usage. + """ + if obj is None: + # used as a decorator + def deco(func_or_class): + name = func_or_class.__name__ + self._do_register(name, func_or_class, suffix) + return func_or_class + + return deco + + # used as a function call + name = obj.__name__ + self._do_register(name, obj, suffix) + + def get(self, name, suffix='basicsr'): + ret = self._obj_map.get(name) + if ret is None: + ret = self._obj_map.get(name + '_' + suffix) + print(f'Name {name} is not found, use name: {name}_{suffix}!') + if ret is None: + raise KeyError(f"No object named '{name}' found in '{self._name}' registry!") + return ret + + def __contains__(self, name): + return name in self._obj_map + + def __iter__(self): + return iter(self._obj_map.items()) + + def keys(self): + return self._obj_map.keys() + + +DATASET_REGISTRY = Registry('dataset') +ARCH_REGISTRY = Registry('arch') +MODEL_REGISTRY = Registry('model') +LOSS_REGISTRY = Registry('loss') +METRIC_REGISTRY = Registry('metric') diff --git a/official/dataset.py b/official/dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..b6107ce8d0675a63bc4e971105101d401fb6fe70 --- /dev/null +++ b/official/dataset.py @@ -0,0 +1,280 @@ +import torch, random, cv2, os, math, glob +import torch.nn.functional as F +import numpy as np +from bsr.degradations import circular_lowpass_kernel, random_mixed_kernels, random_add_gaussian_noise_pt, random_add_poisson_noise_pt +from bsr.transforms import augment, paired_random_crop +from bsr.utils import FileClient, imfrombytes, img2tensor, DiffJPEG +from bsr.utils.img_process_util import filter2D + +class RealESRGANDataset(torch.utils.data.Dataset): + def __init__(self, opt, bsz): + super(RealESRGANDataset, self).__init__() + self.opt = opt + self.file_client = FileClient("disk") + self.gt_folder = opt["dataroot_gt"] + self.len = bsz * opt["iter_num"] + self.paths = glob.glob(os.path.join(self.gt_folder, "**/*"), recursive=True) + + # blur settings for the first degradation + self.blur_kernel_size = opt["blur_kernel_size"] + self.kernel_list = opt["kernel_list"] + self.kernel_prob = opt["kernel_prob"] # a list for each kernel probability + self.blur_sigma = opt["blur_sigma"] + self.betag_range = opt["betag_range"] # betag used in generalized Gaussian blur kernels + self.betap_range = opt["betap_range"] # betap used in plateau blur kernels + self.sinc_prob = opt["sinc_prob"] # the probability for sinc filters + + # blur settings for the second degradation + self.blur_kernel_size2 = opt["blur_kernel_size2"] + self.kernel_list2 = opt["kernel_list2"] + self.kernel_prob2 = opt["kernel_prob2"] + self.blur_sigma2 = opt["blur_sigma2"] + self.betag_range2 = opt["betag_range2"] + self.betap_range2 = opt["betap_range2"] + self.sinc_prob2 = opt["sinc_prob2"] + + # a final sinc filter + self.final_sinc_prob = opt["final_sinc_prob"] + + self.kernel_range = [2 * v + 1 for v in range(3, 11)] # kernel size ranges from 7 to 21 + # TODO: kernel range is now hard-coded, should be in the configure file + self.pulse_tensor = torch.zeros(21, 21).float() # convolving with pulse tensor brings no blurry effect + self.pulse_tensor[10, 10] = 1 + + def __getitem__(self, index): + index = random.randint(0, len(self.paths) - 1) + gt_path = self.paths[index] + img_gt = imfrombytes(self.file_client.get(gt_path, "gt"), float32=True) + img_gt = augment(img_gt, self.opt["use_hflip"], self.opt["use_rot"]) + h, w = img_gt.shape[0:2] + crop_pad_size = self.opt.gt_size + if h < crop_pad_size or w < crop_pad_size: + pad_h = max(0, crop_pad_size - h) + pad_w = max(0, crop_pad_size - w) + img_gt = cv2.copyMakeBorder(img_gt, 0, pad_h, 0, pad_w, cv2.BORDER_REFLECT_101) + if img_gt.shape[0] > crop_pad_size or img_gt.shape[1] > crop_pad_size: + h, w = img_gt.shape[0:2] + top = random.randint(0, h - crop_pad_size) + left = random.randint(0, w - crop_pad_size) + img_gt = img_gt[top:top + crop_pad_size, left:left + crop_pad_size, ...] + + # ------------------------ Generate kernels (used in the first degradation) ------------------------ # + kernel_size = random.choice(self.kernel_range) + if np.random.uniform() < self.opt["sinc_prob"]: + # this sinc filter setting is for kernels ranging from [7, 21] + if kernel_size < 13: + omega_c = np.random.uniform(np.pi / 3, np.pi) + else: + omega_c = np.random.uniform(np.pi / 5, np.pi) + kernel = circular_lowpass_kernel(omega_c, kernel_size, pad_to=False) + else: + kernel = random_mixed_kernels( + self.kernel_list, + self.kernel_prob, + kernel_size, + self.blur_sigma, + self.blur_sigma, [-math.pi, math.pi], + self.betag_range, + self.betap_range, + noise_range=None) + # pad kernel + pad_size = (21 - kernel_size) // 2 + kernel = np.pad(kernel, ((pad_size, pad_size), (pad_size, pad_size))) + + # ------------------------ Generate kernels (used in the second degradation) ------------------------ # + kernel_size = random.choice(self.kernel_range) + if np.random.uniform() < self.opt["sinc_prob2"]: + if kernel_size < 13: + omega_c = np.random.uniform(np.pi / 3, np.pi) + else: + omega_c = np.random.uniform(np.pi / 5, np.pi) + kernel2 = circular_lowpass_kernel(omega_c, kernel_size, pad_to=False) + else: + kernel2 = random_mixed_kernels( + self.kernel_list2, + self.kernel_prob2, + kernel_size, + self.blur_sigma2, + self.blur_sigma2, [-math.pi, math.pi], + self.betag_range2, + self.betap_range2, + noise_range=None) + + # pad kernel + pad_size = (21 - kernel_size) // 2 + kernel2 = np.pad(kernel2, ((pad_size, pad_size), (pad_size, pad_size))) + + # ------------------------------------- the final sinc kernel ------------------------------------- # + if np.random.uniform() < self.opt["final_sinc_prob"]: + kernel_size = random.choice(self.kernel_range) + omega_c = np.random.uniform(np.pi / 3, np.pi) + sinc_kernel = circular_lowpass_kernel(omega_c, kernel_size, pad_to=21) + sinc_kernel = torch.FloatTensor(sinc_kernel) + else: + sinc_kernel = self.pulse_tensor + + # BGR to RGB, HWC to CHW, numpy to tensor + img_gt = img2tensor([img_gt], bgr2rgb=True, float32=True)[0] + kernel = torch.FloatTensor(kernel) + kernel2 = torch.FloatTensor(kernel2) + + return_d = {"gt": img_gt, "kernel1": kernel, "kernel2": kernel2, "sinc_kernel": sinc_kernel, "gt_path": gt_path} + return return_d + + def __len__(self): + return self.len + +class RealESRGANDegrader: + def __init__(self, opt, device): + self.opt = opt + self.device = device + self.jpeger = DiffJPEG(differentiable=False).to(device) # simulate JPEG compression artifacts + self.queue_size = 1200 + + @torch.no_grad() + def _dequeue_and_enqueue(self): + """It is the training pair pool for increasing the diversity in a batch. + + Batch processing limits the diversity of synthetic degradations in a batch. For example, samples in a + batch could not have different resize scaling factors. Therefore, we employ this training pair pool + to increase the degradation diversity in a batch. + """ + # initialize + b, c, h, w = self.lq.size() + if not hasattr(self, "queue_lr"): + assert self.queue_size % b == 0, f"queue size {self.queue_size} should be divisible by batch size {b}" + self.queue_lr = torch.zeros(self.queue_size, c, h, w).to(self.device) + _, c, h, w = self.gt.size() + self.queue_gt = torch.zeros(self.queue_size, c, h, w).to(self.device) + self.queue_ptr = 0 + if self.queue_ptr == self.queue_size: # the pool is full + # do dequeue and enqueue + # shuffle + idx = torch.randperm(self.queue_size) + self.queue_lr = self.queue_lr[idx] + self.queue_gt = self.queue_gt[idx] + # get first b samples + lq_dequeue = self.queue_lr[0:b, :, :, :].clone() + gt_dequeue = self.queue_gt[0:b, :, :, :].clone() + # update the queue + self.queue_lr[0:b, :, :, :] = self.lq.clone() + self.queue_gt[0:b, :, :, :] = self.gt.clone() + + self.lq = lq_dequeue + self.gt = gt_dequeue + else: + # only do enqueue + self.queue_lr[self.queue_ptr:self.queue_ptr + b, :, :, :] = self.lq.clone() + self.queue_gt[self.queue_ptr:self.queue_ptr + b, :, :, :] = self.gt.clone() + self.queue_ptr = self.queue_ptr + b + + @torch.no_grad() + def degrade(self, data): + """Accept data from dataloader, and then add two-order degradations to obtain LQ images. + """ + # training data synthesis + self.gt = data["gt"].to(self.device) + + self.kernel1 = data["kernel1"].to(self.device) + self.kernel2 = data["kernel2"].to(self.device) + self.sinc_kernel = data["sinc_kernel"].to(self.device) + + ori_h, ori_w = self.gt.size()[2:4] + + # ----------------------- The first degradation process ----------------------- # + # blur + out = filter2D(self.gt, self.kernel1) + # random resize + updown_type = random.choices(["up", "down", "keep"], self.opt["resize_prob"])[0] + if updown_type == "up": + scale = np.random.uniform(1, self.opt["resize_range"][1]) + elif updown_type == "down": + scale = np.random.uniform(self.opt["resize_range"][0], 1) + else: + scale = 1 + mode = random.choice(["area", "bilinear", "bicubic"]) + out = F.interpolate(out, scale_factor=scale, mode=mode) + # add noise + gray_noise_prob = self.opt["gray_noise_prob"] + if np.random.uniform() < self.opt["gaussian_noise_prob"]: + out = random_add_gaussian_noise_pt( + out, sigma_range=self.opt["noise_range"], clip=True, rounds=False, gray_prob=gray_noise_prob) + else: + out = random_add_poisson_noise_pt( + out, + scale_range=self.opt["poisson_scale_range"], + gray_prob=gray_noise_prob, + clip=True, + rounds=False) + # JPEG compression + jpeg_p = out.new_zeros(out.size(0)).uniform_(*self.opt["jpeg_range"]) + out = torch.clamp(out, 0, 1) # clamp to [0, 1], otherwise JPEGer will result in unpleasant artifacts + out = self.jpeger(out, quality=jpeg_p) + + # ----------------------- The second degradation process ----------------------- # + # blur + if np.random.uniform() < self.opt["second_blur_prob"]: + out = filter2D(out, self.kernel2) + # random resize + updown_type = random.choices(["up", "down", "keep"], self.opt["resize_prob2"])[0] + if updown_type == "up": + scale = np.random.uniform(1, self.opt["resize_range2"][1]) + elif updown_type == "down": + scale = np.random.uniform(self.opt["resize_range2"][0], 1) + else: + scale = 1 + mode = random.choice(["area", "bilinear", "bicubic"]) + out = F.interpolate( + out, size=(int(ori_h / self.opt["scale"] * scale), int(ori_w / self.opt["scale"] * scale)), mode=mode) + # add noise + gray_noise_prob = self.opt["gray_noise_prob2"] + if np.random.uniform() < self.opt["gaussian_noise_prob2"]: + out = random_add_gaussian_noise_pt( + out, sigma_range=self.opt["noise_range2"], clip=True, rounds=False, gray_prob=gray_noise_prob) + else: + out = random_add_poisson_noise_pt( + out, + scale_range=self.opt["poisson_scale_range2"], + gray_prob=gray_noise_prob, + clip=True, + rounds=False) + + # JPEG compression + the final sinc filter + # We also need to resize images to desired sizes. We group [resize back + sinc filter] together + # as one operation. + # We consider two orders: + # 1. [resize back + sinc filter] + JPEG compression + # 2. JPEG compression + [resize back + sinc filter] + # Empirically, we find other combinations (sinc + JPEG + Resize) will introduce twisted lines. + if np.random.uniform() < 0.5: + # resize back + the final sinc filter + mode = random.choice(["area", "bilinear", "bicubic"]) + out = F.interpolate(out, size=(ori_h // self.opt["scale"], ori_w // self.opt["scale"]), mode=mode) + out = filter2D(out, self.sinc_kernel) + # JPEG compression + jpeg_p = out.new_zeros(out.size(0)).uniform_(*self.opt["jpeg_range2"]) + out = torch.clamp(out, 0, 1) + out = self.jpeger(out, quality=jpeg_p) + else: + # JPEG compression + jpeg_p = out.new_zeros(out.size(0)).uniform_(*self.opt["jpeg_range2"]) + out = torch.clamp(out, 0, 1) + out = self.jpeger(out, quality=jpeg_p) + # resize back + the final sinc filter + mode = random.choice(["area", "bilinear", "bicubic"]) + out = F.interpolate(out, size=(ori_h // self.opt["scale"], ori_w // self.opt["scale"]), mode=mode) + out = filter2D(out, self.sinc_kernel) + + # clamp and round + self.lq = torch.clamp((out * 255.0).round(), 0, 255) / 255. + + # random crop + gt_size = self.opt["gt_size"] + self.gt, self.lq = paired_random_crop(self.gt, self.lq, gt_size, self.opt["scale"]) + + # training pair pool + self._dequeue_and_enqueue() + # sharpen self.gt again, as we have changed the self.gt with self._dequeue_and_enqueue + self.lq = self.lq.contiguous() # for the warning: grad and param do not obey the gradient layout contract + + return self.lq, self.gt diff --git a/official/evaluate.py b/official/evaluate.py new file mode 100644 index 0000000000000000000000000000000000000000..025f9a24023b5ed1710131a8585a01bb509b714a --- /dev/null +++ b/official/evaluate.py @@ -0,0 +1,55 @@ +import torch, os, glob, pyiqa +from argparse import ArgumentParser +import numpy as np +from PIL import Image +from tqdm import tqdm +from torchvision import transforms + +parser = ArgumentParser() +parser.add_argument("--HR_dir", type=str, default="testset/RealSR/HR") +parser.add_argument("--SR_dir", type=str, default="result/RealSR") +args = parser.parse_args() + +device = torch.device("cuda") + +psnr = pyiqa.create_metric("psnr", test_y_channel=True, color_space="ycbcr", device=device) +ssim = pyiqa.create_metric("ssim", test_y_channel=True, color_space="ycbcr", device=device) +lpips = pyiqa.create_metric("lpips", device=device) +dists = pyiqa.create_metric("dists", device=device) +fid = pyiqa.create_metric("fid", device=device) +niqe = pyiqa.create_metric("niqe", device=device) +maniqa = pyiqa.create_metric("maniqa-pipal", device=device) +clipiqa = pyiqa.create_metric("clipiqa", device=device) +musiq = pyiqa.create_metric("musiq", device=device) + +test_SR_paths = list(sorted(glob.glob(os.path.join(args.SR_dir, "*")))) +test_HR_paths = list(sorted(glob.glob(os.path.join(args.HR_dir, "*")))) + +metrics = {"psnr": [], "ssim": [], "lpips": [], "dists": [], "niqe": [], "maniqa": [], "musiq": [], "clipiqa": []} + +for i, (SR_path, HR_path) in tqdm(enumerate(zip(test_SR_paths, test_HR_paths))): + SR = Image.open(SR_path).convert("RGB") + SR = transforms.ToTensor()(SR).to(device).unsqueeze(0) + HR = Image.open(HR_path).convert("RGB") + HR = transforms.ToTensor()(HR).to(device).unsqueeze(0) + metrics["psnr"].append(psnr(SR, HR).item()) + metrics["ssim"].append(ssim(SR, HR).item()) + metrics["lpips"].append(lpips(SR, HR).item()) + metrics["dists"].append(dists(SR, HR).item()) + metrics["niqe"].append(niqe(SR).item()) + metrics["maniqa"].append(maniqa(SR).item()) + metrics["clipiqa"].append(clipiqa(SR).item()) + metrics["musiq"].append(musiq(SR).item()) + +for k in metrics.keys(): + metrics[k] = np.mean(metrics[k]) + +metrics["fid"] = fid(args.SR_dir, args.HR_dir) + +for k, v in metrics.items(): + if k == "niqe": + print(k, f"{v:.3g}") + elif k == "fid": + print(k, f"{v:.5g}") + else: + print(k, f"{v:.4g}") \ No newline at end of file diff --git a/official/forward.py b/official/forward.py new file mode 100644 index 0000000000000000000000000000000000000000..5c3ad96cd2b02bd65a82c4676156687b3536ce47 --- /dev/null +++ b/official/forward.py @@ -0,0 +1,67 @@ +import torch + +def MyUNet2DConditionModel_SD_forward(self, x): + global skip + x = self.conv_in(x) + skip = [x] + x = self.body(x) + return x + +def MyCrossAttnDownBlock2D_SD_forward(self, x): + for i in range(2): + x = self.resnets[i](x) + x = self.attentions[i](x) + skip.append(x) + if self.downsamplers is not None: + x = self.downsamplers[0](x) + skip.append(x) + return x + +def MyCrossAttnUpBlock2D_SD_forward(self, x): + for i in range(3): + x = self.resnets[i](torch.cat([x, skip.pop()], dim=1)) + x = self.attentions[i](x) + if self.upsamplers is not None: + x = self.upsamplers[0](x) + return x + +def MyDownBlock2D_SD_forward(self, x): + for i in range(2): + x = self.resnets[i](x) + skip.append(x) + return x + +def MyUNetMidBlock2DCrossAttn_SD_forward(self, x): + x = self.resnets[0](x) + x = self.attentions[0](x) + x = self.resnets[1](x) + return x + +def MyUpBlock2D_SD_forward(self, x): + for i in range(3): + x = self.resnets[i](torch.cat([x, skip.pop()], dim=1)) + x = self.upsamplers[0](x) + return x + +def MyResnetBlock2D_SD_forward(self, x_in): + x = self.norm1(x_in) + x = self.nonlinearity(x) + x = self.conv1(x) + x = self.norm2(x) + x = self.nonlinearity(x) + x = self.conv2(x) + if self.in_channels == self.out_channels: + return x + x_in + return x + self.conv_shortcut(x_in) + +def MyTransformer2DModel_SD_forward(self, x_in): + b, c, h, w = x_in.shape + x = self.norm(x_in) + x = x.permute(0, 2, 3, 1).reshape(b, h * w, c).contiguous() + x = self.proj_in(x) + for block in self.transformer_blocks: + x = x + block.attn1(block.norm1(x)) + x = x + block.ff(block.norm3(x)) + x = self.proj_out(x) + x = x.reshape(b, h, w, c).permute(0, 3, 1, 2).contiguous() + return x + x_in \ No newline at end of file diff --git a/official/model.py b/official/model.py new file mode 100644 index 0000000000000000000000000000000000000000..ae6773f6281738f73990905dc67f3955e5961eaa --- /dev/null +++ b/official/model.py @@ -0,0 +1,152 @@ +import torch, types, copy +from torch import nn +import torch.nn.functional as F +from diffusers.models.unets.unet_2d_blocks import CrossAttnDownBlock2D, \ + CrossAttnUpBlock2D, \ + DownBlock2D, \ + UpBlock2D, \ + UNetMidBlock2DCrossAttn +from diffusers.models.resnet import ResnetBlock2D +from diffusers.models.transformers.transformer_2d import Transformer2DModel +from diffusers.models.attention import BasicTransformerBlock +from diffusers.models.downsampling import Downsample2D +from diffusers.models.upsampling import Upsample2D +from forward import MyUNet2DConditionModel_SD_forward, \ + MyCrossAttnDownBlock2D_SD_forward, \ + MyDownBlock2D_SD_forward, \ + MyUNetMidBlock2DCrossAttn_SD_forward, \ + MyCrossAttnUpBlock2D_SD_forward, \ + MyUpBlock2D_SD_forward, \ + MyResnetBlock2D_SD_forward, \ + MyTransformer2DModel_SD_forward + +def find_parent(model, module_name): + components = module_name.split(".") + parent = model + for comp in components[:-1]: + parent = getattr(parent, comp) + return parent, components[-1] + +def halve_channels(model): + for name, module in model.named_modules(): + if hasattr(module, "pruned"): + continue + if isinstance(module, nn.Conv2d): + in_channels = int(module.in_channels * 0.75) + out_channels = int(module.out_channels * 0.75) + new_conv = nn.Conv2d(in_channels=in_channels, + out_channels=out_channels, + kernel_size=module.kernel_size, + stride=module.stride, + padding=module.padding, + dilation=module.dilation, + groups=module.groups, + bias=module.bias is not None) + with torch.no_grad(): + new_conv.weight.copy_(module.weight[:out_channels, :in_channels]) + if module.bias is not None: + new_conv.bias.copy_(module.bias[:out_channels]) + parent, last_name = find_parent(model, name) + setattr(parent, last_name, new_conv) + new_conv.pruned = True + elif isinstance(module, nn.Linear): + in_features = int(module.in_features * 0.75) + out_features = int(module.out_features * 0.75) + new_linear = nn.Linear(in_features=in_features, + out_features=out_features, + bias=module.bias is not None) + with torch.no_grad(): + new_linear.weight.copy_(module.weight[:out_features, :in_features]) + if module.bias is not None: + new_linear.bias.copy_(module.bias[:out_features]) + parent, last_name = find_parent(model, name) + setattr(parent, last_name, new_linear) + new_linear.pruned = True + elif isinstance(module, nn.GroupNorm): + num_channels = int(module.num_channels * 0.75) + for num_groups in [32, 24, 16, 12, 8, 6, 4, 2, 1]: + if num_channels % num_groups == 0: + break + new_gn = nn.GroupNorm(num_groups=num_groups, + num_channels=num_channels, + eps=module.eps, + affine=module.affine) + with torch.no_grad(): + new_gn.weight.copy_(module.weight[:num_channels]) + new_gn.bias.copy_(module.bias[:num_channels]) + parent, last_name = find_parent(model, name) + setattr(parent, last_name, new_gn) + new_gn.pruned = True + elif isinstance(module, nn.LayerNorm): + normalized_shape = int(module.normalized_shape[0] * 0.75) + new_ln = nn.LayerNorm(normalized_shape, + eps=module.eps, + elementwise_affine=module.elementwise_affine) + with torch.no_grad(): + new_ln.weight.copy_(module.weight[:normalized_shape]) + new_ln.bias.copy_(module.bias[:normalized_shape]) + parent, last_name = find_parent(model, name) + setattr(parent, last_name, new_ln) + new_ln.pruned = True + elif isinstance(module, Downsample2D) or isinstance(module, Upsample2D): + module.channels = int(module.channels * 0.75) + +class Net(nn.Module): + def __init__(self, unet, decoder): + super().__init__() + del unet.time_embedding + new_conv_in = nn.Conv2d(16, 320, 3, padding=1) + new_conv_in.weight.data = unet.conv_in.weight.data.repeat(1, 4, 1, 1) + new_conv_in.bias.data = unet.conv_in.bias.data + unet.conv_in = new_conv_in + new_conv_out = nn.Conv2d(320, 342, 3, padding=1) + new_conv_out.weight.data = unet.conv_out.weight.data.repeat(86, 1, 1, 1)[:342] + new_conv_out.bias.data = unet.conv_out.bias.data.repeat(86,)[:342] + unet.conv_out = new_conv_out + def ResnetBlock2D_remove_time_emb_proj(module): + if isinstance(module, ResnetBlock2D): + del module.time_emb_proj + unet.apply(ResnetBlock2D_remove_time_emb_proj) + def BasicTransformerBlock_remove_cross_attn(module): + if isinstance(module, BasicTransformerBlock): + del module.attn2, module.norm2 + unet.apply(BasicTransformerBlock_remove_cross_attn) + def set_inplace_to_true(module): + if isinstance(module, nn.Dropout) or isinstance(module, nn.SiLU): + module.inplace = True + unet.apply(set_inplace_to_true) + def replace_forward_methods(module): + if isinstance(module, CrossAttnDownBlock2D): + module.forward = types.MethodType(MyCrossAttnDownBlock2D_SD_forward, module) + elif isinstance(module, DownBlock2D): + module.forward = types.MethodType(MyDownBlock2D_SD_forward, module) + elif isinstance(module, UNetMidBlock2DCrossAttn): + module.forward = types.MethodType(MyUNetMidBlock2DCrossAttn_SD_forward, module) + elif isinstance(module, UpBlock2D): + module.forward = types.MethodType(MyUpBlock2D_SD_forward, module) + elif isinstance(module, CrossAttnUpBlock2D): + module.forward = types.MethodType(MyCrossAttnUpBlock2D_SD_forward, module) + elif isinstance(module, ResnetBlock2D): + module.forward = types.MethodType(MyResnetBlock2D_SD_forward, module) + elif isinstance(module, Transformer2DModel): + module.forward = types.MethodType(MyTransformer2DModel_SD_forward, module) + unet.apply(replace_forward_methods) + unet.forward = types.MethodType(MyUNet2DConditionModel_SD_forward, unet) + halve_channels(unet) + unet.body = nn.Sequential( + *unet.down_blocks, + unet.mid_block, + *unet.up_blocks, + unet.conv_norm_out, + unet.conv_act, + unet.conv_out, + ) + del decoder.conv_in, decoder.up_blocks, decoder.conv_norm_out, decoder.conv_act, decoder.conv_out + self.body = nn.Sequential( + nn.PixelUnshuffle(2), + unet, + decoder.mid_block, + ) + + def forward(self, x): + return self.body(x) \ No newline at end of file diff --git a/official/test.py b/official/test.py new file mode 100644 index 0000000000000000000000000000000000000000..054c35f4452091d05bce0626750c384dff1a3a3c --- /dev/null +++ b/official/test.py @@ -0,0 +1,70 @@ +import torch, os, glob, copy +import torch.nn.functional as F +import numpy as np +from PIL import Image +from argparse import ArgumentParser +from torchvision import transforms +from model import Net + +parser = ArgumentParser() +parser.add_argument("--epoch", type=int, default=200) +parser.add_argument("--model_dir", type=str, default="weight") +parser.add_argument("--LR_dir", type=str, default="testset/RealSR/LR") +parser.add_argument("--HR_dir", type=str, default="testset/RealSR/HR") +parser.add_argument("--SR_dir", type=str, default="result/RealSR") +args = parser.parse_args() + +device = torch.device("cuda") + +from diffusers import StableDiffusionPipeline +model_id = "stabilityai/stable-diffusion-2-1-base" +pipe = StableDiffusionPipeline.from_pretrained(model_id).to(device) + +vae = pipe.vae +tokenizer = pipe.tokenizer +unet = pipe.unet +noise_scheduler = pipe.scheduler +text_encoder = pipe.text_encoder + +from diffusers.models.autoencoders.vae import Decoder +ckpt_halfdecoder = torch.load("./weight/pretrained/halfDecoder.ckpt", weights_only=False) +decoder = Decoder(in_channels=4, + out_channels=3, + up_block_types=["UpDecoderBlock2D" for _ in range(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) +decoder_ckpt = {} +for k,v in ckpt_halfdecoder["state_dict"].items(): + if "decoder" in k: + new_k = k.replace("decoder.", "") + decoder_ckpt[new_k] = v +decoder.load_state_dict(decoder_ckpt, strict=True) + +model = torch.nn.DataParallel(Net(unet, copy.deepcopy(decoder))) +model.load_state_dict(torch.load("./%s/net_params_%d.pkl" % (args.model_dir, args.epoch), weights_only=False)) +model = torch.nn.Sequential( + model.module, + *decoder.up_blocks, + decoder.conv_norm_out, + decoder.conv_act, + decoder.conv_out, +).to(device) + +test_LR_paths = list(sorted(glob.glob(os.path.join(args.LR_dir, "*.png")))) +test_HR_paths = list(sorted(glob.glob(os.path.join(args.HR_dir, "*.png")))) + +os.makedirs(args.SR_dir, exist_ok=True) + +with torch.no_grad(): + for i, path in enumerate(test_LR_paths): + LR = Image.open(path).convert("RGB") + LR = transforms.ToTensor()(LR).to(device).unsqueeze(0) * 2 - 1 + SR = model(LR) + SR = (SR - SR.mean(dim=[2,3],keepdim=True)) / SR.std(dim=[2,3],keepdim=True) \ + * LR.std(dim=[2,3],keepdim=True) + LR.mean(dim=[2,3],keepdim=True) + SR = transforms.ToPILImage()((SR[0] / 2 + 0.5).clamp(0, 1).cpu()) + SR.save(os.path.join(args.SR_dir, os.path.basename(path))) diff --git a/official/train.py b/official/train.py new file mode 100644 index 0000000000000000000000000000000000000000..e39307510e1d17ed993b18fed64292625bfe8572 --- /dev/null +++ b/official/train.py @@ -0,0 +1,225 @@ +import torch, os, glob, random, copy +import torch.nn.functional as F +from torch.utils.data import DataLoader +import torch.distributed as dist +from torch.nn.parallel import DistributedDataParallel as DDP +import numpy as np +from argparse import ArgumentParser +from time import time +from tqdm import tqdm +from omegaconf import OmegaConf +from dataset import RealESRGANDataset, RealESRGANDegrader +from model import Net +from ram.models.ram_lora import ram +from torchvision import transforms +from utils import add_lora_to_unet + +dist.init_process_group(backend="nccl", init_method="env://") +rank = dist.get_rank() +world_size = dist.get_world_size() + +parser = ArgumentParser() +parser.add_argument("--epoch", type=int, default=200) +parser.add_argument("--batch_size", type=int, default=12) +parser.add_argument("--learning_rate", type=float, default=1e-4) +parser.add_argument("--model_dir", type=str, default="weight") +parser.add_argument("--log_dir", type=str, default="log") +parser.add_argument("--save_interval", type=int, default=10) + +args = parser.parse_args() + +# fixed seed for reproduction +seed = rank +random.seed(seed) +np.random.seed(seed) +torch.manual_seed(seed) +torch.cuda.manual_seed_all(seed) + +config = OmegaConf.load("config.yml") + +epoch = args.epoch +learning_rate = args.learning_rate +bsz = args.batch_size + +device = torch.device(f"cuda:{rank}" if torch.cuda.is_available() else "cpu") +torch.backends.cudnn.allow_tf32 = True +torch.backends.cuda.matmul.allow_tf32 = True + +if rank == 0: + print("batch size per gpu =", bsz) + +from diffusers import StableDiffusionPipeline +model_id = "stabilityai/stable-diffusion-2-1-base" +pipe = StableDiffusionPipeline.from_pretrained(model_id).to(device) + +vae = pipe.vae +tokenizer = pipe.tokenizer +unet = pipe.unet +text_encoder = pipe.text_encoder + +unet_D = copy.deepcopy(unet) +new_conv_in = torch.nn.Conv2d(256, 320, 3, padding=1).to(device) +new_conv_in.weight.data = unet_D.conv_in.weight.data.repeat(1, 64, 1, 1) / 64 +new_conv_in.bias.data = unet_D.conv_in.bias.data +unet_D.conv_in = new_conv_in +unet_D = add_lora_to_unet(unet_D) +unet_D.set_adapters(["default_encoder", "default_decoder", "default_others"]) + +vae_teacher = copy.deepcopy(vae) +unet_teacher = copy.deepcopy(unet) + +osediff = torch.load("./weight/pretrained/osediff.pkl", weights_only=False) +vae_teacher.load_state_dict(osediff["vae"]) +unet_teacher.load_state_dict(osediff["unet"]) + +from diffusers.models.autoencoders.vae import Decoder +ckpt_halfdecoder = torch.load("./weight/pretrained/halfDecoder.ckpt", weights_only=False) +decoder = Decoder(in_channels=4, + out_channels=3, + up_block_types=["UpDecoderBlock2D" for _ in range(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) +decoder_ckpt = {} +for k, v in ckpt_halfdecoder["state_dict"].items(): + if "decoder" in k: + new_k = k.replace("decoder.", "") + decoder_ckpt[new_k] = v +decoder.load_state_dict(decoder_ckpt, strict=True) + +ram_transforms = transforms.Compose([ + transforms.Resize((384, 384)), + transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), +]) + +DAPE = ram(pretrained="./weight/pretrained/ram_swin_large_14m.pth", + pretrained_condition="./weight/pretrained/DAPE.pth", + image_size=384, + vit="swin_l").eval().to(device) + +vae.requires_grad_(False) +unet.requires_grad_(False) +text_encoder.requires_grad_(False) +vae_teacher.requires_grad_(False) +unet_teacher.requires_grad_(False) +decoder.requires_grad_(False) +DAPE.requires_grad_(False) + +model = DDP(Net(unet, copy.deepcopy(decoder)).to(device), device_ids=[rank]) +model_D = DDP(unet_D.to(device), device_ids=[rank]) +model.requires_grad_(True) +model_D.requires_grad_(False) +params_to_opt = [] +for n, p in model_D.named_parameters(): + if "lora" in n or "conv_in" in n: + p.requires_grad = True + params_to_opt.append(p) + +if rank == 0: + param_cnt = sum(p.numel() for p in model.parameters() if p.requires_grad) + print("#Param.", param_cnt/1e6, "M") + +dataset = RealESRGANDataset(config, bsz) +degrader = RealESRGANDegrader(config, device) +dataloader = DataLoader(dataset, batch_size=bsz, num_workers=8) +optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate) +optimizer_D = torch.optim.Adam(params_to_opt, lr=1e-6) +scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, milestones=[100,], gamma=0.5) +scaler = torch.cuda.amp.GradScaler() + +model_dir = "./%s" % (args.model_dir,) +log_path = "./%s/log.txt" % (args.log_dir,) +os.makedirs(model_dir, exist_ok=True) +os.makedirs(args.log_dir, exist_ok=True) + +print("start training...") +timesteps = torch.tensor([999], device=device).long().expand(bsz,) +alpha = pipe.scheduler.alphas_cumprod[999] +for epoch_i in range(1, epoch + 1): + start_time = time() + loss_avg = 0.0 + loss_distil_avg = 0.0 + loss_adv_avg = 0.0 + loss_D_avg = 0.0 + iter_num = 0 + dist.barrier() + for batch in tqdm(dataloader): + with torch.cuda.amp.autocast(enabled=True): + with torch.no_grad(): + LR, HR = degrader.degrade(batch) + text_input = tokenizer(DAPE.generate_tag(ram_transforms(LR))[0], + max_length=tokenizer.model_max_length, + padding="max_length", return_tensors="pt").to(device) + encoder_hidden_states = text_encoder(text_input.input_ids, return_dict=False)[0] + LR, HR = LR * 2 - 1, HR * 2 - 1 + LR_ = F.interpolate(LR, scale_factor=4, mode="bicubic") + LR_latents = vae_teacher.encode(LR_).latent_dist.mean * vae_teacher.config.scaling_factor + HR_latents = vae.encode(HR).latent_dist.mean + pred_teacher = unet_teacher( + LR_latents, + timesteps, + encoder_hidden_states=encoder_hidden_states, + return_dict=False, + )[0] + z0_teacher = (LR_latents-((1-alpha)**0.5)*pred_teacher)/(alpha**0.5) + z0_teacher = vae_teacher.post_quant_conv(z0_teacher / vae_teacher.config.scaling_factor) + z0_teacher = decoder.conv_in(z0_teacher) + z0_teacher = decoder.mid_block(z0_teacher) + z0_gt = vae.post_quant_conv(HR_latents) + z0_gt = decoder.conv_in(z0_gt) + z0_gt = decoder.mid_block(z0_gt) + z0_student = model(LR) + loss_distil = (z0_student - z0_teacher).abs().mean() + loss_adv = F.softplus(-model_D( + z0_student, + timesteps, + encoder_hidden_states=encoder_hidden_states, + return_dict=False, + )[0]).mean() + loss = loss_distil + loss_adv + optimizer.zero_grad(set_to_none=True) + scaler.scale(loss).backward() + scaler.step(optimizer) + scaler.update() + with torch.cuda.amp.autocast(enabled=True): + pred_real = model_D( + z0_gt.detach(), + timesteps, + encoder_hidden_states=encoder_hidden_states, + return_dict=False, + )[0] + pred_fake = model_D( + z0_student.detach(), + timesteps, + encoder_hidden_states=encoder_hidden_states, + return_dict=False, + )[0] + loss_D = F.softplus(pred_fake).mean() + F.softplus(-pred_real).mean() + optimizer_D.zero_grad(set_to_none=True) + scaler.scale(loss_D).backward() + scaler.step(optimizer_D) + scaler.update() + loss_avg += loss.item() + loss_distil_avg += loss_distil.item() + loss_adv_avg += loss_adv.item() + loss_D_avg += loss_D.item() + iter_num += 1 + # print("loss", loss.item()) + # print("loss_distil", loss_distil.item()) + # print("loss_adv", loss_adv.item()) + # print("loss_D", loss_D.item()) + scheduler.step() + loss_avg /= iter_num + loss_distil_avg /= iter_num + loss_adv_avg /= iter_num + loss_D_avg /= iter_num + log_data = "[%d/%d] Average loss: %f, distil loss: %f, adv loss: %f, D loss: %f, time cost: %.2fs, cur lr is %f." % (epoch_i, epoch, loss_avg, loss_distil_avg, loss_adv_avg, loss_D_avg, time() - start_time, scheduler.get_last_lr()[0]) + if rank == 0: + print(log_data) + with open(log_path, "a") as log_file: + log_file.write(log_data + "\n") + if epoch_i % args.save_interval == 0: + torch.save(model.state_dict(), "./%s/net_params_%d.pkl" % (model_dir, epoch_i)) diff --git a/official/utils.py b/official/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..7d4c2f25b3bbf1d598163631e97cad9bfac89008 --- /dev/null +++ b/official/utils.py @@ -0,0 +1,24 @@ +import torch +from peft import LoraConfig + +def add_lora_to_unet(unet, rank=4): + l_target_modules_encoder, l_target_modules_decoder, l_modules_others = [], [], [] + l_grep = ["to_k", "to_q", "to_v", "to_out.0", "conv", "conv1", "conv2", "conv_shortcut", "conv_out", "proj_out", "proj_in", "ff.net.2", "ff.net.0.proj"] + for n, p in unet.named_parameters(): + check_flag = 0 + if "bias" in n or "norm" in n: + continue + for pattern in l_grep: + if pattern in n and ("down_blocks" in n or "conv_in" in n): + l_target_modules_encoder.append(n.replace(".weight","")) + break + elif pattern in n and ("up_blocks" in n or "conv_out" in n): + l_target_modules_decoder.append(n.replace(".weight","")) + break + elif pattern in n: + l_modules_others.append(n.replace(".weight","")) + break + unet.add_adapter(LoraConfig(r=rank,init_lora_weights="gaussian",target_modules=l_target_modules_encoder), adapter_name="default_encoder") + unet.add_adapter(LoraConfig(r=rank,init_lora_weights="gaussian",target_modules=l_target_modules_decoder), adapter_name="default_decoder") + unet.add_adapter(LoraConfig(r=rank,init_lora_weights="gaussian",target_modules=l_modules_others), adapter_name="default_others") + return unet \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..f276501089b5ef8b50efa97a390d8e64c39d1991 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,17 @@ +# A100 训练环境依赖 (Python 3.10/3.11, CUDA 12.x) +torch>=2.1 +torchvision>=0.16 +diffusers>=0.27 +accelerate>=0.28 +omegaconf>=2.3 +einops +lpips +pyiqa +opencv-python +scipy +numpy +pillow +safetensors +huggingface_hub +pyyaml +requests diff --git a/scripts/archive_run.py b/scripts/archive_run.py new file mode 100644 index 0000000000000000000000000000000000000000..54fb30312b19e61e4c34433c6ccab6bea9d6cff3 --- /dev/null +++ b/scripts/archive_run.py @@ -0,0 +1,31 @@ +#!/usr/bin/env python +"""归档: 把 logs/ 与权重副本整理到 archive/_/, 生成 summary.txt +用法: python scripts/archive_run.py --name s2_best --src_log logs/s2 --weights weight/s2/net_params_*.pkl +""" +import argparse, os, shutil, json, glob, datetime + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--name", required=True) + ap.add_argument("--src_log", default="logs") + ap.add_argument("--weights", nargs="*", default=[]) + ap.add_argument("--eval_json", default="", help="eval_val.py 输出") + args = ap.parse_args() + dst = os.path.join("archive", f"{datetime.date.today().isoformat()}_{args.name}") + os.makedirs(dst, exist_ok=True) + if os.path.isdir(args.src_log): + shutil.copytree(args.src_log, os.path.join(dst, "log"), dirs_exist_ok=True) + wd = os.path.join(dst, "weights"); os.makedirs(wd, exist_ok=True) + for pat in args.weights: + for f in glob.glob(pat): + shutil.copy2(f, os.path.join(wd, os.path.basename(f))) + summary = {"name": args.name, "date": datetime.date.today().isoformat(), "weights": [os.path.basename(f) for pat in args.weights for f in glob.glob(pat)]} + if args.eval_json and os.path.exists(args.eval_json): + with open(args.eval_json, encoding="utf-8") as fh: + summary["eval_mean"] = json.load(fh).get("mean") + with open(os.path.join(dst, "summary.txt"), "w", encoding="utf-8") as fh: + fh.write(json.dumps(summary, ensure_ascii=False, indent=1)) + print("archived to", dst) + +if __name__ == "__main__": + main() diff --git a/scripts/download_weights.sh b/scripts/download_weights.sh new file mode 100644 index 0000000000000000000000000000000000000000..a07d7e6fb1314a58ae7e136c23d21c2d524c833f --- /dev/null +++ b/scripts/download_weights.sh @@ -0,0 +1,32 @@ +#!/usr/bin/env bash +# Plan A: 服务器有外网时下载权重(在项目根执行) +# 双方案: 无外网时在本地跑同脚本并把 weight/ 与 models/ 传上来(见 pack_upload) +set -euo pipefail +cd "$(dirname "$0")/.." +export HF_ENDPOINT=https://hf-mirror.com +mkdir -p weight/pretrained models/stable-diffusion-2-1-base weight/gdpo +HF=https://huggingface.co/Guaishou74851/AdcSR/resolve/main +for f in weight/net_params_200.pkl weight/pretrained/DAPE.pth weight/pretrained/halfDecoder.ckpt \ + weight/pretrained/osediff.pkl weight/pretrained/ram_swin_large_14m.pth; do + echo "==> $f"; curl -fL --retry 5 --retry-delay 3 -o "$f" "$HF/$f" +done +MS=https://modelscope.cn/models/AI-ModelScope/stable-diffusion-2-1-base/resolve/master +for f in model_index.json feature_extractor/preprocessor_config.json scheduler/scheduler_config.json \ + tokenizer/merges.txt tokenizer/special_tokens_map.json tokenizer/tokenizer_config.json \ + tokenizer/vocab.json text_encoder/config.json text_encoder/model.fp16.safetensors \ + unet/config.json unet/diffusion_pytorch_model.fp16.safetensors \ + vae/config.json vae/diffusion_pytorch_model.fp16.safetensors; do + echo "==> $f"; curl -fL --retry 5 --retry-delay 3 --create-dirs -o "models/stable-diffusion-2-1-base/$f" "$MS/$f" +done +# GDPO 教师 (S2 前下载; 失败可 --teacher osediff) +echo "==> GDPO (Joypop/GDPO)" +python - <<'EOF' +import os +try: + from huggingface_hub import snapshot_download + snapshot_download(repo_id="Joypop/GDPO", repo_type="model", local_dir="weight/gdpo", allow_patterns=["ckp/*", "ram/*"]) + print("GDPO downloaded") +except Exception as e: + print("GDPO download failed (可回退 osediff):", e) +EOF +du -sh weight models diff --git a/scripts/env_check.py b/scripts/env_check.py new file mode 100644 index 0000000000000000000000000000000000000000..77f58c585e5cafc961d35509fbfee6f3e27838c2 --- /dev/null +++ b/scripts/env_check.py @@ -0,0 +1,25 @@ +#!/usr/bin/env python +"""环境自检: GPU/CUDA/pyiqa/dwt/官方import 冒烟(不加载大权重)""" +import os, sys +sys.path.insert(0, "official"); sys.path.insert(0, "src") +import torch +print("torch", torch.__version__, "cuda", torch.cuda.is_available()) +if torch.cuda.is_available(): + print("gpu", torch.cuda.get_device_name(0), "mem_GB", round(torch.cuda.get_device_properties(0).total_memory/1e9,1)) +try: + import pyiqa; print("pyiqa ok") +except Exception as e: print("pyiqa FAIL", e) +try: + import pywt; print("pywt ok") +except Exception as e: print("pywt missing(可选)", e) +try: + import diffusers, transformers, peft, omegaconf + print("diffusers/transformers/peft ok") +except Exception as e: print("libs FAIL", e) +# 官方模块可 import +try: + import model, dataset, utils, forward + print("official imports ok") +except Exception as e: + print("official imports FAIL", e) +print("env_check done") diff --git a/scripts/probe_gdpo.py b/scripts/probe_gdpo.py new file mode 100644 index 0000000000000000000000000000000000000000..392d5d775c29e365e2ec76afc9a9ef25675850c5 --- /dev/null +++ b/scripts/probe_gdpo.py @@ -0,0 +1,13 @@ +#!/usr/bin/env python +"""GDPO 权重 probe: 打印键结构, 供 src/common.load_gdpo_teacher 排错。 +用法: python scripts/probe_gdpo.py weight/gdpo/ckp/diffusion_pytorch_model.safetensors +""" +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) +from common import probe_gdpo + +if __name__ == "__main__": + if len(sys.argv) < 2: + print("用法: probe_gdpo.py "); sys.exit(1) + probe_gdpo(sys.argv[1]) diff --git a/scripts/run_smoke.sh b/scripts/run_smoke.sh new file mode 100644 index 0000000000000000000000000000000000000000..e7bf22a7f12b13f292b403299e04971288ded73c --- /dev/null +++ b/scripts/run_smoke.sh @@ -0,0 +1,12 @@ +#!/usr/bin/env bash +# S0 冒烟(每个训练配置 1 迭代) +set -euo pipefail +cd "$(dirname "$0")/.." +source "$(conda info --base)/etc/profile.d/conda.sh"; conda activate AdcSR +CUDA_VISIBLE_DEVICES=0 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 || true +CUDA_VISIBLE_DEVICES=0 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 || true +echo smoke done diff --git a/scripts/run_stage1.sh b/scripts/run_stage1.sh new file mode 100644 index 0000000000000000000000000000000000000000..c1bff9e1c8e36def8b8c64a19d4a883557a76566 --- /dev/null +++ b/scripts/run_stage1.sh @@ -0,0 +1,14 @@ +#!/usr/bin/env bash +# S1 LoRA 像素域适配 (A100 40G) +set -euo pipefail +cd "$(dirname "$0")/.." +source "$(conda info --base)/etc/profile.d/conda.sh"; conda activate AdcSR +CUDA_VISIBLE_DEVICES=0 python src/train_lora.py \ + --config configs/config_s1_lora.yml \ + --manifest data/manifest_train.json --real_prob 0.35 \ + --model_id models/stable-diffusion-2-1-base \ + --half_decoder weight/pretrained/halfDecoder.ckpt \ + --init_net weight/net_params_200.pkl \ + --out weight/s1 --log_dir logs/s1 \ + --steps 20000 --batch_size 8 --grad_accum 2 --lr 5e-5 --lora_rank 64 \ + --save_every 2000 --w_l1 1.0 --w_lpips 1.0 --w_dists 0.3 --w_wave 0.5 --w_color 0.2 diff --git a/scripts/run_stage2.sh b/scripts/run_stage2.sh new file mode 100644 index 0000000000000000000000000000000000000000..5eaaab19f02675ae77cd65fee653fa9341b47651 --- /dev/null +++ b/scripts/run_stage2.sh @@ -0,0 +1,18 @@ +#!/usr/bin/env bash +# S2 GDPO 教师蒸馏(先 probe; GDPO 不可用回退 osediff) +set -euo pipefail +cd "$(dirname "$0")/.." +source "$(conda info --base)/etc/profile.d/conda.sh"; conda activate AdcSR +TEACHER="${1:-gdpo}" +if [ "$TEACHER" = "gdpo" ] && [ ! -d weight/gdpo ]; then + echo "未找到 weight/gdpo, 回退 osediff"; TEACHER=osediff +fi +CUDA_VISIBLE_DEVICES=0 python src/train_stage2.py \ + --config configs/config_s2_distill.yml \ + --model_id models/stable-diffusion-2-1-base \ + --half_decoder weight/pretrained/halfDecoder.ckpt \ + --teacher "$TEACHER" --gdpo_dir weight/gdpo \ + --init_net weight/s1/net_params_BEST.pkl \ + --out weight/s2 --log_dir logs/s2 \ + --steps 30000 --batch_size 8 --grad_accum 2 --lr 5e-5 --lr_D 1e-6 \ + --save_every 2000 --w_wave 0.5 --w_lpips 0.5 --w_dists 0.3 --w_color 0.2 diff --git a/scripts/run_stage3.sh b/scripts/run_stage3.sh new file mode 100644 index 0000000000000000000000000000000000000000..7b865a1e29ede4f00a294946887a6d92a95aa348 --- /dev/null +++ b/scripts/run_stage3.sh @@ -0,0 +1,14 @@ +#!/usr/bin/env bash +# S3 域微调(S2 最优权重, 小 LR) +set -euo pipefail +cd "$(dirname "$0")/.." +source "$(conda info --base)/etc/profile.d/conda.sh"; conda activate AdcSR +CUDA_VISIBLE_DEVICES=0 python src/train_stage2.py \ + --config configs/config_s3_ft.yml \ + --model_id models/stable-diffusion-2-1-base \ + --half_decoder weight/pretrained/halfDecoder.ckpt \ + --teacher gdpo --gdpo_dir weight/gdpo \ + --init_net weight/s2/net_params_BEST.pkl \ + --out weight/s3 --log_dir logs/s3 \ + --steps 20000 --batch_size 8 --grad_accum 2 --lr 2e-5 --lr_D 5e-7 \ + --save_every 2000 --w_wave 0.5 --w_lpips 0.5 --w_dists 0.3 --w_color 0.2 diff --git a/scripts/setup_env.sh b/scripts/setup_env.sh new file mode 100644 index 0000000000000000000000000000000000000000..f865c7c4f6b5cdb90a0e514e1463656e9a38cf9f --- /dev/null +++ b/scripts/setup_env.sh @@ -0,0 +1,21 @@ +#!/usr/bin/env bash +# aspire2a 环境准备(A100 40G 单卡; python 3.10) +set -euo pipefail +cd "$(dirname "$0")/.." +if ! command -v conda >/dev/null 2>&1; then echo "需要 conda"; exit 1; fi +conda create -n AdcSR python=3.10 -y || true +source "$(conda info --base)/etc/profile.d/conda.sh" +conda activate AdcSR +pip install --upgrade pip +pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121 +pip install -r official/requirements.txt +pip install diffusers transformers accelerate peft omegaconf tqdm pyyaml pyiqa lpips scikit-image imageio pillow numpy +pip install huggingface_hub hf_transfer safetensors gdown +pip install "dwt-pytorch" 2>/dev/null || pip install pywavelets || true +python - <<'EOF' +try: + import pywt; print("pywt ok") +except Exception as e: + print("pywt missing:", e) +EOF +echo "setup done" diff --git a/src/check_submission.py b/src/check_submission.py new file mode 100644 index 0000000000000000000000000000000000000000..b39f006bd88839fb956d63aad89545a4c273beda --- /dev/null +++ b/src/check_submission.py @@ -0,0 +1,48 @@ +#!/usr/bin/env python +"""提交包检查: 命名/数量/大小/结构/可加载/确定性。 +用法: python src/check_submission.py --zip xxx.zip [--names data/test_names.txt] [--test_runner] +""" +import argparse, json, os, sys, tempfile, zipfile +import torch + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--zip", required=True) + ap.add_argument("--names", default="", help="每行一个期望文件名") + ap.add_argument("--test_runner", action="store_true", help="解压并实际加载 jit 跑 1 次") + args = ap.parse_args() + size = os.path.getsize(args.zip) + checks = {"exists": True, "size_gb": round(size / 1e9, 2), "le_10gb": size <= 10e9} + with zipfile.ZipFile(args.zip) as z: + names = z.namelist() + outs = [n for n in names if "/output_dir/" in n or n.startswith("output_dir/")] + jpgs = [n for n in outs if n.lower().endswith(".jpg")] + mods = [n for n in names if "/model_dir/" in n or n.startswith("model_dir/")] + checks["output_jpg_count"] = len(jpgs) + checks["has_model"] = any(n.endswith(".pt") for n in mods) + checks["has_runner"] = any(n.endswith("runner.py") for n in mods) + if args.names: + expect = [l.strip() for l in open(args.names, encoding="utf-8") if l.strip()] + got = {os.path.basename(n) for n in jpgs} + checks["missing"] = [e for e in expect if e not in got] + checks["extra"] = sorted(got - set(expect))[:10] + if args.test_runner: + tmp = tempfile.mkdtemp() + z.extractall(tmp) + pt = [os.path.join(tmp, n) for n in mods if n.endswith(".pt")][0] + m = torch.jit.load(pt, map_location="cpu") + m.eval() + x = torch.randn(1, 3, 512, 512).half() * 0.5 + with torch.no_grad(): + o1 = m(x); o2 = m(x) + checks["jit_load_ok"] = True + checks["out_shape"] = list(o1.shape) + checks["deterministic_maxdiff"] = float((o1 - o2).abs().max()) + for k, v in checks.items(): + print(f"{k}: {v}") + bad = [k for k, v in checks.items() if v is False] or \ + (checks.get("missing") if checks.get("missing") else []) + print("PASS" if not bad else f"FAIL items: {bad}") + +if __name__ == "__main__": + main() diff --git a/src/common.py b/src/common.py new file mode 100644 index 0000000000000000000000000000000000000000..3ffbbdff66662f3fa0448b55aa5ba8b70ca9b92d --- /dev/null +++ b/src/common.py @@ -0,0 +1,213 @@ +#!/usr/bin/env python +"""AdcSR 工程公共工具:路径、模型装配、手工 LoRA 注入、教师加载、GDPO probe。 +训练脚本统一从这里 import,禁止各自重复实现装配逻辑。 +""" +import os, sys, copy, json, types +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"): + from diffusers import StableDiffusionPipeline + pipe = StableDiffusionPipeline.from_pretrained(model_id, torch_dtype=dtype).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 ?????" + "??? --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 + cin, cout = conv.in_channels, conv.out_channels + 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.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, cin, 1, 1)) + z = F.conv2d(z, self.lora_b.t().view(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): + setattr(parent, attr, LoRAConv2d(m, rank, alpha)) + elif isinstance(m, nn.Linear): + setattr(parent, attr, LoRALinear(m, rank, alpha)) + 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 + +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) diff --git a/src/eval_val.py b/src/eval_val.py new file mode 100644 index 0000000000000000000000000000000000000000..1da644ca4ecf4ef6f291c5fe4a5c34036eaeb3f0 --- /dev/null +++ b/src/eval_val.py @@ -0,0 +1,104 @@ +#!/usr/bin/env python +"""Val 评测 + 代理分: 在隔离 val(真实对/官方对) 上计算 8 指标 + proxy + 保真护栏。 +用法: + python src/eval_val.py --weights weight/s2/net_params_X.pkl --pairs_json data/manifest_val.json + python src/eval_val.py --weights ... --official_lq --official_gt +""" +import argparse, json, os, sys +from pathlib import Path +REPO = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(REPO)); sys.path.insert(0, str(REPO / "src")); sys.path.insert(0, str(REPO / "official")) +import numpy as np, torch +from PIL import Image +from inference_4k import build_model, infer_4k_single + +def load_metrics(device): + import pyiqa + return {k: pyiqa.create_metric(k, device=device) for k in + ["psnr", "ssim", "lpips", "dists", "niqe", "maniqa", "musiq", "clipiqa"]} + +def resize_pair(sr, hr, max_side=1024): + sr = sr.convert("RGB"); hr = hr.convert("RGB") + for im in (sr, hr): + w, h = im.size + if max(w, h) > max_side: + s = max_side / max(w, h) + im.thumbnail((int(w * s), int(h * s)), Image.LANCZOS) + return sr, hr + +def metric_dict(metrics, sr, hr, device): + import torchvision.transforms as T + def t(im): + return T.ToTensor()(im).to(device).unsqueeze(0) + a, b = t(sr), t(hr) + out = {"psnr": metrics["psnr"](a, b).item(), "ssim": metrics["ssim"](a, b).item(), + "lpips": metrics["lpips"](a, b).item(), "dists": metrics["dists"](a, b).item(), + "niqe": metrics["niqe"](a).item(), "maniqa": metrics["maniqa"](a).item(), + "musiq": metrics["musiq"](a).item(), "clipiqa": metrics["clipiqa"](a).item()} + return out + +def proxy(m): + return (0.25 * m["clipiqa"] + 0.2 * (m["musiq"] / 100.0) + 0.15 * m["maniqa"] + + 0.15 * (1 - m["niqe"] / 10.0) + 0.25 * (1 - m["lpips"])) + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--weights", required=True) + ap.add_argument("--half_decoder", default="weight/pretrained/halfDecoder.ckpt") + ap.add_argument("--model_id", default="models/stable-diffusion-2-1-base") + ap.add_argument("--pairs_json", default="", help='[{lr,hr}] 或 manifest 结构') + ap.add_argument("--official_lq", default="", help="同分辨率官方对(整图滑窗)") + ap.add_argument("--official_gt", default="") + ap.add_argument("--out", default="logs/eval_result.json") + ap.add_argument("--max_side", type=int, default=1024) + args = ap.parse_args() + device = "cuda" if torch.cuda.is_available() else "cpu" + net, tail = build_model(args.weights, args.half_decoder, args.model_id, device, bf16=True) + metrics = load_metrics(device) + rows = [] + if args.pairs_json: + with open(args.pairs_json, encoding="utf-8") as fh: + data = json.load(fh) + pairs = data.get("real_pairs", data if isinstance(data, list) else []) + for i, p in enumerate(pairs): + with torch.no_grad(): + lr = Image.open(p["lr"]).convert("RGB") + hr = Image.open(p["hr"]).convert("RGB") + lr128 = lr.resize((max(1, lr.width // 4), max(1, lr.height // 4)), Image.LANCZOS) + t = torch.from_numpy(np.asarray(lr128, dtype=np.float32).transpose(2, 0, 1) / 255.0 * 2 - 1)[None].to(device) + with torch.autocast("cuda", dtype=torch.bfloat16): + z = net(t); sr_arr = tail(z) + sr = Image.fromarray(((sr_arr[0].float().cpu().numpy().transpose(1, 2, 0) + 1) / 2 * 255).clip(0, 255).astype(np.uint8)) + sr, hr = resize_pair(sr, hr, args.max_side) + m = metric_dict(metrics, sr, hr, device) + m["name"] = os.path.basename(p["lr"]); m["proxy"] = proxy(m) + rows.append(m) + print(f" [{i+1}] {m['name']} proxy {m['proxy']:.4f}", flush=True) + if args.official_lq and args.official_gt: + lq_files = sorted(os.listdir(args.official_lq)) + for f in lq_files: + lq_p = os.path.join(args.official_lq, f) + gt_p = os.path.join(args.official_gt, f.replace("_lq", "_gt").replace("lq.jpg", "gt.png")) + if not os.path.exists(gt_p): + gt_p = os.path.join(args.official_gt, f.replace("_lq.jpg", "_gt.jpg")) + if not os.path.exists(gt_p): + continue + sr = infer_4k_single(lq_p, net, tail, device) + gt = Image.open(gt_p).convert("RGB") + sr, gt = resize_pair(sr, gt, args.max_side) + m = metric_dict(metrics, sr, gt, device) + m["name"] = f; m["proxy"] = proxy(m) + rows.append(m) + print(f" [official] {f} proxy {m['proxy']:.4f}", flush=True) + if not rows: + print("无评测样本"); return + keys = ["psnr", "ssim", "lpips", "dists", "niqe", "maniqa", "musiq", "clipiqa", "proxy"] + agg = {k: float(np.mean([r[k] for r in rows])) for k in keys} + result = {"per_image": rows, "mean": agg} + os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True) + with open(args.out, "w", encoding="utf-8") as fh: + json.dump(result, fh, ensure_ascii=False, indent=1) + print(json.dumps(agg, indent=1)) + +if __name__ == "__main__": + main() diff --git a/src/export_jit.py b/src/export_jit.py new file mode 100644 index 0000000000000000000000000000000000000000..240b1b61132b87517a6287c9d7a2e13238bafa92 --- /dev/null +++ b/src/export_jit.py @@ -0,0 +1,62 @@ +#!/usr/bin/env python +"""导出 torch.jit(fp16) 512x512 模型: 输入 LR[1,3,512,512][-1,1] -> 输出同尺寸。 +内部: bicubic 512->128 -> 官方 4x 学生全链 -> 512。 (AdaIN 后处理不进测速模型) +用法: python src/export_jit.py --net weight/s2/net_params_X.pkl --out model_dir +""" +import argparse, os, sys +from pathlib import Path +REPO = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(REPO)); sys.path.insert(0, str(REPO / "src")); sys.path.insert(0, str(REPO / "official")) +import torch, torch.nn as nn, torch.nn.functional as F +from common import load_diffusers_sd, load_pruned_decoder, build_net + +class SR512(nn.Module): + def __init__(self, net, tail): + super().__init__() + self.net = net + self.tail = tail + + def forward(self, x512): + x128 = F.interpolate(x512, size=(128, 128), mode="bicubic", align_corners=False) + z = self.net(x128) + return self.tail(z) + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--net", required=True) + ap.add_argument("--out", default="model_dir") + ap.add_argument("--half_decoder", default="weight/pretrained/halfDecoder.ckpt") + ap.add_argument("--model_id", default="models/stable-diffusion-2-1-base") + ap.add_argument("--name", default="your_model.pt") + args = ap.parse_args() + os.makedirs(args.out, exist_ok=True) + device = "cuda" if torch.cuda.is_available() else "cpu" + vae, unet, _, _ = load_diffusers_sd(args.model_id, dtype=torch.float32, device="cpu") + del vae + decoder = load_pruned_decoder(args.half_decoder, device="cpu", dtype=torch.float32) + net = build_net(unet, decoder) + sd = torch.load(args.net, 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.eval() + tail = nn.Sequential(*decoder.up_blocks, decoder.conv_norm_out, + decoder.conv_act, decoder.conv_out).eval() + model = SR512(net, tail).to(device).half().eval() + # trace on fixed 512 fp16 + dummy = torch.randn(1, 3, 512, 512, device=device).half() * 0.5 + with torch.no_grad(): + traced = torch.jit.trace(model, dummy, check_trace=False) + traced = torch.jit.freeze(traced) + out_path = os.path.join(args.out, args.name) + traced.save(out_path) + # 自检: 两次前向一致性 + 形状 + with torch.no_grad(): + o1 = traced(dummy); o2 = traced(dummy) + assert o1.shape == dummy.shape, o1.shape + err = (o1 - o2).abs().max().item() + print("saved", out_path, "| deterministic max-diff:", err) + print("output range sample:", float(o1.min()), float(o1.max())) + +if __name__ == "__main__": + main() diff --git a/src/inference_4k.py b/src/inference_4k.py new file mode 100644 index 0000000000000000000000000000000000000000..ad69060ae5ed85c191703fcaed2b8fb2dd3db6c5 --- /dev/null +++ b/src/inference_4k.py @@ -0,0 +1,105 @@ +#!/usr/bin/env python +"""4K 整图增强: 512 窗 + overlap 128 线性融合 + AdaIN 对齐 LR(禁止 zero-padding, 边缘 reflect)。 +x1 语义: 对每个 512 窗先缩到 128, 走官方 4x 学生模型, 回 512。 +用法: python src/inference_4k.py --lr_dir --net weight/s2/net_params_X.pkl --out output_dir +""" +import argparse, copy, os, sys +from pathlib import Path +REPO = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(REPO)); sys.path.insert(0, str(REPO / "src")); sys.path.insert(0, str(REPO / "official")) +import torch, torch.nn.functional as F +import numpy as np +from PIL import Image +from common import load_diffusers_sd, load_pruned_decoder, build_net + +def build_model(net_pkl, half_decoder, model_id, device, bf16=True): + dtype = torch.bfloat16 if (bf16 and device == "cuda") else torch.float32 + vae, unet, _, _ = load_diffusers_sd(model_id, dtype=torch.float32, device="cpu") + del vae + decoder = load_pruned_decoder(half_decoder, device="cpu", dtype=torch.float32) + net = build_net(unet, decoder) + net.to(device=device, dtype=dtype) + sd = torch.load(net_pkl, 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.eval() + tail = torch.nn.Sequential(*decoder.up_blocks, decoder.conv_norm_out, + decoder.conv_act, decoder.conv_out).to(device=device, dtype=dtype).eval() + return net, tail + +def infer_tile(net, tail, lr128, lr_stats): + """lr128: [1,3,128,128] [-1,1]; 返回 [1,3,512,512] [-1,1] 并 AdaIN 对齐 lr_stats""" + with torch.no_grad(): + z = net(lr128) + out = tail(z) + out = out.float() + mu, std = out.mean(dim=(2, 3), keepdim=True), out.std(dim=(2, 3), keepdim=True) + lmu, lstd = lr_stats + out = (out - mu) / (std + 1e-6) * lstd + lmu + return out.clamp(-1, 1) + +def infer_4k_single(lr_path, net, tail, device, tile=512, overlap=128, bf16=True): + im = Image.open(lr_path).convert("RGB") + w, h = im.size + a = np.asarray(im, dtype=np.float32) / 255.0 # [H,W,3] + # reflect pad so tile 整除 + stride = tile - overlap + pad_r = (stride - w % stride) % stride + (tile - stride) if (w % stride) else max(0, tile - stride) + pad_b = (stride - h % stride) % stride + (tile - stride) if (h % stride) else max(0, tile - stride) + a = np.pad(a, ((0, pad_b), (0, pad_r), (0, 0)), mode="reflect") + H, W = a.shape[:2] + acc = np.zeros((H, W, 3), dtype=np.float64) + wsum = np.zeros((H, W, 1), dtype=np.float64) + rows = list(range(0, H - tile + 1, stride)) or [0] + cols = list(range(0, W - tile + 1, stride)) or [0] + if rows[-1] + tile < H: rows.append(H - tile) + if cols[-1] + tile < W: cols.append(W - tile) + # 线性窗权重 + ramp = np.minimum(np.arange(tile), np.arange(tile)[::-1]) / (tile // 2) + w2d = (ramp[None, :] * ramp[:, None])[..., None].astype(np.float64) + for y in rows: + for x in cols: + crop = a[y:y + tile, x:x + tile] + lr128 = np.asarray(Image.fromarray((crop * 255).astype(np.uint8)).resize((128, 128), Image.BICUBIC), + dtype=np.float32) / 255.0 + lr_t = torch.from_numpy(lr128.transpose(2, 0, 1))[None].to(device) * 2 - 1 + stats = (torch.tensor(crop.mean(axis=(0, 1))[None, :, None, None], device=device) * 2 - 1, + torch.tensor(crop.std(axis=(0, 1))[None, :, None, None], device=device)) + with torch.autocast("cuda", dtype=torch.bfloat16, enabled=bf16): + out = infer_tile(net, tail, lr_t, stats) + o = out[0].float().cpu().numpy().transpose(1, 2, 0) # [512,512,3] [-1,1] + o = (o + 1) / 2.0 + acc[y:y + tile, x:x + tile] += o.astype(np.float64) * w2d + wsum[y:y + tile, x:x + tile] += w2d + res = acc / (wsum + 1e-9) + res = res[:h, :w] + return Image.fromarray((np.clip(res, 0, 1) * 255).astype(np.uint8)) + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--lr_dir", required=True) + ap.add_argument("--net", required=True) + ap.add_argument("--out", required=True) + ap.add_argument("--half_decoder", default="weight/pretrained/halfDecoder.ckpt") + ap.add_argument("--model_id", default="models/stable-diffusion-2-1-base") + ap.add_argument("--tile", type=int, default=512) + ap.add_argument("--overlap", type=int, default=128) + ap.add_argument("--no_bf16", action="store_true") + args = ap.parse_args() + device = "cuda" if torch.cuda.is_available() else "cpu" + os.makedirs(args.out, exist_ok=True) + net, tail = build_model(args.net, args.half_decoder, args.model_id, device, bf16=not args.no_bf16) + files = sorted([f for f in os.listdir(args.lr_dir) if f.lower().endswith((".jpg", ".jpeg", ".png"))]) + import time + t0 = time.time() + for i, f in enumerate(files): + out_im = infer_4k_single(os.path.join(args.lr_dir, f), net, tail, device, + tile=args.tile, overlap=args.overlap, bf16=not args.no_bf16) + out_im.save(os.path.join(args.out, f), quality=95) + if (i + 1) % 10 == 0: + print(f" {i+1}/{len(files)} elapsed {time.time()-t0:.1f}s", flush=True) + print(f"done {len(files)} imgs in {time.time()-t0:.1f}s") + +if __name__ == "__main__": + main() diff --git a/src/latency_test.py b/src/latency_test.py new file mode 100644 index 0000000000000000000000000000000000000000..1671e55063bc4cb2e650bbcbff9722be56356651 --- /dev/null +++ b/src/latency_test.py @@ -0,0 +1,40 @@ +#!/usr/bin/env python +"""512x512 fp16 时延测速: 加载 torch.jit 模型, warmup + N 次前向。 +用法: python src/latency_test.py --model model_dir/your_model.pt [--osediff_latency 0.168] +""" +import argparse, time, statistics +import torch + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--model", required=True) + ap.add_argument("--n", type=int, default=100) + ap.add_argument("--warmup", type=int, default=10) + ap.add_argument("--osediff_latency", type=float, default=0.0, help="秒; 提供则打印加速比") + args = ap.parse_args() + device = "cuda" if torch.cuda.is_available() else "cpu" + m = torch.jit.load(args.model, map_location=device) + m.eval() + x = torch.randn(1, 3, 512, 512, device=device).half() + with torch.no_grad(): + for _ in range(args.warmup): + m(x) + torch.cuda.synchronize() if device == "cuda" else None + times = [] + for _ in range(args.n): + if device == "cuda": + torch.cuda.synchronize() + t0 = time.perf_counter() + with torch.no_grad(): + m(x) + if device == "cuda": + torch.cuda.synchronize() + times.append(time.perf_counter() - t0) + mean = statistics.mean(times) + med = statistics.median(times) + print(f"mean {mean*1000:.3f} ms | median {med*1000:.3f} ms | n={args.n}") + if args.osediff_latency > 0: + print(f"speedup vs OSEDiff({args.osediff_latency*1000:.1f}ms): {args.osediff_latency/mean:.2f}x") + +if __name__ == "__main__": + main() diff --git a/src/train_lora.py b/src/train_lora.py new file mode 100644 index 0000000000000000000000000000000000000000..66f115d825d5c7209a045c4009a53536e4322ff2 --- /dev/null +++ b/src/train_lora.py @@ -0,0 +1,254 @@ +#!/usr/bin/env python +"""S1: LoRA 像素域监督适配(单卡)。 +- 数据: synthetic(在线 Real-ESRGAN 退化) + real pairs(manifest) 混合(--real_prob) +- 模型: 官方 Net+halfDecoder 全链(输入 LR 128 -> 输出 RGB 512) +- 训练: 仅 LoRA 参数(手工注入, rank 可设), bf16, grad accum, save net/full state +用法示例见 scripts/run_stage1.sh +""" +import argparse, json, math, os, random, sys, time, copy +from pathlib import Path + +REPO = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(REPO)) +sys.path.insert(0, str(REPO / "src")) +sys.path.insert(0, str(REPO / "official")) + +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch.utils.data import Dataset, DataLoader +from omegaconf import OmegaConf + +from common import (ensure_official, load_diffusers_sd, load_pruned_decoder, + assemble_full_student, inject_lora, lora_params, + count_params, build_net) +ensure_official() +from dataset import RealESRGANDataset, RealESRGANDegrader # official + +# --------------------------------------------------------------------------- +# 小波高频 mask(纯 torch 单层 Haar,无额外依赖) +# --------------------------------------------------------------------------- +def haar_decomp(x): + """x: [B,C,H,W]; 返回 dict(ll,lh,hl,hh) 每个 [B,C,H/2,W/2]""" + B, C, H, W = x.shape + H2, W2 = H // 2, W // 2 + if H % 2 or W % 2: + x = F.pad(x, (0, W % 2, 0, H % 2)) + a = x.view(B, C, H2, 2, W2, 2) + ll = (a[:, :, :, 0, :, 0] + a[:, :, :, 0, :, 1] + a[:, :, :, 1, :, 0] + a[:, :, :, 1, :, 1]) / 4 + lh = (a[:, :, :, 0, :, 0] - a[:, :, :, 0, :, 1] + a[:, :, :, 1, :, 0] - a[:, :, :, 1, :, 1]) / 4 + hl = (a[:, :, :, 0, :, 0] + a[:, :, :, 0, :, 1] - a[:, :, :, 1, :, 0] - a[:, :, :, 1, :, 1]) / 4 + hh = (a[:, :, :, 0, :, 0] - a[:, :, :, 0, :, 1] - a[:, :, :, 1, :, 0] + a[:, :, :, 1, :, 1]) / 4 + return ll, lh, hl, hh + +def highfreq_mask(x): + with torch.no_grad(): + _, lh, hl, hh = haar_decomp(x.detach().float()) + e = (lh ** 2 + hl ** 2 + hh ** 2).sqrt() + m = e / (e.flatten(2).mean(dim=2, keepdim=True) + 1e-6).unsqueeze(-1) + m = F.interpolate(m, size=x.shape[-2:], mode="bilinear", align_corners=False) + return m + +# --------------------------------------------------------------------------- +# Real pairs dataset: manifest {"real_pairs":[{lr,hr}]} +# --------------------------------------------------------------------------- +class RealPairDataset(Dataset): + """?? x4 ?(RealSR/DRealSR): ?? HR 512 crop + ??? LR 128 crop? + ??? LQ~1K??/GT?4K ? 4x ??????????? 128 LR -> 512 HR? + manifest: {"real_pairs":[{lr,hr}]}?lr/hr ??? 4x ??lr ???hr/4?? + """ + + def __init__(self, manifest_path, patch=512, scale=4, seed=0): + with open(manifest_path, encoding="utf-8") as fh: + m = json.load(fh) + self.pairs = m.get("real_pairs", []) + self.patch = patch + self.scale = scale + self.lr_patch = patch // scale + self.rng = random.Random(seed) + self._cache = {} + + def __len__(self): + return max(1, len(self.pairs) * 40) + + def __getitem__(self, idx): + from PIL import Image + from torchvision import transforms + p = self.pairs[idx % len(self.pairs)] + key = p["hr"] + if key not in self._cache: + hr = Image.open(p["hr"]).convert("RGB") + lr = Image.open(p["lr"]).convert("RGB") + self._cache[key] = (lr, hr) + lr, hr = self._cache[key] + w, h = hr.size + if w < self.patch or h < self.patch: + raise RuntimeError(f"HR ?? patch: {p['hr']} {hr.size}") + x = self.rng.randint(0, w - self.patch) + y = self.rng.randint(0, h - self.patch) + hr_c = hr.crop((x, y, x + self.patch, y + self.patch)) + # LR ??????: ??? lr/hr ????, ????? 128 + sc_w, sc_h = lr.width / w, lr.height / h + lx0, ly0 = int(x * sc_w), int(y * sc_h) + lx1, ly1 = int((x + self.patch) * sc_w), int((y + self.patch) * sc_h) + lx1 = min(lx1, lr.width); ly1 = min(ly1, lr.height) + lr_c = lr.crop((lx0, ly0, lx1, ly1)).resize((self.lr_patch, self.lr_patch), Image.BICUBIC) + to_t = transforms.ToTensor() + lr_t = to_t(lr_c) * 2 - 1 + hr_t = to_t(hr_c) * 2 - 1 + if self.rng.random() < 0.5: + lr_t = torch.flip(lr_t, dims=[2]); hr_t = torch.flip(hr_t, dims=[2]) + return lr_t, hr_t + +# --------------------------------------------------------------------------- +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--config", default="configs/config_s1_lora.yml") + ap.add_argument("--manifest", default="data/manifest_train.json", help="含 real_pairs 的训练清单") + ap.add_argument("--real_prob", type=float, default=0.35) + ap.add_argument("--model_id", default="models/stable-diffusion-2-1-base") + ap.add_argument("--half_decoder", default="weight/pretrained/halfDecoder.ckpt") + ap.add_argument("--init_net", default="weight/net_params_200.pkl", help="官方学生权重(Net state)") + ap.add_argument("--out", default="weight/s1") + ap.add_argument("--log_dir", default="logs/s1") + ap.add_argument("--steps", type=int, default=20000) + ap.add_argument("--batch_size", type=int, default=8) + ap.add_argument("--grad_accum", type=int, default=2) + ap.add_argument("--lr", type=float, default=5e-5) + ap.add_argument("--lora_rank", type=int, default=64) + ap.add_argument("--lora_alpha", type=float, default=1.0) + ap.add_argument("--save_every", type=int, default=2000) + ap.add_argument("--w_l1", type=float, default=1.0) + ap.add_argument("--w_lpips", type=float, default=1.0) + ap.add_argument("--w_dists", type=float, default=0.3) + ap.add_argument("--w_wave", type=float, default=0.5) + ap.add_argument("--w_color", type=float, default=0.2) + ap.add_argument("--bf16", action="store_true", default=True) + ap.add_argument("--no_bf16", dest="bf16", action="store_false") + ap.add_argument("--seed", type=int, default=123) + ap.add_argument("--num_workers", type=int, default=8) + args = ap.parse_args() + random.seed(args.seed); torch.manual_seed(args.seed) + device = "cuda" if torch.cuda.is_available() else "cpu" + cfg = OmegaConf.load(args.config) + os.makedirs(args.out, exist_ok=True); os.makedirs(args.log_dir, exist_ok=True) + log_path = os.path.join(args.log_dir, "train_lora.log") + logf = open(log_path, "a", encoding="utf-8") + def log(msg): + print(msg, flush=True); logf.write(msg + "\n"); logf.flush() + + # ---- data ---- + syn_ds = RealESRGANDataset(cfg, args.batch_size) + syn_dl = DataLoader(syn_ds, batch_size=args.batch_size, num_workers=args.num_workers, shuffle=True) + degrader = RealESRGANDegrader(cfg, device) + real_ds = RealPairDataset(args.manifest) if args.real_prob > 0 else None + real_dl = DataLoader(real_ds, batch_size=args.batch_size, num_workers=args.num_workers, shuffle=True) if real_ds else None + + # ---- model ---- + dtype = torch.bfloat16 if (args.bf16 and device == "cuda") else torch.float32 + vae, unet, text_encoder, tokenizer = load_diffusers_sd(args.model_id, dtype=torch.float32, device="cpu") + del text_encoder, tokenizer, vae + decoder = load_pruned_decoder(args.half_decoder, device="cpu", dtype=torch.float32) + full = assemble_full_student(unet, decoder, net_weights=args.init_net, device=device, dtype=dtype) + inject_lora(full, rank=args.lora_rank, alpha=args.lora_alpha) + params = list(lora_params(full)) + log(f"trainable params: {count_params(full, only_trainable=True)/1e6:.2f}M") + optimizer = torch.optim.AdamW(params, lr=args.lr) + sched = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.steps) + scaler = torch.amp.GradScaler("cuda", enabled=(dtype == torch.float16)) if (device == "cuda" and dtype == torch.float16) else None + + # lpips / dists 可选 + lpips_fn = None + if args.w_lpips > 0: + try: + import lpips + lpips_fn = lpips.LPIPS(net="alex").to(device).eval() + for p in lpips_fn.parameters(): p.requires_grad_(False) + except Exception as e: + log(f"[warn] lpips 不可用: {e}; 该损失置 0") + dists_fn = None + if args.w_dists > 0: + try: + import pyiqa + dists_fn = pyiqa.create_metric("dists", device=device) + for p in dists_fn.parameters(): p.requires_grad_(False) + except Exception as e: + log(f"[warn] pyiqa dists 不可用: {e}; 该损失置 0") + + # ---- train loop ---- + syn_iter = iter(syn_dl); real_iter = iter(real_dl) if real_dl else None + step = 0 + full.train() + optimizer.zero_grad(set_to_none=True) + while step < args.steps: + use_real = real_dl is not None and random.random() < args.real_prob + try: + if use_real: + lr_t, hr_t = next(real_iter) + else: + batch = next(syn_iter) + lr_t, hr_t = degrader.degrade(batch) + except StopIteration: + syn_iter = iter(syn_dl) + real_iter = iter(real_dl) if real_dl else None + continue + lr_t, hr_t = lr_t.to(device), hr_t.to(device) + if dtype == torch.bfloat16: + with torch.autocast("cuda", dtype=torch.bfloat16): + out = full(lr_t) + loss, items = _losses(out.float(), hr_t.float(), full, args, lpips_fn, dists_fn) + else: + out = full(lr_t) + loss, items = _losses(out, hr_t, full, args, lpips_fn, dists_fn) + (loss / args.grad_accum).backward() + if (step + 1) % args.grad_accum == 0: + if scaler is not None: + scaler.step(optimizer); scaler.update() + else: + optimizer.step() + optimizer.zero_grad(set_to_none=True) + sched.step() + if (step + 1) % 50 == 0: + log(f"step {step+1}/{args.steps} loss {loss.item():.4f} " + + " ".join(f"{k}:{v:.4f}" for k, v in items.items())) + if (step + 1) % args.save_every == 0: + _save(full, args.out, step + 1) + step += 1 + _save(full, args.out, step) + log("S1 done") + +def _losses(out, hr, full, args, lpips_fn, dists_fn): + out = out.float(); hr = hr.float() + l1 = F.l1_loss(out, hr) + items = {"l1": l1.item()} + total = args.w_l1 * l1 + if lpips_fn is not None: + try: + lp = lpips_fn(out.clamp(-1, 1), hr.clamp(-1, 1)).mean() + total = total + args.w_lpips * lp; items["lpips"] = lp.item() + except Exception: + pass + if dists_fn is not None: + try: + d = dists_fn((out.clamp(-1,1)+1)/2, (hr.clamp(-1,1)+1)/2).mean() + total = total + args.w_dists * d; items["dists"] = d.item() + except Exception: + pass + if args.w_wave > 0: + mask = highfreq_mask(out.detach()) + wav = (mask * (out - hr).abs()).mean() + total = total + args.w_wave * wav; items["wave"] = wav.item() + if args.w_color > 0: + mo, so = out.mean(dim=(2,3)), out.std(dim=(2,3)) + mh, sh = hr.mean(dim=(2,3)), hr.std(dim=(2,3)) + col = (mo - mh).abs().mean() + (so - sh).abs().mean() + total = total + args.w_color * col; items["color"] = col.item() + return total, items + +def _save(full, out_dir, step): + net = full[0] # Net (first module of full chain) + torch.save(net.state_dict(), os.path.join(out_dir, f"net_params_{step}.pkl")) + torch.save(full.state_dict(), os.path.join(out_dir, f"full_params_{step}.pkl")) + +if __name__ == "__main__": + main() diff --git a/src/train_stage2.py b/src/train_stage2.py new file mode 100644 index 0000000000000000000000000000000000000000..1f3c95b696c2847eebf975cbf570d62719781989 --- /dev/null +++ b/src/train_stage2.py @@ -0,0 +1,221 @@ +#!/usr/bin/env python +"""S2/S3: 特征对抗蒸馏(官方 Stage-2 扩展;单卡, bf16, grad_accum)。 +teacher: osediff(默认) | gdpo | osediff_ft +损失: base(特征L1蒸馏+对抗) 可选 + wavelet / 像素 LPIPS+DISTS / AdaIN 色彩 +""" +import argparse, copy, os, random, sys +from pathlib import Path + +REPO = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(REPO)); sys.path.insert(0, str(REPO / "src")); sys.path.insert(0, str(REPO / "official")) + +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch.utils.data import DataLoader +from omegaconf import OmegaConf + +from common import (ensure_official, load_diffusers_sd, load_pruned_decoder, + build_net, build_discriminator_unet, load_osediff_teacher, + load_gdpo_teacher, count_params) +ensure_official() +from dataset import RealESRGANDataset, RealESRGANDegrader # official +from train_lora import highfreq_mask # 复用小波高频 mask + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--config", default="configs/config_s2_distill.yml") + ap.add_argument("--model_id", default="models/stable-diffusion-2-1-base") + ap.add_argument("--half_decoder", default="weight/pretrained/halfDecoder.ckpt") + ap.add_argument("--osediff_pkl", default="weight/pretrained/osediff.pkl") + ap.add_argument("--gdpo_dir", default="", help="GDPO 权重(teacher=gdpo)") + ap.add_argument("--teacher", choices=["osediff", "gdpo", "osediff_ft"], default="osediff") + ap.add_argument("--teacher_ft_pkl", default="", help="teacher=osediff_ft 的 unet state pkl") + ap.add_argument("--init_net", default="weight/net_params_200.pkl", help="学生 Net 初始化/续训") + ap.add_argument("--out", default="weight/s2"); ap.add_argument("--log_dir", default="logs/s2") + ap.add_argument("--steps", type=int, default=30000) + ap.add_argument("--batch_size", type=int, default=8); ap.add_argument("--grad_accum", type=int, default=2) + ap.add_argument("--lr", type=float, default=5e-5); ap.add_argument("--lr_D", type=float, default=1e-6) + ap.add_argument("--save_every", type=int, default=2000) + ap.add_argument("--w_distil", type=float, default=1.0); ap.add_argument("--w_adv", type=float, default=1.0) + ap.add_argument("--w_wave", type=float, default=0.0) + ap.add_argument("--w_lpips", type=float, default=0.0); ap.add_argument("--w_dists", type=float, default=0.0) + ap.add_argument("--w_color", type=float, default=0.0) + ap.add_argument("--skip_ram", action="store_true") + ap.add_argument("--bf16", action="store_true", default=True); ap.add_argument("--no_bf16", dest="bf16", action="store_false") + ap.add_argument("--seed", type=int, default=123); ap.add_argument("--num_workers", type=int, default=8) + args = ap.parse_args() + random.seed(args.seed); torch.manual_seed(args.seed) + device = "cuda" if torch.cuda.is_available() else "cpu" + cfg = OmegaConf.load(args.config) + os.makedirs(args.out, exist_ok=True); os.makedirs(args.log_dir, exist_ok=True) + logf = open(os.path.join(args.log_dir, "train_stage2.log"), "a", encoding="utf-8") + def log(msg): + print(msg, flush=True); logf.write(msg + "\n"); logf.flush() + log(f"teacher={args.teacher} steps={args.steps} bs={args.batch_size} accum={args.grad_accum}") + + use_bf16 = args.bf16 and device == "cuda" + amp = torch.autocast("cuda", dtype=torch.bfloat16) if use_bf16 else torch.autocast("cuda", enabled=False) + + # ---- SD2.1 组件(先保留 pristine unet 副本, 供教师/判别器) ---- + vae, unet, text_encoder, tokenizer = load_diffusers_sd(args.model_id, dtype=torch.float32, device="cpu") + vae = vae.to(device); unet = unet.to(device); text_encoder = text_encoder.to(device) + unet_pristine = copy.deepcopy(unet) # 判别器/教师用(Net 会修改 unet) + vae_teacher = copy.deepcopy(vae) + unet_teacher = copy.deepcopy(unet_pristine) + + # 教师权重 + if args.teacher == "osediff": + ck = load_osediff_teacher(args.osediff_pkl) + vae_teacher.load_state_dict(ck["vae"]); unet_teacher.load_state_dict(ck["unet"]) + elif args.teacher == "gdpo": + g = load_gdpo_teacher(args.gdpo_dir) + unet_teacher.load_state_dict(g.state_dict() if hasattr(g, "state_dict") else g) + elif args.teacher == "osediff_ft": + sd = torch.load(args.teacher_ft_pkl, map_location="cpu", weights_only=False) + if any(k.startswith("unet.") for k in sd): + sd = {k.replace("unet.", "", 1): v for k, v in sd.items()} + unet_teacher.load_state_dict(sd) + + # 判别器(基于 pristine unet) 先建, 再让 Net 修改 unet + unet_D = build_discriminator_unet(unet_pristine, rank=4, device=device, + dtype=torch.bfloat16 if use_bf16 else torch.float32) + + decoder = load_pruned_decoder(args.half_decoder, device="cpu", + dtype=torch.bfloat16 if use_bf16 else torch.float32).to(device) + net = build_net(unet, decoder) + net.to(device=device, dtype=torch.bfloat16 if use_bf16 else torch.float32) + if args.init_net: + sd = torch.load(args.init_net, 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); log("student init " + args.init_net) + tail = nn.Sequential(*decoder.up_blocks, decoder.conv_norm_out, + decoder.conv_act, decoder.conv_out) + tail.to(device=device, dtype=torch.bfloat16 if use_bf16 else torch.float32) + + for m in (vae, unet, text_encoder, vae_teacher, unet_teacher, decoder): + m.requires_grad_(False); m.eval() + if use_bf16: + for m in (vae, unet, vae_teacher, unet_teacher, text_encoder): + m.to(torch.bfloat16) + for n, p in unet_D.named_parameters(): + p.requires_grad_("lora" in n or "conv_in" in n) + log(f"student trainable: {(count_params(net)+count_params(tail))/1e6:.1f}M") + + # ---- data ---- + dataset = RealESRGANDataset(cfg, args.batch_size) + degrader = RealESRGANDegrader(cfg, device) + dataloader = DataLoader(dataset, batch_size=args.batch_size, num_workers=args.num_workers, shuffle=True) + + # ---- DAPE ---- + from torchvision import transforms + ram_tf = transforms.Compose([transforms.Resize((384, 384)), + transforms.Normalize(mean=[0.485, 0.456, 0.406], + std=[0.229, 0.224, 0.225])]) + if args.skip_ram: + class _Dummy: + def generate_tag(self, image): + return ["a photo of a city street, clean, high-resolution"] + DAPE = _Dummy() + else: + from ram.models.ram_lora import ram + DAPE = ram(pretrained="weight/pretrained/ram_swin_large_14m.pth", + pretrained_condition="weight/pretrained/DAPE.pth", + image_size=384, vit="swin_l").eval().to(device) + if use_bf16: DAPE = DAPE.to(torch.bfloat16) + + opt_s = torch.optim.AdamW(list(net.parameters()) + list(tail.parameters()), lr=args.lr) + opt_D = torch.optim.AdamW([p for p in unet_D.parameters() if p.requires_grad], lr=args.lr_D) + sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt_s, T_max=args.steps) + from diffusers import DDIMScheduler + sched_cfg = DDIMScheduler.from_pretrained(args.model_id, subfolder="scheduler") + alpha = sched_cfg.alphas_cumprod[999].to(device) + del sched_cfg + + lpips_fn = dists_fn = None + if args.w_lpips > 0: + try: + import lpips + lpips_fn = lpips.LPIPS(net="alex").to(device).eval() + for p in lpips_fn.parameters(): p.requires_grad_(False) + except Exception as e: log(f"[warn] lpips 不可用 {e}") + if args.w_dists > 0: + try: + import pyiqa + dists_fn = pyiqa.create_metric("dists", device=device) + for p in dists_fn.parameters(): p.requires_grad_(False) + except Exception as e: log(f"[warn] pyiqa dists 不可用 {e}") + + dl_iter = iter(dataloader); step = 0 + opt_s.zero_grad(set_to_none=True); opt_D.zero_grad(set_to_none=True) + loss_D = torch.tensor(0.0, device=device) + while step < args.steps: + try: + batch = next(dl_iter) + except StopIteration: + dl_iter = iter(dataloader); continue + LR, HR = degrader.degrade(batch) # [0,1] range + LR, HR = LR.to(device), HR.to(device) + with torch.no_grad(), amp: + tag = DAPE.generate_tag(ram_tf(LR))[0] + text_input = tokenizer(tag, max_length=tokenizer.model_max_length, + padding="max_length", truncation=True, + return_tensors="pt").to(device) + enc = text_encoder(text_input.input_ids, return_dict=False)[0] + LRn, HRn = LR * 2 - 1, HR * 2 - 1 + LR_ = F.interpolate(LRn, scale_factor=4, mode="bicubic") + LR_lat = vae_teacher.encode(LR_).latent_dist.mean * vae_teacher.config.scaling_factor + HR_lat = vae.encode(HRn).latent_dist.mean + t_all = torch.full((LR_lat.shape[0],), 999, dtype=torch.long, device=device) + pred_t = unet_teacher(LR_lat, t_all, encoder_hidden_states=enc, return_dict=False)[0] + z0_t = (LR_lat - ((1 - alpha) ** 0.5) * pred_t) / (alpha ** 0.5) + z0_t = vae_teacher.post_quant_conv(z0_t / vae_teacher.config.scaling_factor) + z0_t = decoder.conv_in(z0_t); z0_t = decoder.mid_block(z0_t) + z0_gt = vae.post_quant_conv(HR_lat) + z0_gt = decoder.conv_in(z0_gt); z0_gt = decoder.mid_block(z0_gt) + with amp: + z0_s = net(LRn) + tB = torch.full((z0_s.shape[0],), 999, dtype=torch.long, device=device) + d_s = unet_D(z0_s, tB, encoder_hidden_states=enc, return_dict=False)[0] + loss_distil = (z0_s - z0_t).abs().mean() + loss_adv = F.softplus(-d_s).mean() + if args.w_wave > 0: + m = highfreq_mask(z0_s.detach()) + loss_distil = (m * (z0_s - z0_t).abs()).mean() + loss_adv = (m * F.softplus(-d_s)).mean() + loss = args.w_distil * loss_distil + args.w_adv * loss_adv + if args.w_lpips + args.w_dists + args.w_color > 0: + img_s = tail(z0_s) + if lpips_fn is not None and args.w_lpips > 0: + loss = loss + args.w_lpips * lpips_fn(img_s.float().clamp(-1, 1), HRn.float().clamp(-1, 1)).mean() + if dists_fn is not None and args.w_dists > 0: + loss = loss + args.w_dists * dists_fn((img_s.float().clamp(-1, 1) + 1) / 2, + (HRn.float().clamp(-1, 1) + 1) / 2).mean() + if args.w_color > 0: + loss = loss + args.w_color * ( + (img_s.float().mean(dim=(2, 3)) - HRn.float().mean(dim=(2, 3))).abs().mean() + + (img_s.float().std(dim=(2, 3)) - HRn.float().std(dim=(2, 3))).abs().mean()) + (loss / args.grad_accum).backward() + if (step + 1) % args.grad_accum == 0: + opt_s.step(); opt_s.zero_grad(set_to_none=True); sched.step() + with amp: + tB = torch.full((z0_s.shape[0],), 999, dtype=torch.long, device=device) + pred_real = unet_D(z0_gt.detach(), tB, encoder_hidden_states=enc, return_dict=False)[0] + pred_fake = unet_D(z0_s.detach(), tB, encoder_hidden_states=enc, return_dict=False)[0] + loss_D = F.softplus(pred_fake).mean() + F.softplus(-pred_real).mean() + (loss_D / args.grad_accum).backward() + opt_D.step(); opt_D.zero_grad(set_to_none=True) + if (step + 1) % 50 == 0: + log(f"step {step+1}/{args.steps} g {loss.item():.4f} distil {loss_distil.item():.4f} " + f"adv {loss_adv.item():.4f} D {loss_D.item():.4f}") + if (step + 1) % args.save_every == 0: + torch.save(net.state_dict(), os.path.join(args.out, f"net_params_{step+1}.pkl")) + torch.save(tail.state_dict(), os.path.join(args.out, f"tail_params_{step+1}.pkl")) + step += 1 + torch.save(net.state_dict(), os.path.join(args.out, f"net_params_{step}.pkl")) + torch.save(tail.state_dict(), os.path.join(args.out, f"tail_params_{step}.pkl")) + log("done") + +if __name__ == "__main__": + main()