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

CSIGv3 AdcSR train scripts + A100 runbook

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. NOTES.md +107 -0
  2. README.md +36 -0
  3. RUNBOOK_A100.md +51 -0
  4. configs/config_base.yml +35 -0
  5. configs/config_s1_lora.yml +37 -0
  6. configs/config_s2_distill.yml +37 -0
  7. configs/config_s3_ft.yml +37 -0
  8. configs/config_smoke.yml +37 -0
  9. official/bsr/__pycache__/degradations.cpython-312.pyc +0 -0
  10. official/bsr/degradations.py +764 -0
  11. official/bsr/transforms.py +179 -0
  12. official/bsr/utils/__init__.py +47 -0
  13. official/bsr/utils/color_util.py +208 -0
  14. official/bsr/utils/diffjpeg.py +515 -0
  15. official/bsr/utils/dist_util.py +82 -0
  16. official/bsr/utils/download_util.py +98 -0
  17. official/bsr/utils/file_client.py +167 -0
  18. official/bsr/utils/flow_util.py +170 -0
  19. official/bsr/utils/img_process_util.py +83 -0
  20. official/bsr/utils/img_util.py +172 -0
  21. official/bsr/utils/lmdb_util.py +199 -0
  22. official/bsr/utils/logger.py +213 -0
  23. official/bsr/utils/matlab_functions.py +178 -0
  24. official/bsr/utils/misc.py +141 -0
  25. official/bsr/utils/options.py +218 -0
  26. official/bsr/utils/plot_util.py +83 -0
  27. official/bsr/utils/registry.py +88 -0
  28. official/dataset.py +280 -0
  29. official/evaluate.py +55 -0
  30. official/forward.py +67 -0
  31. official/model.py +152 -0
  32. official/test.py +70 -0
  33. official/train.py +225 -0
  34. official/utils.py +24 -0
  35. requirements.txt +17 -0
  36. scripts/archive_run.py +31 -0
  37. scripts/download_weights.sh +32 -0
  38. scripts/env_check.py +25 -0
  39. scripts/probe_gdpo.py +13 -0
  40. scripts/run_smoke.sh +12 -0
  41. scripts/run_stage1.sh +14 -0
  42. scripts/run_stage2.sh +18 -0
  43. scripts/run_stage3.sh +14 -0
  44. scripts/setup_env.sh +21 -0
  45. src/check_submission.py +48 -0
  46. src/common.py +213 -0
  47. src/eval_val.py +104 -0
  48. src/export_jit.py +62 -0
  49. src/inference_4k.py +105 -0
  50. src/latency_test.py +40 -0
NOTES.md ADDED
@@ -0,0 +1,107 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # 实施与 S0 冒烟要点(给下一执行者/服务器环境)
2
+
3
+ ## 已实现(本地已验证)
4
+ - 工程脚手架 + 官方 clone(official/) + git 已提交(5a1afb7)
5
+ - tools/*.py 全部 CPU 可跑: clean_images(阈值/去重/可复制到 clean) / tag_categories(7类启发式) /
6
+ build_manifest / split_val(70组隔离) / build_patches(>=20k patch) / download_batch(编排)
7
+ 已用合成小图冒烟: 模糊/全黑/过曝/纯色正确剔除, 打标与 patch 构建正常。
8
+ - src/*.py 与 scripts/*.py 语法编译通过(py_compile)。
9
+
10
+ ## 必须上 A100 用真实权重验证的点(S0, 1 迭代/配置)
11
+ 1. scripts/setup_env.sh + scripts/download_weights.sh(Plan A) 后运行 scripts/env_check.py。
12
+ 2. scripts/run_smoke.sh: train_lora 与 train_stage2 各 1 迭代; 重点核对:
13
+ - Net+halfDecoder 全链装配 forward 形状: LR[1,3,128,128] -> RGB[1,3,512,512]
14
+ (若 shape 报错, 以 official/test.py 官方装配为唯一基准修正 src/common.build_net/assemble_full_student)
15
+ - teacher z0 前向与判别器 dtype/shape 匹配(bf16 + autocast)
16
+ - loss 有限值且可保存 checkpoint、可续训
17
+ 3. 基线复现: 用官方 net_params_200.pkl + official/test.py 在 RealSR 复现论文量级指标;
18
+ 并在 A100 记录 OSEDiff 512 fp16 时延(基准分母)。
19
+ 4. GDPO 教师: 先 python scripts/probe_gdpo.py <权重路径> 确认键格式; 若格式不兼容,
20
+ 按 probe 输出在 src/common.load_gdpo_teacher 补映射, 否则回退 --teacher osediff(默认已支持)。
21
+
22
+ ## 执行顺序(每阶段 gate 见方案 §6)
23
+ 1. tools 数据构建 B1-B4 -> clean -> tag -> patches>=20k -> manifest -> split_val(70)
24
+ 2. pack_upload.ps1 上传代码+数据; 服务器下载权重
25
+ 3. S0 smoke -> S1(LoRA) -> eval_val 选 best -> S2(GDPO蒸馏) -> eval_val 选 best -> S3
26
+ 4. 每次选点: 用 src/eval_val.py 输出 proxy, 把最优权重软链/复制为
27
+ weight/s1/net_params_BEST.pkl (run_stage2 引用) / weight/s2/net_params_BEST.pkl (run_stage3 引用)
28
+ 5. 速度: src/latency_test.py; 出图: src/inference_4k.py; 导出: src/export_jit.py; 提交包检查: src/check_submission.py
29
+
30
+ ## 已知待办/说明
31
+ - 训练脚本仅在服务器有 torch/权重时可运行; 本地 3050 4G 只做代码与数据开发。
32
+ - run_stage1/2/3.sh 里的 net_params_BEST.pkl 为占位名, 按实际 eval 结果更新。
33
+ - inference_4k/export_jit 采用“512窗->128->官方4x链->512”的 x1 包装语义;
34
+ 若官方评测机 runner 口径不同, 以官方样例 runner 实测为准调整 export_jit.SR512。
35
+ - 类别打标为启发式(文件名关键词+亮度/绿色占比), 抽样人工复核后再上传训练。
36
+
37
+ ## 试点批次进度 (B3 NKUSR8K 首 100 张, 2026-09-05)
38
+ - raw 100 张(~1.9GB) -> clean 90 张(10 张 aHash 重复剔除; 8K 图质量均达标) -> 已删 raw 释放空间
39
+ - 分类: city_view 41 / night_building 26 / plant 23 (skyscraper/shop_street/residential 为 0, 需 4KLSDB/RealSR 等源补充)
40
+ - patches: 2998 张 512x512 (city_view 1366 / night_building 866 / plant 766), manifest: data/manifest_train.json
41
+ - 人工复核图: data/review/NKUSR8K/*.jpg (每类 20 张拼图)
42
+ - 注意: synthetic patch 的 val 划分须按"源图"分组(文件名前缀)而非单 patch, 防同源泄露; 最终 val 主要来自 real pairs。
43
+
44
+ ## 数据集构建进度 (2026-09-05 晚 更新, 本批次完成)
45
+ - NKUSR8K 全量 1000 张: 下载 900(跳过试点100) -> 清洗 kept 639 (dup 259/过曝 2) -> 已删 raw(17.9GB 释放)
46
+ - 合并 clean_all/NKUSR8K = 729 张(试点90 + 全量639, 4096 长边 jpg, 文件名前缀 NKUSR8K_)
47
+ - 启发式分类: night_building 231 / plant 255 / city_view 243 (skyscraper/shop/residential 启发式分不出, 无碍训练)
48
+ - 源级 val 划分: 70 源留作 val (split_val_sources.json), train 源 659
49
+ - patches: data/patches 23,999 (city 7866/night 7684/plant 8449) + data/patches_val 560 (val 源)
50
+ - manifest_train.json: synthetic 23,999 + RealSR 训练对 400; manifest_val.json: synthetic 560 + RealSR 测试对 99
51
+ - RealSR 对清单: clean_pairs 配对命名 Canon_001_LR4.png <-> Canon_001_HR.png (build_manifest 已按 scene key 匹配)
52
+ - 人工复核拼图: data/review/NKUSR8K_all/{city_view,night_building,plant}.jpg
53
+
54
+ ## 待办
55
+ - [ ] DRealSR 训练集: 需人工(百度网盘 osiy / Google Drive 文件夹), 见 data/SOURCES.md; 到手后走 clean_pairs 入 real_pairs
56
+ - [ ] 4KLSDB: 默认跳过(稀疏类抽取不划算, 结论见 SOURCES.md); 若需补量可用 tools 的 RangeReader 思路实现 extract_klsdb_rg.py
57
+ - [ ] 打包上传: scripts/pack_upload.ps1 (需密码), 服务器 setup_env.sh + download_weights.sh + S0 smoke -> S1/S2/S3
58
+ - [ ] 本地清理: 确认后删除 data/clean/NKUSR8K* / data/clean_all(已并入 patches, 仅留存备查)
59
+
60
+ ## 2026-09-05 深夜: CLIP/4KLSDB/稀疏类抓取 进展
61
+ - 修正: NKUSR8K 无夜景(CLIP 确认 night_building=0); tag_categories 夜景仅允许文件名关键词, 不再按暗度推断
62
+ - CLIP(openai/clip-vit-base-patch32, CPU, transformers5.x 需 .pooler_output) 打标 NKUSR8K 729:
63
+ skyscraper 204 / plant 354 / residential 63 / other_building 41 / storefront 22 / mall 21 / city_view 5, reject 19
64
+ - 分类升级为 8 类: night_building / skyscraper / storefront(店铺门牌) / mall(商场外拍) / plant /
65
+ residential / city_view / other_building (tools: clip_filter.py, tag_categories.py)
66
+ - 4KLSDB 抽取管线落地(全部本地可跑):
67
+ klsdb_index.py(206 分片 id->shard/rg 索引, ~27min) + klsdb_extract.py(rank 按 caption 评分排序 + extract 按 row-group 拉 hr)
68
+ 第一批: 18 row-group, 11.6GB 流量 -> 1060 张 4K 原始 -> clean 802 张(data/clean/4KLSDB_b1)
69
+ - 稀疏类原图 URL 抓取(fetch_klsdb_urls.py, 成功率 ~96%): night_building 575 / mall 355 /
70
+ storefront 501 / residential 495 / skyscraper 486 -> 后台逐类清洗中(data/clean/4KLSDB_url/<cls>)
71
+ - DRealSR: 用户给的 GDrive 文件夹(1tP5m4k1...)只含 Test_x2/x3/x4.zip(测试集); 训练集在百度网盘 osiy(需登录)
72
+ - SSDLite: 本地 Windows torchvision CPU wheel ABI 不兼容(2.7/2.6+cpu 均报 torchvision::nms), 已放弃本地;
73
+ 方案: 本地用 CLIP; SSDLite person/物过滤代码留到 A100(Linux+torchvision 正常) 或后续换环境再启用
74
+ - 待办: URL 清洗完成 -> CLIP 复核(去 people/indoor/other) -> 全池 pHash 去重 -> 合并 clean_all(带源前缀)
75
+ -> 8 类平衡报表 -> sample_review 拼图人工复核 -> build_patches(>=? 目标 5000-7000 HR 源后 50k+ patch)
76
+ -> 再跑 4KLSDB b2/b3 补量到总池 5000-7000
77
+
78
+ ## 2026-09-06 凌晨: LIU4K v2 引入 + 策略更新
79
+ - 4KLSDB 政策: 建筑/风景/花草均可入池(增强鲁棒), 夜景适量即可(不堆量)
80
+ - LIU4K v2 (GDrive 1FtVQtY2t_ecuy_gzJqZ-CatqrJBAdq_d): 4 组多段zip(Animal/Building/Street/Mountain),
81
+ 只取 Building/Street/Mountain; Building = 400 张高分辨率 PNG(3992x2242~6208x4139, Pexels 建筑图, 10.2GB)
82
+ 解压用 NVIDIA App 7-Zip 22.01(自带 7z.exe); 旧 7za 9.20 不支持 spanned zip
83
+ - DRealSR: 放弃训练(用户确认); SSDLite: 取消(用户确认); 过滤以 CLIP 为主
84
+ - URL 稀疏类清洗后: night_building 409 / storefront 446 / mall 248 / residential 339 / skyscraper 331 (共 1773)
85
+ - 4KLSDB b1 CLIP 后 accept 604 (plant172/other_building171/residential124/city_view66/storefront35/skyscraper29/mall6/night1), reject 198
86
+ - NKUSR8K CLIP accept 710 (skyscraper204/plant354/residential63/other_building41/storefront22/mall21/city_view5), reject 19
87
+ - 新工具: merge_reports.py(多源合并/剔 reject/去重) + download_gdrive_parts.py(分卷断点下载)
88
+
89
+ ## 2026-09-06 凌晨2: 四源合并池(CLIP accept + URL文件夹标签 + reject剔除)
90
+ - 池: logs/cat_pool_all.json, 共 3072 张清洗后 HR 原图
91
+ night_building 332 / skyscraper 575 / storefront 445 / mall 252 / plant 540 / residential 526 /
92
+ city_view 102 / other_building 300
93
+ - 人工复核拼图: data/review/pool_all/{8类}.jpg (每类 20 张)
94
+ - LIU4K Building: 400 PNG(10.2GB) -> clean 320 -> jpg(clean_all/LIU4K) -> CLIP accept 235(reject 85 室内/人像)
95
+ - 剩余待补: city_view(102, 目标>=400), mall(252); 总池 3072 -> 目标 5000-7000
96
+ - 下一批: 4KLSDB b2/b3(城景/商场/其他建筑加量) + LIU4K Street/Mountain(城市街道/山地风景)
97
+
98
+ ## 2026-09-06 凌晨: 收尾流水线启动(finalize_all.ps1 后台自动链)
99
+ - 补量URL抓取: plant 790+383 / city_view 398+240 / other_building 554+259 / storefront 129 / mall 160 (raw)
100
+ (skyscraper/residential 候选枯竭未新增; fetch 去重由 jsonl 断点保证)
101
+ - 4KLSDB b2 事故: klsdb_extract 未跳过已用 row-group -> b2 与 b1 完全重复(802/802), 已删 b2 目录;
102
+ 已修复 extract 记录 used_rgs.json, 后续 b3 不再重复
103
+ - clean_url_vol.ps1 正清洗 5 类 raw(plant/city_view/other_building/storefront/mall)
104
+ - finalize_all.ps1(等清洗完): clip_url3(全URL目录) -> merge_reports(cat_pool_final, id去重)
105
+ -> split_pool(70源val) -> build_patches2(train 45k, val 560) -> build_manifest(train/val + RealSR对)
106
+ - 产出: logs/cat_pool_final.json / cat_train.json / cat_val.json / data/patches / data/patches_val /
107
+ data/manifest_train.json / data/manifest_val.json / data/README.md
README.md ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # AdcSR 竞赛冲分工程(CSIG 2026 Camera学术之星 赛道二)
2
+
3
+ 以官方 Guaishou74851/AdcSR (CVPR 2025) 为基座的分阶段训练工程:
4
+ S1 LoRA/QLoRA 像素域适配 -> S2 GDPO-SR 教师升级蒸馏 -> S3 域微调 -> 速度优化 -> 赛题提交包。
5
+
6
+ ## 目录
7
+ - `official/` 官方仓库原样 clone(只读参照;训练脚本通过 sys.path 引用其 model/dataset/utils/ram)
8
+ - `src/` 训练/蒸馏/导出/推理/评测
9
+ - `tools/` 数据下载、清洗、打标、建 patch、清单、val 划分(纯 CPU)
10
+ - `configs/` 各阶段 yaml
11
+ - `scripts/` 服务器环境、权重、运行、打包上传、归档
12
+ - `weight/ models/ logs/ data/` 运行产物(不入库)
13
+
14
+ ## 本地 vs 服务器
15
+ - 本地(RTX 3050 4G / F 盘 ~44G):代码开发、CPU 数据清洗、gt<=128 的 1 迭代冒烟。
16
+ - aspire2a(单卡 A100 40G):S1/S2/S3 全量训练、评测、导出、100 张 4K 出图。
17
+
18
+ ## 快速开始(服务器)
19
+ ```bash
20
+ bash scripts/setup_env.sh # conda + 依赖
21
+ bash scripts/download_weights.sh # Plan A: 服务器拉权重(AdcSR 5件套+SD2.1+GDPO)
22
+ python scripts/env_check.py
23
+ bash scripts/run_smoke.sh # S0 冒烟(每个训练配置 1 迭代)
24
+ bash scripts/run_stage1.sh # S1 LoRA
25
+ bash scripts/run_stage2.sh # S2 GDPO 教师蒸馏
26
+ bash scripts/run_stage3.sh # S3 域微调
27
+ python src/eval_val.py --weights <pkl> # val 评测 + 代理分
28
+ python src/export_jit.py --weights <pkl> --out model_dir
29
+ python src/inference_4k.py --lr_dir <testLR> --weights <pkl> --out output_dir
30
+ python src/check_submission.py --zip <file.zip>
31
+ ```
32
+
33
+ ## 训练/评测口径(与官方一致)
34
+ - 输入 LR 128x128 (scale=4, gt 512) -> 输出 512x512 RGB;官方 runner 用 512x512 包装/分块。
35
+ - val 代理分 proxy = 0.25*CLIPIQA + 0.2*(MUSIQ/100) + 0.15*MANIQA + 0.15*(1-NIQE/10) + 0.25*(1-LPIPS)
36
+ - 保真护栏:PSNR 不低于官方基线 0.5dB,SSIM 不低于 0.01。
RUNBOOK_A100.md ADDED
@@ -0,0 +1,51 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # CSIGv3 AdcSR 训练脚本包 (A100 runbook)
2
+
3
+ ## 0) 数据
4
+ ```bash
5
+ hf download XenderYang/CSIGv2_SR-Train --repo-type dataset --local-dir ./data
6
+ cd data && for f in *.tar; do tar -xf "$f"; done && cd ..
7
+ # 重新生成清单(manifest 内为 Windows 绝对路径, 不可直接用)
8
+ python tools/build_manifest.py --hr_dirs data/patches \
9
+ --real_lr data/real_pairs/RealSR/train_lr --real_hr data/real_pairs/RealSR/train_hr \
10
+ --out data/manifest_train.json
11
+ python tools/build_manifest.py --hr_dirs data/patches_val \
12
+ --real_lr data/real_pairs/RealSR/test_lr --real_hr data/real_pairs/RealSR/test_hr \
13
+ --out data/manifest_val.json
14
+ ```
15
+
16
+ ## 1) 环境
17
+ ```bash
18
+ conda create -n AdcSR python=3.10 -y && conda activate AdcSR
19
+ pip install -r requirements.txt
20
+ # 数据仍以本仓库 scripts/setup_env.sh 指引; 权重:
21
+ bash scripts/download_weights.sh # 官方 AdcSR 5件 + SD2.1(fp16) + GDPO(可选)
22
+ python scripts/env_check.py # GPU/驱动/关键 import/单迭代自检
23
+ ```
24
+
25
+ ## 2) S0 冒烟(每个配置 2 迭代, 必须先过)
26
+ ```bash
27
+ bash scripts/run_smoke.sh
28
+ # 或手动:
29
+ 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
30
+ 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
31
+ # GDPO 教师冒烟(先 probe)
32
+ python scripts/probe_gdpo.py weight/gdpo
33
+ 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
34
+ ```
35
+
36
+ ## 3) 正式训练
37
+ ```bash
38
+ bash scripts/run_stage1.sh # S1 LoRA 域适配 (20k step)
39
+ # eval 选 best -> weight/s1/net_params_BEST.pkl
40
+ bash scripts/run_stage2.sh gdpo # S2 GDPO 蒸馏(回退 osediff) (30k step)
41
+ bash scripts/run_stage3.sh # S3 域微调 (20k step)
42
+ ```
43
+
44
+ ## 4) 验证/导出/提交
45
+ 见 src/eval_val.py / export_jit.py / inference_4k.py / check_submission.py
46
+ (赛题验证 3 对: data/val_official; 打分测试集: 见工程 README)
47
+
48
+ ## 注意
49
+ - 训练脚本依赖 torchvision(官方 bsr/degradations.py import)。A100(Linux+torchvision)正常;
50
+ 本机 Windows 沙箱若 torchvision 报 `operator torchvision::nms does not exist`, 请在自己终端(非沙箱)运行。
51
+ - 数据清单中的真实退化对仅 RealSR 400/99; 其余母本在 data/clean_all, data/clean。
configs/config_base.yml ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # base degradation config (override dataroot_gt/gt_size/iter_num per stage)
2
+ # Real-ESRGAN style degradation (shared by all stages)
3
+ scale: 4
4
+ resize_prob: [0.2, 0.7, 0.1]
5
+ resize_range: [0.3, 1.5]
6
+ gaussian_noise_prob: 0.5
7
+ noise_range: [1, 15]
8
+ poisson_scale_range: [0.05, 2.0]
9
+ gray_noise_prob: 0.4
10
+ jpeg_range: [60, 95]
11
+ second_blur_prob: 0.5
12
+ resize_prob2: [0.3, 0.4, 0.3]
13
+ resize_range2: [0.6, 1.2]
14
+ gaussian_noise_prob2: 0.5
15
+ noise_range2: [1, 12]
16
+ poisson_scale_range2: [0.05, 1.0]
17
+ gray_noise_prob2: 0.4
18
+ jpeg_range2: [60, 100]
19
+ blur_kernel_size: 21
20
+ kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
21
+ kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
22
+ sinc_prob: 0.1
23
+ blur_sigma: [0.2, 1.5]
24
+ betag_range: [0.5, 2.0]
25
+ betap_range: [1, 1.5]
26
+ blur_kernel_size2: 11
27
+ kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
28
+ kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
29
+ sinc_prob2: 0.1
30
+ blur_sigma2: [0.2, 1.0]
31
+ betag_range2: [0.5, 2.0]
32
+ betap_range2: [1, 1.5]
33
+ final_sinc_prob: 0.8
34
+ use_hflip: True
35
+ use_rot: False
configs/config_s1_lora.yml ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ dataroot_gt: data/patches
2
+ gt_size: 512
3
+ iter_num: 1000
4
+ # Real-ESRGAN style degradation (shared by all stages)
5
+ scale: 4
6
+ resize_prob: [0.2, 0.7, 0.1]
7
+ resize_range: [0.3, 1.5]
8
+ gaussian_noise_prob: 0.5
9
+ noise_range: [1, 15]
10
+ poisson_scale_range: [0.05, 2.0]
11
+ gray_noise_prob: 0.4
12
+ jpeg_range: [60, 95]
13
+ second_blur_prob: 0.5
14
+ resize_prob2: [0.3, 0.4, 0.3]
15
+ resize_range2: [0.6, 1.2]
16
+ gaussian_noise_prob2: 0.5
17
+ noise_range2: [1, 12]
18
+ poisson_scale_range2: [0.05, 1.0]
19
+ gray_noise_prob2: 0.4
20
+ jpeg_range2: [60, 100]
21
+ blur_kernel_size: 21
22
+ kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
23
+ kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
24
+ sinc_prob: 0.1
25
+ blur_sigma: [0.2, 1.5]
26
+ betag_range: [0.5, 2.0]
27
+ betap_range: [1, 1.5]
28
+ blur_kernel_size2: 11
29
+ kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
30
+ kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
31
+ sinc_prob2: 0.1
32
+ blur_sigma2: [0.2, 1.0]
33
+ betag_range2: [0.5, 2.0]
34
+ betap_range2: [1, 1.5]
35
+ final_sinc_prob: 0.8
36
+ use_hflip: True
37
+ use_rot: False
configs/config_s2_distill.yml ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ dataroot_gt: data/patches
2
+ gt_size: 512
3
+ iter_num: 1000
4
+ # Real-ESRGAN style degradation (shared by all stages)
5
+ scale: 4
6
+ resize_prob: [0.2, 0.7, 0.1]
7
+ resize_range: [0.3, 1.5]
8
+ gaussian_noise_prob: 0.5
9
+ noise_range: [1, 15]
10
+ poisson_scale_range: [0.05, 2.0]
11
+ gray_noise_prob: 0.4
12
+ jpeg_range: [60, 95]
13
+ second_blur_prob: 0.5
14
+ resize_prob2: [0.3, 0.4, 0.3]
15
+ resize_range2: [0.6, 1.2]
16
+ gaussian_noise_prob2: 0.5
17
+ noise_range2: [1, 12]
18
+ poisson_scale_range2: [0.05, 1.0]
19
+ gray_noise_prob2: 0.4
20
+ jpeg_range2: [60, 100]
21
+ blur_kernel_size: 21
22
+ kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
23
+ kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
24
+ sinc_prob: 0.1
25
+ blur_sigma: [0.2, 1.5]
26
+ betag_range: [0.5, 2.0]
27
+ betap_range: [1, 1.5]
28
+ blur_kernel_size2: 11
29
+ kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
30
+ kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
31
+ sinc_prob2: 0.1
32
+ blur_sigma2: [0.2, 1.0]
33
+ betag_range2: [0.5, 2.0]
34
+ betap_range2: [1, 1.5]
35
+ final_sinc_prob: 0.8
36
+ use_hflip: True
37
+ use_rot: False
configs/config_s3_ft.yml ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ dataroot_gt: data/patches
2
+ gt_size: 512
3
+ iter_num: 1000
4
+ # Real-ESRGAN style degradation (shared by all stages)
5
+ scale: 4
6
+ resize_prob: [0.2, 0.7, 0.1]
7
+ resize_range: [0.3, 1.5]
8
+ gaussian_noise_prob: 0.5
9
+ noise_range: [1, 15]
10
+ poisson_scale_range: [0.05, 2.0]
11
+ gray_noise_prob: 0.4
12
+ jpeg_range: [60, 95]
13
+ second_blur_prob: 0.5
14
+ resize_prob2: [0.3, 0.4, 0.3]
15
+ resize_range2: [0.6, 1.2]
16
+ gaussian_noise_prob2: 0.5
17
+ noise_range2: [1, 12]
18
+ poisson_scale_range2: [0.05, 1.0]
19
+ gray_noise_prob2: 0.4
20
+ jpeg_range2: [60, 100]
21
+ blur_kernel_size: 21
22
+ kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
23
+ kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
24
+ sinc_prob: 0.1
25
+ blur_sigma: [0.2, 1.5]
26
+ betag_range: [0.5, 2.0]
27
+ betap_range: [1, 1.5]
28
+ blur_kernel_size2: 11
29
+ kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
30
+ kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
31
+ sinc_prob2: 0.1
32
+ blur_sigma2: [0.2, 1.0]
33
+ betag_range2: [0.5, 2.0]
34
+ betap_range2: [1, 1.5]
35
+ final_sinc_prob: 0.8
36
+ use_hflip: True
37
+ use_rot: False
configs/config_smoke.yml ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ dataroot_gt: data/patches
2
+ gt_size: 128
3
+ iter_num: 2
4
+ # Real-ESRGAN style degradation (shared by all stages)
5
+ scale: 4
6
+ resize_prob: [0.2, 0.7, 0.1]
7
+ resize_range: [0.3, 1.5]
8
+ gaussian_noise_prob: 0.5
9
+ noise_range: [1, 15]
10
+ poisson_scale_range: [0.05, 2.0]
11
+ gray_noise_prob: 0.4
12
+ jpeg_range: [60, 95]
13
+ second_blur_prob: 0.5
14
+ resize_prob2: [0.3, 0.4, 0.3]
15
+ resize_range2: [0.6, 1.2]
16
+ gaussian_noise_prob2: 0.5
17
+ noise_range2: [1, 12]
18
+ poisson_scale_range2: [0.05, 1.0]
19
+ gray_noise_prob2: 0.4
20
+ jpeg_range2: [60, 100]
21
+ blur_kernel_size: 21
22
+ kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
23
+ kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
24
+ sinc_prob: 0.1
25
+ blur_sigma: [0.2, 1.5]
26
+ betag_range: [0.5, 2.0]
27
+ betap_range: [1, 1.5]
28
+ blur_kernel_size2: 11
29
+ kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso']
30
+ kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03]
31
+ sinc_prob2: 0.1
32
+ blur_sigma2: [0.2, 1.0]
33
+ betag_range2: [0.5, 2.0]
34
+ betap_range2: [1, 1.5]
35
+ final_sinc_prob: 0.8
36
+ use_hflip: True
37
+ use_rot: False
official/bsr/__pycache__/degradations.cpython-312.pyc ADDED
Binary file (33.6 kB). View file
 
official/bsr/degradations.py ADDED
@@ -0,0 +1,764 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import math
3
+ import numpy as np
4
+ import random
5
+ import torch
6
+ from scipy import special
7
+ from scipy.stats import multivariate_normal
8
+ from torchvision.transforms._functional_tensor import rgb_to_grayscale
9
+
10
+ # -------------------------------------------------------------------- #
11
+ # --------------------------- blur kernels --------------------------- #
12
+ # -------------------------------------------------------------------- #
13
+
14
+
15
+ # --------------------------- util functions --------------------------- #
16
+ def sigma_matrix2(sig_x, sig_y, theta):
17
+ """Calculate the rotated sigma matrix (two dimensional matrix).
18
+
19
+ Args:
20
+ sig_x (float):
21
+ sig_y (float):
22
+ theta (float): Radian measurement.
23
+
24
+ Returns:
25
+ ndarray: Rotated sigma matrix.
26
+ """
27
+ d_matrix = np.array([[sig_x**2, 0], [0, sig_y**2]])
28
+ u_matrix = np.array([[np.cos(theta), -np.sin(theta)], [np.sin(theta), np.cos(theta)]])
29
+ return np.dot(u_matrix, np.dot(d_matrix, u_matrix.T))
30
+
31
+
32
+ def mesh_grid(kernel_size):
33
+ """Generate the mesh grid, centering at zero.
34
+
35
+ Args:
36
+ kernel_size (int):
37
+
38
+ Returns:
39
+ xy (ndarray): with the shape (kernel_size, kernel_size, 2)
40
+ xx (ndarray): with the shape (kernel_size, kernel_size)
41
+ yy (ndarray): with the shape (kernel_size, kernel_size)
42
+ """
43
+ ax = np.arange(-kernel_size // 2 + 1., kernel_size // 2 + 1.)
44
+ xx, yy = np.meshgrid(ax, ax)
45
+ xy = np.hstack((xx.reshape((kernel_size * kernel_size, 1)), yy.reshape(kernel_size * kernel_size,
46
+ 1))).reshape(kernel_size, kernel_size, 2)
47
+ return xy, xx, yy
48
+
49
+
50
+ def pdf2(sigma_matrix, grid):
51
+ """Calculate PDF of the bivariate Gaussian distribution.
52
+
53
+ Args:
54
+ sigma_matrix (ndarray): with the shape (2, 2)
55
+ grid (ndarray): generated by :func:`mesh_grid`,
56
+ with the shape (K, K, 2), K is the kernel size.
57
+
58
+ Returns:
59
+ kernel (ndarrray): un-normalized kernel.
60
+ """
61
+ inverse_sigma = np.linalg.inv(sigma_matrix)
62
+ kernel = np.exp(-0.5 * np.sum(np.dot(grid, inverse_sigma) * grid, 2))
63
+ return kernel
64
+
65
+
66
+ def cdf2(d_matrix, grid):
67
+ """Calculate the CDF of the standard bivariate Gaussian distribution.
68
+ Used in skewed Gaussian distribution.
69
+
70
+ Args:
71
+ d_matrix (ndarrasy): skew matrix.
72
+ grid (ndarray): generated by :func:`mesh_grid`,
73
+ with the shape (K, K, 2), K is the kernel size.
74
+
75
+ Returns:
76
+ cdf (ndarray): skewed cdf.
77
+ """
78
+ rv = multivariate_normal([0, 0], [[1, 0], [0, 1]])
79
+ grid = np.dot(grid, d_matrix)
80
+ cdf = rv.cdf(grid)
81
+ return cdf
82
+
83
+
84
+ def bivariate_Gaussian(kernel_size, sig_x, sig_y, theta, grid=None, isotropic=True):
85
+ """Generate a bivariate isotropic or anisotropic Gaussian kernel.
86
+
87
+ In the isotropic mode, only `sig_x` is used. `sig_y` and `theta` is ignored.
88
+
89
+ Args:
90
+ kernel_size (int):
91
+ sig_x (float):
92
+ sig_y (float):
93
+ theta (float): Radian measurement.
94
+ grid (ndarray, optional): generated by :func:`mesh_grid`,
95
+ with the shape (K, K, 2), K is the kernel size. Default: None
96
+ isotropic (bool):
97
+
98
+ Returns:
99
+ kernel (ndarray): normalized kernel.
100
+ """
101
+ if grid is None:
102
+ grid, _, _ = mesh_grid(kernel_size)
103
+ if isotropic:
104
+ sigma_matrix = np.array([[sig_x**2, 0], [0, sig_x**2]])
105
+ else:
106
+ sigma_matrix = sigma_matrix2(sig_x, sig_y, theta)
107
+ kernel = pdf2(sigma_matrix, grid)
108
+ kernel = kernel / np.sum(kernel)
109
+ return kernel
110
+
111
+
112
+ def bivariate_generalized_Gaussian(kernel_size, sig_x, sig_y, theta, beta, grid=None, isotropic=True):
113
+ """Generate a bivariate generalized Gaussian kernel.
114
+
115
+ ``Paper: Parameter Estimation For Multivariate Generalized Gaussian Distributions``
116
+
117
+ In the isotropic mode, only `sig_x` is used. `sig_y` and `theta` is ignored.
118
+
119
+ Args:
120
+ kernel_size (int):
121
+ sig_x (float):
122
+ sig_y (float):
123
+ theta (float): Radian measurement.
124
+ beta (float): shape parameter, beta = 1 is the normal distribution.
125
+ grid (ndarray, optional): generated by :func:`mesh_grid`,
126
+ with the shape (K, K, 2), K is the kernel size. Default: None
127
+
128
+ Returns:
129
+ kernel (ndarray): normalized kernel.
130
+ """
131
+ if grid is None:
132
+ grid, _, _ = mesh_grid(kernel_size)
133
+ if isotropic:
134
+ sigma_matrix = np.array([[sig_x**2, 0], [0, sig_x**2]])
135
+ else:
136
+ sigma_matrix = sigma_matrix2(sig_x, sig_y, theta)
137
+ inverse_sigma = np.linalg.inv(sigma_matrix)
138
+ kernel = np.exp(-0.5 * np.power(np.sum(np.dot(grid, inverse_sigma) * grid, 2), beta))
139
+ kernel = kernel / np.sum(kernel)
140
+ return kernel
141
+
142
+
143
+ def bivariate_plateau(kernel_size, sig_x, sig_y, theta, beta, grid=None, isotropic=True):
144
+ """Generate a plateau-like anisotropic kernel.
145
+
146
+ 1 / (1+x^(beta))
147
+
148
+ Reference: https://stats.stackexchange.com/questions/203629/is-there-a-plateau-shaped-distribution
149
+
150
+ In the isotropic mode, only `sig_x` is used. `sig_y` and `theta` is ignored.
151
+
152
+ Args:
153
+ kernel_size (int):
154
+ sig_x (float):
155
+ sig_y (float):
156
+ theta (float): Radian measurement.
157
+ beta (float): shape parameter, beta = 1 is the normal distribution.
158
+ grid (ndarray, optional): generated by :func:`mesh_grid`,
159
+ with the shape (K, K, 2), K is the kernel size. Default: None
160
+
161
+ Returns:
162
+ kernel (ndarray): normalized kernel.
163
+ """
164
+ if grid is None:
165
+ grid, _, _ = mesh_grid(kernel_size)
166
+ if isotropic:
167
+ sigma_matrix = np.array([[sig_x**2, 0], [0, sig_x**2]])
168
+ else:
169
+ sigma_matrix = sigma_matrix2(sig_x, sig_y, theta)
170
+ inverse_sigma = np.linalg.inv(sigma_matrix)
171
+ kernel = np.reciprocal(np.power(np.sum(np.dot(grid, inverse_sigma) * grid, 2), beta) + 1)
172
+ kernel = kernel / np.sum(kernel)
173
+ return kernel
174
+
175
+
176
+ def random_bivariate_Gaussian(kernel_size,
177
+ sigma_x_range,
178
+ sigma_y_range,
179
+ rotation_range,
180
+ noise_range=None,
181
+ isotropic=True):
182
+ """Randomly generate bivariate isotropic or anisotropic Gaussian kernels.
183
+
184
+ In the isotropic mode, only `sigma_x_range` is used. `sigma_y_range` and `rotation_range` is ignored.
185
+
186
+ Args:
187
+ kernel_size (int):
188
+ sigma_x_range (tuple): [0.6, 5]
189
+ sigma_y_range (tuple): [0.6, 5]
190
+ rotation range (tuple): [-math.pi, math.pi]
191
+ noise_range(tuple, optional): multiplicative kernel noise,
192
+ [0.75, 1.25]. Default: None
193
+
194
+ Returns:
195
+ kernel (ndarray):
196
+ """
197
+ assert kernel_size % 2 == 1, 'Kernel size must be an odd number.'
198
+ assert sigma_x_range[0] < sigma_x_range[1], 'Wrong sigma_x_range.'
199
+ sigma_x = np.random.uniform(sigma_x_range[0], sigma_x_range[1])
200
+ if isotropic is False:
201
+ assert sigma_y_range[0] < sigma_y_range[1], 'Wrong sigma_y_range.'
202
+ assert rotation_range[0] < rotation_range[1], 'Wrong rotation_range.'
203
+ sigma_y = np.random.uniform(sigma_y_range[0], sigma_y_range[1])
204
+ rotation = np.random.uniform(rotation_range[0], rotation_range[1])
205
+ else:
206
+ sigma_y = sigma_x
207
+ rotation = 0
208
+
209
+ kernel = bivariate_Gaussian(kernel_size, sigma_x, sigma_y, rotation, isotropic=isotropic)
210
+
211
+ # add multiplicative noise
212
+ if noise_range is not None:
213
+ assert noise_range[0] < noise_range[1], 'Wrong noise range.'
214
+ noise = np.random.uniform(noise_range[0], noise_range[1], size=kernel.shape)
215
+ kernel = kernel * noise
216
+ kernel = kernel / np.sum(kernel)
217
+ return kernel
218
+
219
+
220
+ def random_bivariate_generalized_Gaussian(kernel_size,
221
+ sigma_x_range,
222
+ sigma_y_range,
223
+ rotation_range,
224
+ beta_range,
225
+ noise_range=None,
226
+ isotropic=True):
227
+ """Randomly generate bivariate generalized Gaussian kernels.
228
+
229
+ In the isotropic mode, only `sigma_x_range` is used. `sigma_y_range` and `rotation_range` is ignored.
230
+
231
+ Args:
232
+ kernel_size (int):
233
+ sigma_x_range (tuple): [0.6, 5]
234
+ sigma_y_range (tuple): [0.6, 5]
235
+ rotation range (tuple): [-math.pi, math.pi]
236
+ beta_range (tuple): [0.5, 8]
237
+ noise_range(tuple, optional): multiplicative kernel noise,
238
+ [0.75, 1.25]. Default: None
239
+
240
+ Returns:
241
+ kernel (ndarray):
242
+ """
243
+ assert kernel_size % 2 == 1, 'Kernel size must be an odd number.'
244
+ assert sigma_x_range[0] < sigma_x_range[1], 'Wrong sigma_x_range.'
245
+ sigma_x = np.random.uniform(sigma_x_range[0], sigma_x_range[1])
246
+ if isotropic is False:
247
+ assert sigma_y_range[0] < sigma_y_range[1], 'Wrong sigma_y_range.'
248
+ assert rotation_range[0] < rotation_range[1], 'Wrong rotation_range.'
249
+ sigma_y = np.random.uniform(sigma_y_range[0], sigma_y_range[1])
250
+ rotation = np.random.uniform(rotation_range[0], rotation_range[1])
251
+ else:
252
+ sigma_y = sigma_x
253
+ rotation = 0
254
+
255
+ # assume beta_range[0] < 1 < beta_range[1]
256
+ if np.random.uniform() < 0.5:
257
+ beta = np.random.uniform(beta_range[0], 1)
258
+ else:
259
+ beta = np.random.uniform(1, beta_range[1])
260
+
261
+ kernel = bivariate_generalized_Gaussian(kernel_size, sigma_x, sigma_y, rotation, beta, isotropic=isotropic)
262
+
263
+ # add multiplicative noise
264
+ if noise_range is not None:
265
+ assert noise_range[0] < noise_range[1], 'Wrong noise range.'
266
+ noise = np.random.uniform(noise_range[0], noise_range[1], size=kernel.shape)
267
+ kernel = kernel * noise
268
+ kernel = kernel / np.sum(kernel)
269
+ return kernel
270
+
271
+
272
+ def random_bivariate_plateau(kernel_size,
273
+ sigma_x_range,
274
+ sigma_y_range,
275
+ rotation_range,
276
+ beta_range,
277
+ noise_range=None,
278
+ isotropic=True):
279
+ """Randomly generate bivariate plateau kernels.
280
+
281
+ In the isotropic mode, only `sigma_x_range` is used. `sigma_y_range` and `rotation_range` is ignored.
282
+
283
+ Args:
284
+ kernel_size (int):
285
+ sigma_x_range (tuple): [0.6, 5]
286
+ sigma_y_range (tuple): [0.6, 5]
287
+ rotation range (tuple): [-math.pi/2, math.pi/2]
288
+ beta_range (tuple): [1, 4]
289
+ noise_range(tuple, optional): multiplicative kernel noise,
290
+ [0.75, 1.25]. Default: None
291
+
292
+ Returns:
293
+ kernel (ndarray):
294
+ """
295
+ assert kernel_size % 2 == 1, 'Kernel size must be an odd number.'
296
+ assert sigma_x_range[0] < sigma_x_range[1], 'Wrong sigma_x_range.'
297
+ sigma_x = np.random.uniform(sigma_x_range[0], sigma_x_range[1])
298
+ if isotropic is False:
299
+ assert sigma_y_range[0] < sigma_y_range[1], 'Wrong sigma_y_range.'
300
+ assert rotation_range[0] < rotation_range[1], 'Wrong rotation_range.'
301
+ sigma_y = np.random.uniform(sigma_y_range[0], sigma_y_range[1])
302
+ rotation = np.random.uniform(rotation_range[0], rotation_range[1])
303
+ else:
304
+ sigma_y = sigma_x
305
+ rotation = 0
306
+
307
+ # TODO: this may be not proper
308
+ if np.random.uniform() < 0.5:
309
+ beta = np.random.uniform(beta_range[0], 1)
310
+ else:
311
+ beta = np.random.uniform(1, beta_range[1])
312
+
313
+ kernel = bivariate_plateau(kernel_size, sigma_x, sigma_y, rotation, beta, isotropic=isotropic)
314
+ # add multiplicative noise
315
+ if noise_range is not None:
316
+ assert noise_range[0] < noise_range[1], 'Wrong noise range.'
317
+ noise = np.random.uniform(noise_range[0], noise_range[1], size=kernel.shape)
318
+ kernel = kernel * noise
319
+ kernel = kernel / np.sum(kernel)
320
+
321
+ return kernel
322
+
323
+
324
+ def random_mixed_kernels(kernel_list,
325
+ kernel_prob,
326
+ kernel_size=21,
327
+ sigma_x_range=(0.6, 5),
328
+ sigma_y_range=(0.6, 5),
329
+ rotation_range=(-math.pi, math.pi),
330
+ betag_range=(0.5, 8),
331
+ betap_range=(0.5, 8),
332
+ noise_range=None):
333
+ """Randomly generate mixed kernels.
334
+
335
+ Args:
336
+ kernel_list (tuple): a list name of kernel types,
337
+ support ['iso', 'aniso', 'skew', 'generalized', 'plateau_iso',
338
+ 'plateau_aniso']
339
+ kernel_prob (tuple): corresponding kernel probability for each
340
+ kernel type
341
+ kernel_size (int):
342
+ sigma_x_range (tuple): [0.6, 5]
343
+ sigma_y_range (tuple): [0.6, 5]
344
+ rotation range (tuple): [-math.pi, math.pi]
345
+ beta_range (tuple): [0.5, 8]
346
+ noise_range(tuple, optional): multiplicative kernel noise,
347
+ [0.75, 1.25]. Default: None
348
+
349
+ Returns:
350
+ kernel (ndarray):
351
+ """
352
+ kernel_type = random.choices(kernel_list, kernel_prob)[0]
353
+ if kernel_type == 'iso':
354
+ kernel = random_bivariate_Gaussian(
355
+ kernel_size, sigma_x_range, sigma_y_range, rotation_range, noise_range=noise_range, isotropic=True)
356
+ elif kernel_type == 'aniso':
357
+ kernel = random_bivariate_Gaussian(
358
+ kernel_size, sigma_x_range, sigma_y_range, rotation_range, noise_range=noise_range, isotropic=False)
359
+ elif kernel_type == 'generalized_iso':
360
+ kernel = random_bivariate_generalized_Gaussian(
361
+ kernel_size,
362
+ sigma_x_range,
363
+ sigma_y_range,
364
+ rotation_range,
365
+ betag_range,
366
+ noise_range=noise_range,
367
+ isotropic=True)
368
+ elif kernel_type == 'generalized_aniso':
369
+ kernel = random_bivariate_generalized_Gaussian(
370
+ kernel_size,
371
+ sigma_x_range,
372
+ sigma_y_range,
373
+ rotation_range,
374
+ betag_range,
375
+ noise_range=noise_range,
376
+ isotropic=False)
377
+ elif kernel_type == 'plateau_iso':
378
+ kernel = random_bivariate_plateau(
379
+ kernel_size, sigma_x_range, sigma_y_range, rotation_range, betap_range, noise_range=None, isotropic=True)
380
+ elif kernel_type == 'plateau_aniso':
381
+ kernel = random_bivariate_plateau(
382
+ kernel_size, sigma_x_range, sigma_y_range, rotation_range, betap_range, noise_range=None, isotropic=False)
383
+ return kernel
384
+
385
+
386
+ np.seterr(divide='ignore', invalid='ignore')
387
+
388
+
389
+ def circular_lowpass_kernel(cutoff, kernel_size, pad_to=0):
390
+ """2D sinc filter
391
+
392
+ Reference: https://dsp.stackexchange.com/questions/58301/2-d-circularly-symmetric-low-pass-filter
393
+
394
+ Args:
395
+ cutoff (float): cutoff frequency in radians (pi is max)
396
+ kernel_size (int): horizontal and vertical size, must be odd.
397
+ pad_to (int): pad kernel size to desired size, must be odd or zero.
398
+ """
399
+ assert kernel_size % 2 == 1, 'Kernel size must be an odd number.'
400
+ kernel = np.fromfunction(
401
+ lambda x, y: cutoff * special.j1(cutoff * np.sqrt(
402
+ (x - (kernel_size - 1) / 2)**2 + (y - (kernel_size - 1) / 2)**2)) / (2 * np.pi * np.sqrt(
403
+ (x - (kernel_size - 1) / 2)**2 + (y - (kernel_size - 1) / 2)**2)), [kernel_size, kernel_size])
404
+ kernel[(kernel_size - 1) // 2, (kernel_size - 1) // 2] = cutoff**2 / (4 * np.pi)
405
+ kernel = kernel / np.sum(kernel)
406
+ if pad_to > kernel_size:
407
+ pad_size = (pad_to - kernel_size) // 2
408
+ kernel = np.pad(kernel, ((pad_size, pad_size), (pad_size, pad_size)))
409
+ return kernel
410
+
411
+
412
+ # ------------------------------------------------------------- #
413
+ # --------------------------- noise --------------------------- #
414
+ # ------------------------------------------------------------- #
415
+
416
+ # ----------------------- Gaussian Noise ----------------------- #
417
+
418
+
419
+ def generate_gaussian_noise(img, sigma=10, gray_noise=False):
420
+ """Generate Gaussian noise.
421
+
422
+ Args:
423
+ img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32.
424
+ sigma (float): Noise scale (measured in range 255). Default: 10.
425
+
426
+ Returns:
427
+ (Numpy array): Returned noisy image, shape (h, w, c), range[0, 1],
428
+ float32.
429
+ """
430
+ if gray_noise:
431
+ noise = np.float32(np.random.randn(*(img.shape[0:2]))) * sigma / 255.
432
+ noise = np.expand_dims(noise, axis=2).repeat(3, axis=2)
433
+ else:
434
+ noise = np.float32(np.random.randn(*(img.shape))) * sigma / 255.
435
+ return noise
436
+
437
+
438
+ def add_gaussian_noise(img, sigma=10, clip=True, rounds=False, gray_noise=False):
439
+ """Add Gaussian noise.
440
+
441
+ Args:
442
+ img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32.
443
+ sigma (float): Noise scale (measured in range 255). Default: 10.
444
+
445
+ Returns:
446
+ (Numpy array): Returned noisy image, shape (h, w, c), range[0, 1],
447
+ float32.
448
+ """
449
+ noise = generate_gaussian_noise(img, sigma, gray_noise)
450
+ out = img + noise
451
+ if clip and rounds:
452
+ out = np.clip((out * 255.0).round(), 0, 255) / 255.
453
+ elif clip:
454
+ out = np.clip(out, 0, 1)
455
+ elif rounds:
456
+ out = (out * 255.0).round() / 255.
457
+ return out
458
+
459
+
460
+ def generate_gaussian_noise_pt(img, sigma=10, gray_noise=0):
461
+ """Add Gaussian noise (PyTorch version).
462
+
463
+ Args:
464
+ img (Tensor): Shape (b, c, h, w), range[0, 1], float32.
465
+ scale (float | Tensor): Noise scale. Default: 1.0.
466
+
467
+ Returns:
468
+ (Tensor): Returned noisy image, shape (b, c, h, w), range[0, 1],
469
+ float32.
470
+ """
471
+ b, _, h, w = img.size()
472
+ if not isinstance(sigma, (float, int)):
473
+ sigma = sigma.view(img.size(0), 1, 1, 1)
474
+ if isinstance(gray_noise, (float, int)):
475
+ cal_gray_noise = gray_noise > 0
476
+ else:
477
+ gray_noise = gray_noise.view(b, 1, 1, 1)
478
+ cal_gray_noise = torch.sum(gray_noise) > 0
479
+
480
+ if cal_gray_noise:
481
+ noise_gray = torch.randn(*img.size()[2:4], dtype=img.dtype, device=img.device) * sigma / 255.
482
+ noise_gray = noise_gray.view(b, 1, h, w)
483
+
484
+ # always calculate color noise
485
+ noise = torch.randn(*img.size(), dtype=img.dtype, device=img.device) * sigma / 255.
486
+
487
+ if cal_gray_noise:
488
+ noise = noise * (1 - gray_noise) + noise_gray * gray_noise
489
+ return noise
490
+
491
+
492
+ def add_gaussian_noise_pt(img, sigma=10, gray_noise=0, clip=True, rounds=False):
493
+ """Add Gaussian noise (PyTorch version).
494
+
495
+ Args:
496
+ img (Tensor): Shape (b, c, h, w), range[0, 1], float32.
497
+ scale (float | Tensor): Noise scale. Default: 1.0.
498
+
499
+ Returns:
500
+ (Tensor): Returned noisy image, shape (b, c, h, w), range[0, 1],
501
+ float32.
502
+ """
503
+ noise = generate_gaussian_noise_pt(img, sigma, gray_noise)
504
+ out = img + noise
505
+ if clip and rounds:
506
+ out = torch.clamp((out * 255.0).round(), 0, 255) / 255.
507
+ elif clip:
508
+ out = torch.clamp(out, 0, 1)
509
+ elif rounds:
510
+ out = (out * 255.0).round() / 255.
511
+ return out
512
+
513
+
514
+ # ----------------------- Random Gaussian Noise ----------------------- #
515
+ def random_generate_gaussian_noise(img, sigma_range=(0, 10), gray_prob=0):
516
+ sigma = np.random.uniform(sigma_range[0], sigma_range[1])
517
+ if np.random.uniform() < gray_prob:
518
+ gray_noise = True
519
+ else:
520
+ gray_noise = False
521
+ return generate_gaussian_noise(img, sigma, gray_noise)
522
+
523
+
524
+ def random_add_gaussian_noise(img, sigma_range=(0, 1.0), gray_prob=0, clip=True, rounds=False):
525
+ noise = random_generate_gaussian_noise(img, sigma_range, gray_prob)
526
+ out = img + noise
527
+ if clip and rounds:
528
+ out = np.clip((out * 255.0).round(), 0, 255) / 255.
529
+ elif clip:
530
+ out = np.clip(out, 0, 1)
531
+ elif rounds:
532
+ out = (out * 255.0).round() / 255.
533
+ return out
534
+
535
+
536
+ def random_generate_gaussian_noise_pt(img, sigma_range=(0, 10), gray_prob=0):
537
+ sigma = torch.rand(
538
+ img.size(0), dtype=img.dtype, device=img.device) * (sigma_range[1] - sigma_range[0]) + sigma_range[0]
539
+ gray_noise = torch.rand(img.size(0), dtype=img.dtype, device=img.device)
540
+ gray_noise = (gray_noise < gray_prob).float()
541
+ return generate_gaussian_noise_pt(img, sigma, gray_noise)
542
+
543
+
544
+ def random_add_gaussian_noise_pt(img, sigma_range=(0, 1.0), gray_prob=0, clip=True, rounds=False):
545
+ noise = random_generate_gaussian_noise_pt(img, sigma_range, gray_prob)
546
+ out = img + noise
547
+ if clip and rounds:
548
+ out = torch.clamp((out * 255.0).round(), 0, 255) / 255.
549
+ elif clip:
550
+ out = torch.clamp(out, 0, 1)
551
+ elif rounds:
552
+ out = (out * 255.0).round() / 255.
553
+ return out
554
+
555
+
556
+ # ----------------------- Poisson (Shot) Noise ----------------------- #
557
+
558
+
559
+ def generate_poisson_noise(img, scale=1.0, gray_noise=False):
560
+ """Generate poisson noise.
561
+
562
+ Reference: https://github.com/scikit-image/scikit-image/blob/main/skimage/util/noise.py#L37-L219
563
+
564
+ Args:
565
+ img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32.
566
+ scale (float): Noise scale. Default: 1.0.
567
+ gray_noise (bool): Whether generate gray noise. Default: False.
568
+
569
+ Returns:
570
+ (Numpy array): Returned noisy image, shape (h, w, c), range[0, 1],
571
+ float32.
572
+ """
573
+ if gray_noise:
574
+ img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
575
+ # round and clip image for counting vals correctly
576
+ img = np.clip((img * 255.0).round(), 0, 255) / 255.
577
+ vals = len(np.unique(img))
578
+ vals = 2**np.ceil(np.log2(vals))
579
+ out = np.float32(np.random.poisson(img * vals) / float(vals))
580
+ noise = out - img
581
+ if gray_noise:
582
+ noise = np.repeat(noise[:, :, np.newaxis], 3, axis=2)
583
+ return noise * scale
584
+
585
+
586
+ def add_poisson_noise(img, scale=1.0, clip=True, rounds=False, gray_noise=False):
587
+ """Add poisson noise.
588
+
589
+ Args:
590
+ img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32.
591
+ scale (float): Noise scale. Default: 1.0.
592
+ gray_noise (bool): Whether generate gray noise. Default: False.
593
+
594
+ Returns:
595
+ (Numpy array): Returned noisy image, shape (h, w, c), range[0, 1],
596
+ float32.
597
+ """
598
+ noise = generate_poisson_noise(img, scale, gray_noise)
599
+ out = img + noise
600
+ if clip and rounds:
601
+ out = np.clip((out * 255.0).round(), 0, 255) / 255.
602
+ elif clip:
603
+ out = np.clip(out, 0, 1)
604
+ elif rounds:
605
+ out = (out * 255.0).round() / 255.
606
+ return out
607
+
608
+
609
+ def generate_poisson_noise_pt(img, scale=1.0, gray_noise=0):
610
+ """Generate a batch of poisson noise (PyTorch version)
611
+
612
+ Args:
613
+ img (Tensor): Input image, shape (b, c, h, w), range [0, 1], float32.
614
+ scale (float | Tensor): Noise scale. Number or Tensor with shape (b).
615
+ Default: 1.0.
616
+ gray_noise (float | Tensor): 0-1 number or Tensor with shape (b).
617
+ 0 for False, 1 for True. Default: 0.
618
+
619
+ Returns:
620
+ (Tensor): Returned noisy image, shape (b, c, h, w), range[0, 1],
621
+ float32.
622
+ """
623
+ b, _, h, w = img.size()
624
+ if isinstance(gray_noise, (float, int)):
625
+ cal_gray_noise = gray_noise > 0
626
+ else:
627
+ gray_noise = gray_noise.view(b, 1, 1, 1)
628
+ cal_gray_noise = torch.sum(gray_noise) > 0
629
+ if cal_gray_noise:
630
+ img_gray = rgb_to_grayscale(img, num_output_channels=1)
631
+ # round and clip image for counting vals correctly
632
+ img_gray = torch.clamp((img_gray * 255.0).round(), 0, 255) / 255.
633
+ # use for-loop to get the unique values for each sample
634
+ vals_list = [len(torch.unique(img_gray[i, :, :, :])) for i in range(b)]
635
+ vals_list = [2**np.ceil(np.log2(vals)) for vals in vals_list]
636
+ vals = img_gray.new_tensor(vals_list).view(b, 1, 1, 1)
637
+ out = torch.poisson(img_gray * vals) / vals
638
+ noise_gray = out - img_gray
639
+ noise_gray = noise_gray.expand(b, 3, h, w)
640
+
641
+ # always calculate color noise
642
+ # round and clip image for counting vals correctly
643
+ img = torch.clamp((img * 255.0).round(), 0, 255) / 255.
644
+ # use for-loop to get the unique values for each sample
645
+ vals_list = [len(torch.unique(img[i, :, :, :])) for i in range(b)]
646
+ vals_list = [2**np.ceil(np.log2(vals)) for vals in vals_list]
647
+ vals = img.new_tensor(vals_list).view(b, 1, 1, 1)
648
+ out = torch.poisson(img * vals) / vals
649
+ noise = out - img
650
+ if cal_gray_noise:
651
+ noise = noise * (1 - gray_noise) + noise_gray * gray_noise
652
+ if not isinstance(scale, (float, int)):
653
+ scale = scale.view(b, 1, 1, 1)
654
+ return noise * scale
655
+
656
+
657
+ def add_poisson_noise_pt(img, scale=1.0, clip=True, rounds=False, gray_noise=0):
658
+ """Add poisson noise to a batch of images (PyTorch version).
659
+
660
+ Args:
661
+ img (Tensor): Input image, shape (b, c, h, w), range [0, 1], float32.
662
+ scale (float | Tensor): Noise scale. Number or Tensor with shape (b).
663
+ Default: 1.0.
664
+ gray_noise (float | Tensor): 0-1 number or Tensor with shape (b).
665
+ 0 for False, 1 for True. Default: 0.
666
+
667
+ Returns:
668
+ (Tensor): Returned noisy image, shape (b, c, h, w), range[0, 1],
669
+ float32.
670
+ """
671
+ noise = generate_poisson_noise_pt(img, scale, gray_noise)
672
+ out = img + noise
673
+ if clip and rounds:
674
+ out = torch.clamp((out * 255.0).round(), 0, 255) / 255.
675
+ elif clip:
676
+ out = torch.clamp(out, 0, 1)
677
+ elif rounds:
678
+ out = (out * 255.0).round() / 255.
679
+ return out
680
+
681
+
682
+ # ----------------------- Random Poisson (Shot) Noise ----------------------- #
683
+
684
+
685
+ def random_generate_poisson_noise(img, scale_range=(0, 1.0), gray_prob=0):
686
+ scale = np.random.uniform(scale_range[0], scale_range[1])
687
+ if np.random.uniform() < gray_prob:
688
+ gray_noise = True
689
+ else:
690
+ gray_noise = False
691
+ return generate_poisson_noise(img, scale, gray_noise)
692
+
693
+
694
+ def random_add_poisson_noise(img, scale_range=(0, 1.0), gray_prob=0, clip=True, rounds=False):
695
+ noise = random_generate_poisson_noise(img, scale_range, gray_prob)
696
+ out = img + noise
697
+ if clip and rounds:
698
+ out = np.clip((out * 255.0).round(), 0, 255) / 255.
699
+ elif clip:
700
+ out = np.clip(out, 0, 1)
701
+ elif rounds:
702
+ out = (out * 255.0).round() / 255.
703
+ return out
704
+
705
+
706
+ def random_generate_poisson_noise_pt(img, scale_range=(0, 1.0), gray_prob=0):
707
+ scale = torch.rand(
708
+ img.size(0), dtype=img.dtype, device=img.device) * (scale_range[1] - scale_range[0]) + scale_range[0]
709
+ gray_noise = torch.rand(img.size(0), dtype=img.dtype, device=img.device)
710
+ gray_noise = (gray_noise < gray_prob).float()
711
+ return generate_poisson_noise_pt(img, scale, gray_noise)
712
+
713
+
714
+ def random_add_poisson_noise_pt(img, scale_range=(0, 1.0), gray_prob=0, clip=True, rounds=False):
715
+ noise = random_generate_poisson_noise_pt(img, scale_range, gray_prob)
716
+ out = img + noise
717
+ if clip and rounds:
718
+ out = torch.clamp((out * 255.0).round(), 0, 255) / 255.
719
+ elif clip:
720
+ out = torch.clamp(out, 0, 1)
721
+ elif rounds:
722
+ out = (out * 255.0).round() / 255.
723
+ return out
724
+
725
+
726
+ # ------------------------------------------------------------------------ #
727
+ # --------------------------- JPEG compression --------------------------- #
728
+ # ------------------------------------------------------------------------ #
729
+
730
+
731
+ def add_jpg_compression(img, quality=90):
732
+ """Add JPG compression artifacts.
733
+
734
+ Args:
735
+ img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32.
736
+ quality (float): JPG compression quality. 0 for lowest quality, 100 for
737
+ best quality. Default: 90.
738
+
739
+ Returns:
740
+ (Numpy array): Returned image after JPG, shape (h, w, c), range[0, 1],
741
+ float32.
742
+ """
743
+ img = np.clip(img, 0, 1)
744
+ encode_param = [int(cv2.IMWRITE_JPEG_QUALITY), quality]
745
+ _, encimg = cv2.imencode('.jpg', img * 255., encode_param)
746
+ img = np.float32(cv2.imdecode(encimg, 1)) / 255.
747
+ return img
748
+
749
+
750
+ def random_add_jpg_compression(img, quality_range=(90, 100)):
751
+ """Randomly add JPG compression artifacts.
752
+
753
+ Args:
754
+ img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32.
755
+ quality_range (tuple[float] | list[float]): JPG compression quality
756
+ range. 0 for lowest quality, 100 for best quality.
757
+ Default: (90, 100).
758
+
759
+ Returns:
760
+ (Numpy array): Returned image after JPG, shape (h, w, c), range[0, 1],
761
+ float32.
762
+ """
763
+ quality = np.random.uniform(quality_range[0], quality_range[1])
764
+ return add_jpg_compression(img, quality)
official/bsr/transforms.py ADDED
@@ -0,0 +1,179 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import random
3
+ import torch
4
+
5
+
6
+ def mod_crop(img, scale):
7
+ """Mod crop images, used during testing.
8
+
9
+ Args:
10
+ img (ndarray): Input image.
11
+ scale (int): Scale factor.
12
+
13
+ Returns:
14
+ ndarray: Result image.
15
+ """
16
+ img = img.copy()
17
+ if img.ndim in (2, 3):
18
+ h, w = img.shape[0], img.shape[1]
19
+ h_remainder, w_remainder = h % scale, w % scale
20
+ img = img[:h - h_remainder, :w - w_remainder, ...]
21
+ else:
22
+ raise ValueError(f'Wrong img ndim: {img.ndim}.')
23
+ return img
24
+
25
+
26
+ def paired_random_crop(img_gts, img_lqs, gt_patch_size, scale, gt_path=None):
27
+ """Paired random crop. Support Numpy array and Tensor inputs.
28
+
29
+ It crops lists of lq and gt images with corresponding locations.
30
+
31
+ Args:
32
+ img_gts (list[ndarray] | ndarray | list[Tensor] | Tensor): GT images. Note that all images
33
+ should have the same shape. If the input is an ndarray, it will
34
+ be transformed to a list containing itself.
35
+ img_lqs (list[ndarray] | ndarray): LQ images. Note that all images
36
+ should have the same shape. If the input is an ndarray, it will
37
+ be transformed to a list containing itself.
38
+ gt_patch_size (int): GT patch size.
39
+ scale (int): Scale factor.
40
+ gt_path (str): Path to ground-truth. Default: None.
41
+
42
+ Returns:
43
+ list[ndarray] | ndarray: GT images and LQ images. If returned results
44
+ only have one element, just return ndarray.
45
+ """
46
+
47
+ if not isinstance(img_gts, list):
48
+ img_gts = [img_gts]
49
+ if not isinstance(img_lqs, list):
50
+ img_lqs = [img_lqs]
51
+
52
+ # determine input type: Numpy array or Tensor
53
+ input_type = 'Tensor' if torch.is_tensor(img_gts[0]) else 'Numpy'
54
+
55
+ if input_type == 'Tensor':
56
+ h_lq, w_lq = img_lqs[0].size()[-2:]
57
+ h_gt, w_gt = img_gts[0].size()[-2:]
58
+ else:
59
+ h_lq, w_lq = img_lqs[0].shape[0:2]
60
+ h_gt, w_gt = img_gts[0].shape[0:2]
61
+ lq_patch_size = gt_patch_size // scale
62
+
63
+ if h_gt != h_lq * scale or w_gt != w_lq * scale:
64
+ raise ValueError(f'Scale mismatches. GT ({h_gt}, {w_gt}) is not {scale}x ',
65
+ f'multiplication of LQ ({h_lq}, {w_lq}).')
66
+ if h_lq < lq_patch_size or w_lq < lq_patch_size:
67
+ raise ValueError(f'LQ ({h_lq}, {w_lq}) is smaller than patch size '
68
+ f'({lq_patch_size}, {lq_patch_size}). '
69
+ f'Please remove {gt_path}.')
70
+
71
+ # randomly choose top and left coordinates for lq patch
72
+ top = random.randint(0, h_lq - lq_patch_size)
73
+ left = random.randint(0, w_lq - lq_patch_size)
74
+
75
+ # crop lq patch
76
+ if input_type == 'Tensor':
77
+ img_lqs = [v[:, :, top:top + lq_patch_size, left:left + lq_patch_size] for v in img_lqs]
78
+ else:
79
+ img_lqs = [v[top:top + lq_patch_size, left:left + lq_patch_size, ...] for v in img_lqs]
80
+
81
+ # crop corresponding gt patch
82
+ top_gt, left_gt = int(top * scale), int(left * scale)
83
+ if input_type == 'Tensor':
84
+ img_gts = [v[:, :, top_gt:top_gt + gt_patch_size, left_gt:left_gt + gt_patch_size] for v in img_gts]
85
+ else:
86
+ img_gts = [v[top_gt:top_gt + gt_patch_size, left_gt:left_gt + gt_patch_size, ...] for v in img_gts]
87
+ if len(img_gts) == 1:
88
+ img_gts = img_gts[0]
89
+ if len(img_lqs) == 1:
90
+ img_lqs = img_lqs[0]
91
+ return img_gts, img_lqs
92
+
93
+
94
+ def augment(imgs, hflip=True, rotation=True, flows=None, return_status=False):
95
+ """Augment: horizontal flips OR rotate (0, 90, 180, 270 degrees).
96
+
97
+ We use vertical flip and transpose for rotation implementation.
98
+ All the images in the list use the same augmentation.
99
+
100
+ Args:
101
+ imgs (list[ndarray] | ndarray): Images to be augmented. If the input
102
+ is an ndarray, it will be transformed to a list.
103
+ hflip (bool): Horizontal flip. Default: True.
104
+ rotation (bool): Ratotation. Default: True.
105
+ flows (list[ndarray]: Flows to be augmented. If the input is an
106
+ ndarray, it will be transformed to a list.
107
+ Dimension is (h, w, 2). Default: None.
108
+ return_status (bool): Return the status of flip and rotation.
109
+ Default: False.
110
+
111
+ Returns:
112
+ list[ndarray] | ndarray: Augmented images and flows. If returned
113
+ results only have one element, just return ndarray.
114
+
115
+ """
116
+ hflip = hflip and random.random() < 0.5
117
+ vflip = rotation and random.random() < 0.5
118
+ rot90 = rotation and random.random() < 0.5
119
+
120
+ def _augment(img):
121
+ if hflip: # horizontal
122
+ cv2.flip(img, 1, img)
123
+ if vflip: # vertical
124
+ cv2.flip(img, 0, img)
125
+ if rot90:
126
+ img = img.transpose(1, 0, 2)
127
+ return img
128
+
129
+ def _augment_flow(flow):
130
+ if hflip: # horizontal
131
+ cv2.flip(flow, 1, flow)
132
+ flow[:, :, 0] *= -1
133
+ if vflip: # vertical
134
+ cv2.flip(flow, 0, flow)
135
+ flow[:, :, 1] *= -1
136
+ if rot90:
137
+ flow = flow.transpose(1, 0, 2)
138
+ flow = flow[:, :, [1, 0]]
139
+ return flow
140
+
141
+ if not isinstance(imgs, list):
142
+ imgs = [imgs]
143
+ imgs = [_augment(img) for img in imgs]
144
+ if len(imgs) == 1:
145
+ imgs = imgs[0]
146
+
147
+ if flows is not None:
148
+ if not isinstance(flows, list):
149
+ flows = [flows]
150
+ flows = [_augment_flow(flow) for flow in flows]
151
+ if len(flows) == 1:
152
+ flows = flows[0]
153
+ return imgs, flows
154
+ else:
155
+ if return_status:
156
+ return imgs, (hflip, vflip, rot90)
157
+ else:
158
+ return imgs
159
+
160
+
161
+ def img_rotate(img, angle, center=None, scale=1.0):
162
+ """Rotate image.
163
+
164
+ Args:
165
+ img (ndarray): Image to be rotated.
166
+ angle (float): Rotation angle in degrees. Positive values mean
167
+ counter-clockwise rotation.
168
+ center (tuple[int]): Rotation center. If the center is None,
169
+ initialize it as the center of the image. Default: None.
170
+ scale (float): Isotropic scale factor. Default: 1.0.
171
+ """
172
+ (h, w) = img.shape[:2]
173
+
174
+ if center is None:
175
+ center = (w // 2, h // 2)
176
+
177
+ matrix = cv2.getRotationMatrix2D(center, angle, scale)
178
+ rotated_img = cv2.warpAffine(img, matrix, (w, h))
179
+ return rotated_img
official/bsr/utils/__init__.py ADDED
@@ -0,0 +1,47 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .color_util import bgr2ycbcr, rgb2ycbcr, rgb2ycbcr_pt, ycbcr2bgr, ycbcr2rgb
2
+ from .diffjpeg import DiffJPEG
3
+ from .file_client import FileClient
4
+ from .img_process_util import USMSharp, usm_sharp
5
+ from .img_util import crop_border, imfrombytes, img2tensor, imwrite, tensor2img
6
+ from .logger import AvgTimer, MessageLogger, get_env_info, get_root_logger, init_tb_logger, init_wandb_logger
7
+ from .misc import check_resume, get_time_str, make_exp_dirs, mkdir_and_rename, scandir, set_random_seed, sizeof_fmt
8
+ from .options import yaml_load
9
+
10
+ __all__ = [
11
+ # color_util.py
12
+ 'bgr2ycbcr',
13
+ 'rgb2ycbcr',
14
+ 'rgb2ycbcr_pt',
15
+ 'ycbcr2bgr',
16
+ 'ycbcr2rgb',
17
+ # file_client.py
18
+ 'FileClient',
19
+ # img_util.py
20
+ 'img2tensor',
21
+ 'tensor2img',
22
+ 'imfrombytes',
23
+ 'imwrite',
24
+ 'crop_border',
25
+ # logger.py
26
+ 'MessageLogger',
27
+ 'AvgTimer',
28
+ 'init_tb_logger',
29
+ 'init_wandb_logger',
30
+ 'get_root_logger',
31
+ 'get_env_info',
32
+ # misc.py
33
+ 'set_random_seed',
34
+ 'get_time_str',
35
+ 'mkdir_and_rename',
36
+ 'make_exp_dirs',
37
+ 'scandir',
38
+ 'check_resume',
39
+ 'sizeof_fmt',
40
+ # diffjpeg
41
+ 'DiffJPEG',
42
+ # img_process_util
43
+ 'USMSharp',
44
+ 'usm_sharp',
45
+ # options
46
+ 'yaml_load'
47
+ ]
official/bsr/utils/color_util.py ADDED
@@ -0,0 +1,208 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import torch
3
+
4
+
5
+ def rgb2ycbcr(img, y_only=False):
6
+ """Convert a RGB image to YCbCr image.
7
+
8
+ This function produces the same results as Matlab's `rgb2ycbcr` function.
9
+ It implements the ITU-R BT.601 conversion for standard-definition
10
+ television. See more details in
11
+ https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion.
12
+
13
+ It differs from a similar function in cv2.cvtColor: `RGB <-> YCrCb`.
14
+ In OpenCV, it implements a JPEG conversion. See more details in
15
+ https://en.wikipedia.org/wiki/YCbCr#JPEG_conversion.
16
+
17
+ Args:
18
+ img (ndarray): The input image. It accepts:
19
+ 1. np.uint8 type with range [0, 255];
20
+ 2. np.float32 type with range [0, 1].
21
+ y_only (bool): Whether to only return Y channel. Default: False.
22
+
23
+ Returns:
24
+ ndarray: The converted YCbCr image. The output image has the same type
25
+ and range as input image.
26
+ """
27
+ img_type = img.dtype
28
+ img = _convert_input_type_range(img)
29
+ if y_only:
30
+ out_img = np.dot(img, [65.481, 128.553, 24.966]) + 16.0
31
+ else:
32
+ out_img = np.matmul(
33
+ img, [[65.481, -37.797, 112.0], [128.553, -74.203, -93.786], [24.966, 112.0, -18.214]]) + [16, 128, 128]
34
+ out_img = _convert_output_type_range(out_img, img_type)
35
+ return out_img
36
+
37
+
38
+ def bgr2ycbcr(img, y_only=False):
39
+ """Convert a BGR image to YCbCr image.
40
+
41
+ The bgr version of rgb2ycbcr.
42
+ It implements the ITU-R BT.601 conversion for standard-definition
43
+ television. See more details in
44
+ https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion.
45
+
46
+ It differs from a similar function in cv2.cvtColor: `BGR <-> YCrCb`.
47
+ In OpenCV, it implements a JPEG conversion. See more details in
48
+ https://en.wikipedia.org/wiki/YCbCr#JPEG_conversion.
49
+
50
+ Args:
51
+ img (ndarray): The input image. It accepts:
52
+ 1. np.uint8 type with range [0, 255];
53
+ 2. np.float32 type with range [0, 1].
54
+ y_only (bool): Whether to only return Y channel. Default: False.
55
+
56
+ Returns:
57
+ ndarray: The converted YCbCr image. The output image has the same type
58
+ and range as input image.
59
+ """
60
+ img_type = img.dtype
61
+ img = _convert_input_type_range(img)
62
+ if y_only:
63
+ out_img = np.dot(img, [24.966, 128.553, 65.481]) + 16.0
64
+ else:
65
+ out_img = np.matmul(
66
+ img, [[24.966, 112.0, -18.214], [128.553, -74.203, -93.786], [65.481, -37.797, 112.0]]) + [16, 128, 128]
67
+ out_img = _convert_output_type_range(out_img, img_type)
68
+ return out_img
69
+
70
+
71
+ def ycbcr2rgb(img):
72
+ """Convert a YCbCr image to RGB image.
73
+
74
+ This function produces the same results as Matlab's ycbcr2rgb function.
75
+ It implements the ITU-R BT.601 conversion for standard-definition
76
+ television. See more details in
77
+ https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion.
78
+
79
+ It differs from a similar function in cv2.cvtColor: `YCrCb <-> RGB`.
80
+ In OpenCV, it implements a JPEG conversion. See more details in
81
+ https://en.wikipedia.org/wiki/YCbCr#JPEG_conversion.
82
+
83
+ Args:
84
+ img (ndarray): The input image. It accepts:
85
+ 1. np.uint8 type with range [0, 255];
86
+ 2. np.float32 type with range [0, 1].
87
+
88
+ Returns:
89
+ ndarray: The converted RGB image. The output image has the same type
90
+ and range as input image.
91
+ """
92
+ img_type = img.dtype
93
+ img = _convert_input_type_range(img) * 255
94
+ out_img = np.matmul(img, [[0.00456621, 0.00456621, 0.00456621], [0, -0.00153632, 0.00791071],
95
+ [0.00625893, -0.00318811, 0]]) * 255.0 + [-222.921, 135.576, -276.836] # noqa: E126
96
+ out_img = _convert_output_type_range(out_img, img_type)
97
+ return out_img
98
+
99
+
100
+ def ycbcr2bgr(img):
101
+ """Convert a YCbCr image to BGR image.
102
+
103
+ The bgr version of ycbcr2rgb.
104
+ It implements the ITU-R BT.601 conversion for standard-definition
105
+ television. See more details in
106
+ https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion.
107
+
108
+ It differs from a similar function in cv2.cvtColor: `YCrCb <-> BGR`.
109
+ In OpenCV, it implements a JPEG conversion. See more details in
110
+ https://en.wikipedia.org/wiki/YCbCr#JPEG_conversion.
111
+
112
+ Args:
113
+ img (ndarray): The input image. It accepts:
114
+ 1. np.uint8 type with range [0, 255];
115
+ 2. np.float32 type with range [0, 1].
116
+
117
+ Returns:
118
+ ndarray: The converted BGR image. The output image has the same type
119
+ and range as input image.
120
+ """
121
+ img_type = img.dtype
122
+ img = _convert_input_type_range(img) * 255
123
+ out_img = np.matmul(img, [[0.00456621, 0.00456621, 0.00456621], [0.00791071, -0.00153632, 0],
124
+ [0, -0.00318811, 0.00625893]]) * 255.0 + [-276.836, 135.576, -222.921] # noqa: E126
125
+ out_img = _convert_output_type_range(out_img, img_type)
126
+ return out_img
127
+
128
+
129
+ def _convert_input_type_range(img):
130
+ """Convert the type and range of the input image.
131
+
132
+ It converts the input image to np.float32 type and range of [0, 1].
133
+ It is mainly used for pre-processing the input image in colorspace
134
+ conversion functions such as rgb2ycbcr and ycbcr2rgb.
135
+
136
+ Args:
137
+ img (ndarray): The input image. It accepts:
138
+ 1. np.uint8 type with range [0, 255];
139
+ 2. np.float32 type with range [0, 1].
140
+
141
+ Returns:
142
+ (ndarray): The converted image with type of np.float32 and range of
143
+ [0, 1].
144
+ """
145
+ img_type = img.dtype
146
+ img = img.astype(np.float32)
147
+ if img_type == np.float32:
148
+ pass
149
+ elif img_type == np.uint8:
150
+ img /= 255.
151
+ else:
152
+ raise TypeError(f'The img type should be np.float32 or np.uint8, but got {img_type}')
153
+ return img
154
+
155
+
156
+ def _convert_output_type_range(img, dst_type):
157
+ """Convert the type and range of the image according to dst_type.
158
+
159
+ It converts the image to desired type and range. If `dst_type` is np.uint8,
160
+ images will be converted to np.uint8 type with range [0, 255]. If
161
+ `dst_type` is np.float32, it converts the image to np.float32 type with
162
+ range [0, 1].
163
+ It is mainly used for post-processing images in colorspace conversion
164
+ functions such as rgb2ycbcr and ycbcr2rgb.
165
+
166
+ Args:
167
+ img (ndarray): The image to be converted with np.float32 type and
168
+ range [0, 255].
169
+ dst_type (np.uint8 | np.float32): If dst_type is np.uint8, it
170
+ converts the image to np.uint8 type with range [0, 255]. If
171
+ dst_type is np.float32, it converts the image to np.float32 type
172
+ with range [0, 1].
173
+
174
+ Returns:
175
+ (ndarray): The converted image with desired type and range.
176
+ """
177
+ if dst_type not in (np.uint8, np.float32):
178
+ raise TypeError(f'The dst_type should be np.float32 or np.uint8, but got {dst_type}')
179
+ if dst_type == np.uint8:
180
+ img = img.round()
181
+ else:
182
+ img /= 255.
183
+ return img.astype(dst_type)
184
+
185
+
186
+ def rgb2ycbcr_pt(img, y_only=False):
187
+ """Convert RGB images to YCbCr images (PyTorch version).
188
+
189
+ It implements the ITU-R BT.601 conversion for standard-definition television. See more details in
190
+ https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion.
191
+
192
+ Args:
193
+ img (Tensor): Images with shape (n, 3, h, w), the range [0, 1], float, RGB format.
194
+ y_only (bool): Whether to only return Y channel. Default: False.
195
+
196
+ Returns:
197
+ (Tensor): converted images with the shape (n, 3/1, h, w), the range [0, 1], float.
198
+ """
199
+ if y_only:
200
+ weight = torch.tensor([[65.481], [128.553], [24.966]]).to(img)
201
+ out_img = torch.matmul(img.permute(0, 2, 3, 1), weight).permute(0, 3, 1, 2) + 16.0
202
+ else:
203
+ weight = torch.tensor([[65.481, -37.797, 112.0], [128.553, -74.203, -93.786], [24.966, 112.0, -18.214]]).to(img)
204
+ bias = torch.tensor([16, 128, 128]).view(1, 3, 1, 1).to(img)
205
+ out_img = torch.matmul(img.permute(0, 2, 3, 1), weight).permute(0, 3, 1, 2) + bias
206
+
207
+ out_img = out_img / 255.
208
+ return out_img
official/bsr/utils/diffjpeg.py ADDED
@@ -0,0 +1,515 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Modified from https://github.com/mlomnitz/DiffJPEG
3
+
4
+ For images not divisible by 8
5
+ https://dsp.stackexchange.com/questions/35339/jpeg-dct-padding/35343#35343
6
+ """
7
+ import itertools
8
+ import numpy as np
9
+ import torch
10
+ import torch.nn as nn
11
+ from torch.nn import functional as F
12
+
13
+ # ------------------------ utils ------------------------#
14
+ y_table = np.array(
15
+ [[16, 11, 10, 16, 24, 40, 51, 61], [12, 12, 14, 19, 26, 58, 60, 55], [14, 13, 16, 24, 40, 57, 69, 56],
16
+ [14, 17, 22, 29, 51, 87, 80, 62], [18, 22, 37, 56, 68, 109, 103, 77], [24, 35, 55, 64, 81, 104, 113, 92],
17
+ [49, 64, 78, 87, 103, 121, 120, 101], [72, 92, 95, 98, 112, 100, 103, 99]],
18
+ dtype=np.float32).T
19
+ y_table = nn.Parameter(torch.from_numpy(y_table))
20
+ c_table = np.empty((8, 8), dtype=np.float32)
21
+ c_table.fill(99)
22
+ c_table[:4, :4] = np.array([[17, 18, 24, 47], [18, 21, 26, 66], [24, 26, 56, 99], [47, 66, 99, 99]]).T
23
+ c_table = nn.Parameter(torch.from_numpy(c_table))
24
+
25
+
26
+ def diff_round(x):
27
+ """ Differentiable rounding function
28
+ """
29
+ return torch.round(x) + (x - torch.round(x))**3
30
+
31
+
32
+ def quality_to_factor(quality):
33
+ """ Calculate factor corresponding to quality
34
+
35
+ Args:
36
+ quality(float): Quality for jpeg compression.
37
+
38
+ Returns:
39
+ float: Compression factor.
40
+ """
41
+ if quality < 50:
42
+ quality = 5000. / quality
43
+ else:
44
+ quality = 200. - quality * 2
45
+ return quality / 100.
46
+
47
+
48
+ # ------------------------ compression ------------------------#
49
+ class RGB2YCbCrJpeg(nn.Module):
50
+ """ Converts RGB image to YCbCr
51
+ """
52
+
53
+ def __init__(self):
54
+ super(RGB2YCbCrJpeg, self).__init__()
55
+ matrix = np.array([[0.299, 0.587, 0.114], [-0.168736, -0.331264, 0.5], [0.5, -0.418688, -0.081312]],
56
+ dtype=np.float32).T
57
+ self.shift = nn.Parameter(torch.tensor([0., 128., 128.]))
58
+ self.matrix = nn.Parameter(torch.from_numpy(matrix))
59
+
60
+ def forward(self, image):
61
+ """
62
+ Args:
63
+ image(Tensor): batch x 3 x height x width
64
+
65
+ Returns:
66
+ Tensor: batch x height x width x 3
67
+ """
68
+ image = image.permute(0, 2, 3, 1)
69
+ result = torch.tensordot(image, self.matrix, dims=1) + self.shift
70
+ return result.view(image.shape)
71
+
72
+
73
+ class ChromaSubsampling(nn.Module):
74
+ """ Chroma subsampling on CbCr channels
75
+ """
76
+
77
+ def __init__(self):
78
+ super(ChromaSubsampling, self).__init__()
79
+
80
+ def forward(self, image):
81
+ """
82
+ Args:
83
+ image(tensor): batch x height x width x 3
84
+
85
+ Returns:
86
+ y(tensor): batch x height x width
87
+ cb(tensor): batch x height/2 x width/2
88
+ cr(tensor): batch x height/2 x width/2
89
+ """
90
+ image_2 = image.permute(0, 3, 1, 2).clone()
91
+ cb = F.avg_pool2d(image_2[:, 1, :, :].unsqueeze(1), kernel_size=2, stride=(2, 2), count_include_pad=False)
92
+ cr = F.avg_pool2d(image_2[:, 2, :, :].unsqueeze(1), kernel_size=2, stride=(2, 2), count_include_pad=False)
93
+ cb = cb.permute(0, 2, 3, 1)
94
+ cr = cr.permute(0, 2, 3, 1)
95
+ return image[:, :, :, 0], cb.squeeze(3), cr.squeeze(3)
96
+
97
+
98
+ class BlockSplitting(nn.Module):
99
+ """ Splitting image into patches
100
+ """
101
+
102
+ def __init__(self):
103
+ super(BlockSplitting, self).__init__()
104
+ self.k = 8
105
+
106
+ def forward(self, image):
107
+ """
108
+ Args:
109
+ image(tensor): batch x height x width
110
+
111
+ Returns:
112
+ Tensor: batch x h*w/64 x h x w
113
+ """
114
+ height, _ = image.shape[1:3]
115
+ batch_size = image.shape[0]
116
+ image_reshaped = image.view(batch_size, height // self.k, self.k, -1, self.k)
117
+ image_transposed = image_reshaped.permute(0, 1, 3, 2, 4)
118
+ return image_transposed.contiguous().view(batch_size, -1, self.k, self.k)
119
+
120
+
121
+ class DCT8x8(nn.Module):
122
+ """ Discrete Cosine Transformation
123
+ """
124
+
125
+ def __init__(self):
126
+ super(DCT8x8, self).__init__()
127
+ tensor = np.zeros((8, 8, 8, 8), dtype=np.float32)
128
+ for x, y, u, v in itertools.product(range(8), repeat=4):
129
+ tensor[x, y, u, v] = np.cos((2 * x + 1) * u * np.pi / 16) * np.cos((2 * y + 1) * v * np.pi / 16)
130
+ alpha = np.array([1. / np.sqrt(2)] + [1] * 7)
131
+ self.tensor = nn.Parameter(torch.from_numpy(tensor).float())
132
+ self.scale = nn.Parameter(torch.from_numpy(np.outer(alpha, alpha) * 0.25).float())
133
+
134
+ def forward(self, image):
135
+ """
136
+ Args:
137
+ image(tensor): batch x height x width
138
+
139
+ Returns:
140
+ Tensor: batch x height x width
141
+ """
142
+ image = image - 128
143
+ result = self.scale * torch.tensordot(image, self.tensor, dims=2)
144
+ result.view(image.shape)
145
+ return result
146
+
147
+
148
+ class YQuantize(nn.Module):
149
+ """ JPEG Quantization for Y channel
150
+
151
+ Args:
152
+ rounding(function): rounding function to use
153
+ """
154
+
155
+ def __init__(self, rounding):
156
+ super(YQuantize, self).__init__()
157
+ self.rounding = rounding
158
+ self.y_table = y_table
159
+
160
+ def forward(self, image, factor=1):
161
+ """
162
+ Args:
163
+ image(tensor): batch x height x width
164
+
165
+ Returns:
166
+ Tensor: batch x height x width
167
+ """
168
+ if isinstance(factor, (int, float)):
169
+ image = image.float() / (self.y_table * factor)
170
+ else:
171
+ b = factor.size(0)
172
+ table = self.y_table.expand(b, 1, 8, 8) * factor.view(b, 1, 1, 1)
173
+ image = image.float() / table
174
+ image = self.rounding(image)
175
+ return image
176
+
177
+
178
+ class CQuantize(nn.Module):
179
+ """ JPEG Quantization for CbCr channels
180
+
181
+ Args:
182
+ rounding(function): rounding function to use
183
+ """
184
+
185
+ def __init__(self, rounding):
186
+ super(CQuantize, self).__init__()
187
+ self.rounding = rounding
188
+ self.c_table = c_table
189
+
190
+ def forward(self, image, factor=1):
191
+ """
192
+ Args:
193
+ image(tensor): batch x height x width
194
+
195
+ Returns:
196
+ Tensor: batch x height x width
197
+ """
198
+ if isinstance(factor, (int, float)):
199
+ image = image.float() / (self.c_table * factor)
200
+ else:
201
+ b = factor.size(0)
202
+ table = self.c_table.expand(b, 1, 8, 8) * factor.view(b, 1, 1, 1)
203
+ image = image.float() / table
204
+ image = self.rounding(image)
205
+ return image
206
+
207
+
208
+ class CompressJpeg(nn.Module):
209
+ """Full JPEG compression algorithm
210
+
211
+ Args:
212
+ rounding(function): rounding function to use
213
+ """
214
+
215
+ def __init__(self, rounding=torch.round):
216
+ super(CompressJpeg, self).__init__()
217
+ self.l1 = nn.Sequential(RGB2YCbCrJpeg(), ChromaSubsampling())
218
+ self.l2 = nn.Sequential(BlockSplitting(), DCT8x8())
219
+ self.c_quantize = CQuantize(rounding=rounding)
220
+ self.y_quantize = YQuantize(rounding=rounding)
221
+
222
+ def forward(self, image, factor=1):
223
+ """
224
+ Args:
225
+ image(tensor): batch x 3 x height x width
226
+
227
+ Returns:
228
+ dict(tensor): Compressed tensor with batch x h*w/64 x 8 x 8.
229
+ """
230
+ y, cb, cr = self.l1(image * 255)
231
+ components = {'y': y, 'cb': cb, 'cr': cr}
232
+ for k in components.keys():
233
+ comp = self.l2(components[k])
234
+ if k in ('cb', 'cr'):
235
+ comp = self.c_quantize(comp, factor=factor)
236
+ else:
237
+ comp = self.y_quantize(comp, factor=factor)
238
+
239
+ components[k] = comp
240
+
241
+ return components['y'], components['cb'], components['cr']
242
+
243
+
244
+ # ------------------------ decompression ------------------------#
245
+
246
+
247
+ class YDequantize(nn.Module):
248
+ """Dequantize Y channel
249
+ """
250
+
251
+ def __init__(self):
252
+ super(YDequantize, self).__init__()
253
+ self.y_table = y_table
254
+
255
+ def forward(self, image, factor=1):
256
+ """
257
+ Args:
258
+ image(tensor): batch x height x width
259
+
260
+ Returns:
261
+ Tensor: batch x height x width
262
+ """
263
+ if isinstance(factor, (int, float)):
264
+ out = image * (self.y_table * factor)
265
+ else:
266
+ b = factor.size(0)
267
+ table = self.y_table.expand(b, 1, 8, 8) * factor.view(b, 1, 1, 1)
268
+ out = image * table
269
+ return out
270
+
271
+
272
+ class CDequantize(nn.Module):
273
+ """Dequantize CbCr channel
274
+ """
275
+
276
+ def __init__(self):
277
+ super(CDequantize, self).__init__()
278
+ self.c_table = c_table
279
+
280
+ def forward(self, image, factor=1):
281
+ """
282
+ Args:
283
+ image(tensor): batch x height x width
284
+
285
+ Returns:
286
+ Tensor: batch x height x width
287
+ """
288
+ if isinstance(factor, (int, float)):
289
+ out = image * (self.c_table * factor)
290
+ else:
291
+ b = factor.size(0)
292
+ table = self.c_table.expand(b, 1, 8, 8) * factor.view(b, 1, 1, 1)
293
+ out = image * table
294
+ return out
295
+
296
+
297
+ class iDCT8x8(nn.Module):
298
+ """Inverse discrete Cosine Transformation
299
+ """
300
+
301
+ def __init__(self):
302
+ super(iDCT8x8, self).__init__()
303
+ alpha = np.array([1. / np.sqrt(2)] + [1] * 7)
304
+ self.alpha = nn.Parameter(torch.from_numpy(np.outer(alpha, alpha)).float())
305
+ tensor = np.zeros((8, 8, 8, 8), dtype=np.float32)
306
+ for x, y, u, v in itertools.product(range(8), repeat=4):
307
+ tensor[x, y, u, v] = np.cos((2 * u + 1) * x * np.pi / 16) * np.cos((2 * v + 1) * y * np.pi / 16)
308
+ self.tensor = nn.Parameter(torch.from_numpy(tensor).float())
309
+
310
+ def forward(self, image):
311
+ """
312
+ Args:
313
+ image(tensor): batch x height x width
314
+
315
+ Returns:
316
+ Tensor: batch x height x width
317
+ """
318
+ image = image * self.alpha
319
+ result = 0.25 * torch.tensordot(image, self.tensor, dims=2) + 128
320
+ result.view(image.shape)
321
+ return result
322
+
323
+
324
+ class BlockMerging(nn.Module):
325
+ """Merge patches into image
326
+ """
327
+
328
+ def __init__(self):
329
+ super(BlockMerging, self).__init__()
330
+
331
+ def forward(self, patches, height, width):
332
+ """
333
+ Args:
334
+ patches(tensor) batch x height*width/64, height x width
335
+ height(int)
336
+ width(int)
337
+
338
+ Returns:
339
+ Tensor: batch x height x width
340
+ """
341
+ k = 8
342
+ batch_size = patches.shape[0]
343
+ image_reshaped = patches.view(batch_size, height // k, width // k, k, k)
344
+ image_transposed = image_reshaped.permute(0, 1, 3, 2, 4)
345
+ return image_transposed.contiguous().view(batch_size, height, width)
346
+
347
+
348
+ class ChromaUpsampling(nn.Module):
349
+ """Upsample chroma layers
350
+ """
351
+
352
+ def __init__(self):
353
+ super(ChromaUpsampling, self).__init__()
354
+
355
+ def forward(self, y, cb, cr):
356
+ """
357
+ Args:
358
+ y(tensor): y channel image
359
+ cb(tensor): cb channel
360
+ cr(tensor): cr channel
361
+
362
+ Returns:
363
+ Tensor: batch x height x width x 3
364
+ """
365
+
366
+ def repeat(x, k=2):
367
+ height, width = x.shape[1:3]
368
+ x = x.unsqueeze(-1)
369
+ x = x.repeat(1, 1, k, k)
370
+ x = x.view(-1, height * k, width * k)
371
+ return x
372
+
373
+ cb = repeat(cb)
374
+ cr = repeat(cr)
375
+ return torch.cat([y.unsqueeze(3), cb.unsqueeze(3), cr.unsqueeze(3)], dim=3)
376
+
377
+
378
+ class YCbCr2RGBJpeg(nn.Module):
379
+ """Converts YCbCr image to RGB JPEG
380
+ """
381
+
382
+ def __init__(self):
383
+ super(YCbCr2RGBJpeg, self).__init__()
384
+
385
+ matrix = np.array([[1., 0., 1.402], [1, -0.344136, -0.714136], [1, 1.772, 0]], dtype=np.float32).T
386
+ self.shift = nn.Parameter(torch.tensor([0, -128., -128.]))
387
+ self.matrix = nn.Parameter(torch.from_numpy(matrix))
388
+
389
+ def forward(self, image):
390
+ """
391
+ Args:
392
+ image(tensor): batch x height x width x 3
393
+
394
+ Returns:
395
+ Tensor: batch x 3 x height x width
396
+ """
397
+ result = torch.tensordot(image + self.shift, self.matrix, dims=1)
398
+ return result.view(image.shape).permute(0, 3, 1, 2)
399
+
400
+
401
+ class DeCompressJpeg(nn.Module):
402
+ """Full JPEG decompression algorithm
403
+
404
+ Args:
405
+ rounding(function): rounding function to use
406
+ """
407
+
408
+ def __init__(self, rounding=torch.round):
409
+ super(DeCompressJpeg, self).__init__()
410
+ self.c_dequantize = CDequantize()
411
+ self.y_dequantize = YDequantize()
412
+ self.idct = iDCT8x8()
413
+ self.merging = BlockMerging()
414
+ self.chroma = ChromaUpsampling()
415
+ self.colors = YCbCr2RGBJpeg()
416
+
417
+ def forward(self, y, cb, cr, imgh, imgw, factor=1):
418
+ """
419
+ Args:
420
+ compressed(dict(tensor)): batch x h*w/64 x 8 x 8
421
+ imgh(int)
422
+ imgw(int)
423
+ factor(float)
424
+
425
+ Returns:
426
+ Tensor: batch x 3 x height x width
427
+ """
428
+ components = {'y': y, 'cb': cb, 'cr': cr}
429
+ for k in components.keys():
430
+ if k in ('cb', 'cr'):
431
+ comp = self.c_dequantize(components[k], factor=factor)
432
+ height, width = int(imgh / 2), int(imgw / 2)
433
+ else:
434
+ comp = self.y_dequantize(components[k], factor=factor)
435
+ height, width = imgh, imgw
436
+ comp = self.idct(comp)
437
+ components[k] = self.merging(comp, height, width)
438
+ #
439
+ image = self.chroma(components['y'], components['cb'], components['cr'])
440
+ image = self.colors(image)
441
+
442
+ image = torch.min(255 * torch.ones_like(image), torch.max(torch.zeros_like(image), image))
443
+ return image / 255
444
+
445
+
446
+ # ------------------------ main DiffJPEG ------------------------ #
447
+
448
+
449
+ class DiffJPEG(nn.Module):
450
+ """This JPEG algorithm result is slightly different from cv2.
451
+ DiffJPEG supports batch processing.
452
+
453
+ Args:
454
+ differentiable(bool): If True, uses custom differentiable rounding function, if False, uses standard torch.round
455
+ """
456
+
457
+ def __init__(self, differentiable=True):
458
+ super(DiffJPEG, self).__init__()
459
+ if differentiable:
460
+ rounding = diff_round
461
+ else:
462
+ rounding = torch.round
463
+
464
+ self.compress = CompressJpeg(rounding=rounding)
465
+ self.decompress = DeCompressJpeg(rounding=rounding)
466
+
467
+ def forward(self, x, quality):
468
+ """
469
+ Args:
470
+ x (Tensor): Input image, bchw, rgb, [0, 1]
471
+ quality(float): Quality factor for jpeg compression scheme.
472
+ """
473
+ factor = quality
474
+ if isinstance(factor, (int, float)):
475
+ factor = quality_to_factor(factor)
476
+ else:
477
+ for i in range(factor.size(0)):
478
+ factor[i] = quality_to_factor(factor[i])
479
+ h, w = x.size()[-2:]
480
+ h_pad, w_pad = 0, 0
481
+ # why should use 16
482
+ if h % 16 != 0:
483
+ h_pad = 16 - h % 16
484
+ if w % 16 != 0:
485
+ w_pad = 16 - w % 16
486
+ x = F.pad(x, (0, w_pad, 0, h_pad), mode='constant', value=0)
487
+
488
+ y, cb, cr = self.compress(x, factor=factor)
489
+ recovered = self.decompress(y, cb, cr, (h + h_pad), (w + w_pad), factor=factor)
490
+ recovered = recovered[:, :, 0:h, 0:w]
491
+ return recovered
492
+
493
+
494
+ if __name__ == '__main__':
495
+ import cv2
496
+
497
+ from bsr.utils import img2tensor, tensor2img
498
+
499
+ img_gt = cv2.imread('test.png') / 255.
500
+
501
+ # -------------- cv2 -------------- #
502
+ encode_param = [int(cv2.IMWRITE_JPEG_QUALITY), 20]
503
+ _, encimg = cv2.imencode('.jpg', img_gt * 255., encode_param)
504
+ img_lq = np.float32(cv2.imdecode(encimg, 1))
505
+ cv2.imwrite('cv2_JPEG_20.png', img_lq)
506
+
507
+ # -------------- DiffJPEG -------------- #
508
+ jpeger = DiffJPEG(differentiable=False).cuda()
509
+ img_gt = img2tensor(img_gt)
510
+ img_gt = torch.stack([img_gt, img_gt]).cuda()
511
+ quality = img_gt.new_tensor([20, 40])
512
+ out = jpeger(img_gt, quality=quality)
513
+
514
+ cv2.imwrite('pt_JPEG_20.png', tensor2img(out[0]))
515
+ cv2.imwrite('pt_JPEG_40.png', tensor2img(out[1]))
official/bsr/utils/dist_util.py ADDED
@@ -0,0 +1,82 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Modified from https://github.com/open-mmlab/mmcv/blob/master/mmcv/runner/dist_utils.py # noqa: E501
2
+ import functools
3
+ import os
4
+ import subprocess
5
+ import torch
6
+ import torch.distributed as dist
7
+ import torch.multiprocessing as mp
8
+
9
+
10
+ def init_dist(launcher, backend='nccl', **kwargs):
11
+ if mp.get_start_method(allow_none=True) is None:
12
+ mp.set_start_method('spawn')
13
+ if launcher == 'pytorch':
14
+ _init_dist_pytorch(backend, **kwargs)
15
+ elif launcher == 'slurm':
16
+ _init_dist_slurm(backend, **kwargs)
17
+ else:
18
+ raise ValueError(f'Invalid launcher type: {launcher}')
19
+
20
+
21
+ def _init_dist_pytorch(backend, **kwargs):
22
+ rank = int(os.environ['RANK'])
23
+ num_gpus = torch.cuda.device_count()
24
+ torch.cuda.set_device(rank % num_gpus)
25
+ dist.init_process_group(backend=backend, **kwargs)
26
+
27
+
28
+ def _init_dist_slurm(backend, port=None):
29
+ """Initialize slurm distributed training environment.
30
+
31
+ If argument ``port`` is not specified, then the master port will be system
32
+ environment variable ``MASTER_PORT``. If ``MASTER_PORT`` is not in system
33
+ environment variable, then a default port ``29500`` will be used.
34
+
35
+ Args:
36
+ backend (str): Backend of torch.distributed.
37
+ port (int, optional): Master port. Defaults to None.
38
+ """
39
+ proc_id = int(os.environ['SLURM_PROCID'])
40
+ ntasks = int(os.environ['SLURM_NTASKS'])
41
+ node_list = os.environ['SLURM_NODELIST']
42
+ num_gpus = torch.cuda.device_count()
43
+ torch.cuda.set_device(proc_id % num_gpus)
44
+ addr = subprocess.getoutput(f'scontrol show hostname {node_list} | head -n1')
45
+ # specify master port
46
+ if port is not None:
47
+ os.environ['MASTER_PORT'] = str(port)
48
+ elif 'MASTER_PORT' in os.environ:
49
+ pass # use MASTER_PORT in the environment variable
50
+ else:
51
+ # 29500 is torch.distributed default port
52
+ os.environ['MASTER_PORT'] = '29500'
53
+ os.environ['MASTER_ADDR'] = addr
54
+ os.environ['WORLD_SIZE'] = str(ntasks)
55
+ os.environ['LOCAL_RANK'] = str(proc_id % num_gpus)
56
+ os.environ['RANK'] = str(proc_id)
57
+ dist.init_process_group(backend=backend)
58
+
59
+
60
+ def get_dist_info():
61
+ if dist.is_available():
62
+ initialized = dist.is_initialized()
63
+ else:
64
+ initialized = False
65
+ if initialized:
66
+ rank = dist.get_rank()
67
+ world_size = dist.get_world_size()
68
+ else:
69
+ rank = 0
70
+ world_size = 1
71
+ return rank, world_size
72
+
73
+
74
+ def master_only(func):
75
+
76
+ @functools.wraps(func)
77
+ def wrapper(*args, **kwargs):
78
+ rank, _ = get_dist_info()
79
+ if rank == 0:
80
+ return func(*args, **kwargs)
81
+
82
+ return wrapper
official/bsr/utils/download_util.py ADDED
@@ -0,0 +1,98 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ import os
3
+ import requests
4
+ from torch.hub import download_url_to_file, get_dir
5
+ from tqdm import tqdm
6
+ from urllib.parse import urlparse
7
+
8
+ from .misc import sizeof_fmt
9
+
10
+
11
+ def download_file_from_google_drive(file_id, save_path):
12
+ """Download files from google drive.
13
+
14
+ Reference: https://stackoverflow.com/questions/25010369/wget-curl-large-file-from-google-drive
15
+
16
+ Args:
17
+ file_id (str): File id.
18
+ save_path (str): Save path.
19
+ """
20
+
21
+ session = requests.Session()
22
+ URL = 'https://docs.google.com/uc?export=download'
23
+ params = {'id': file_id}
24
+
25
+ response = session.get(URL, params=params, stream=True)
26
+ token = get_confirm_token(response)
27
+ if token:
28
+ params['confirm'] = token
29
+ response = session.get(URL, params=params, stream=True)
30
+
31
+ # get file size
32
+ response_file_size = session.get(URL, params=params, stream=True, headers={'Range': 'bytes=0-2'})
33
+ if 'Content-Range' in response_file_size.headers:
34
+ file_size = int(response_file_size.headers['Content-Range'].split('/')[1])
35
+ else:
36
+ file_size = None
37
+
38
+ save_response_content(response, save_path, file_size)
39
+
40
+
41
+ def get_confirm_token(response):
42
+ for key, value in response.cookies.items():
43
+ if key.startswith('download_warning'):
44
+ return value
45
+ return None
46
+
47
+
48
+ def save_response_content(response, destination, file_size=None, chunk_size=32768):
49
+ if file_size is not None:
50
+ pbar = tqdm(total=math.ceil(file_size / chunk_size), unit='chunk')
51
+
52
+ readable_file_size = sizeof_fmt(file_size)
53
+ else:
54
+ pbar = None
55
+
56
+ with open(destination, 'wb') as f:
57
+ downloaded_size = 0
58
+ for chunk in response.iter_content(chunk_size):
59
+ downloaded_size += chunk_size
60
+ if pbar is not None:
61
+ pbar.update(1)
62
+ pbar.set_description(f'Download {sizeof_fmt(downloaded_size)} / {readable_file_size}')
63
+ if chunk: # filter out keep-alive new chunks
64
+ f.write(chunk)
65
+ if pbar is not None:
66
+ pbar.close()
67
+
68
+
69
+ def load_file_from_url(url, model_dir=None, progress=True, file_name=None):
70
+ """Load file form http url, will download models if necessary.
71
+
72
+ Reference: https://github.com/1adrianb/face-alignment/blob/master/face_alignment/utils.py
73
+
74
+ Args:
75
+ url (str): URL to be downloaded.
76
+ model_dir (str): The path to save the downloaded model. Should be a full path. If None, use pytorch hub_dir.
77
+ Default: None.
78
+ progress (bool): Whether to show the download progress. Default: True.
79
+ file_name (str): The downloaded file name. If None, use the file name in the url. Default: None.
80
+
81
+ Returns:
82
+ str: The path to the downloaded file.
83
+ """
84
+ if model_dir is None: # use the pytorch hub_dir
85
+ hub_dir = get_dir()
86
+ model_dir = os.path.join(hub_dir, 'checkpoints')
87
+
88
+ os.makedirs(model_dir, exist_ok=True)
89
+
90
+ parts = urlparse(url)
91
+ filename = os.path.basename(parts.path)
92
+ if file_name is not None:
93
+ filename = file_name
94
+ cached_file = os.path.abspath(os.path.join(model_dir, filename))
95
+ if not os.path.exists(cached_file):
96
+ print(f'Downloading: "{url}" to {cached_file}\n')
97
+ download_url_to_file(url, cached_file, hash_prefix=None, progress=progress)
98
+ return cached_file
official/bsr/utils/file_client.py ADDED
@@ -0,0 +1,167 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Modified from https://github.com/open-mmlab/mmcv/blob/master/mmcv/fileio/file_client.py # noqa: E501
2
+ from abc import ABCMeta, abstractmethod
3
+
4
+
5
+ class BaseStorageBackend(metaclass=ABCMeta):
6
+ """Abstract class of storage backends.
7
+
8
+ All backends need to implement two apis: ``get()`` and ``get_text()``.
9
+ ``get()`` reads the file as a byte stream and ``get_text()`` reads the file
10
+ as texts.
11
+ """
12
+
13
+ @abstractmethod
14
+ def get(self, filepath):
15
+ pass
16
+
17
+ @abstractmethod
18
+ def get_text(self, filepath):
19
+ pass
20
+
21
+
22
+ class MemcachedBackend(BaseStorageBackend):
23
+ """Memcached storage backend.
24
+
25
+ Attributes:
26
+ server_list_cfg (str): Config file for memcached server list.
27
+ client_cfg (str): Config file for memcached client.
28
+ sys_path (str | None): Additional path to be appended to `sys.path`.
29
+ Default: None.
30
+ """
31
+
32
+ def __init__(self, server_list_cfg, client_cfg, sys_path=None):
33
+ if sys_path is not None:
34
+ import sys
35
+ sys.path.append(sys_path)
36
+ try:
37
+ import mc
38
+ except ImportError:
39
+ raise ImportError('Please install memcached to enable MemcachedBackend.')
40
+
41
+ self.server_list_cfg = server_list_cfg
42
+ self.client_cfg = client_cfg
43
+ self._client = mc.MemcachedClient.GetInstance(self.server_list_cfg, self.client_cfg)
44
+ # mc.pyvector servers as a point which points to a memory cache
45
+ self._mc_buffer = mc.pyvector()
46
+
47
+ def get(self, filepath):
48
+ filepath = str(filepath)
49
+ import mc
50
+ self._client.Get(filepath, self._mc_buffer)
51
+ value_buf = mc.ConvertBuffer(self._mc_buffer)
52
+ return value_buf
53
+
54
+ def get_text(self, filepath):
55
+ raise NotImplementedError
56
+
57
+
58
+ class HardDiskBackend(BaseStorageBackend):
59
+ """Raw hard disks storage backend."""
60
+
61
+ def get(self, filepath):
62
+ filepath = str(filepath)
63
+ with open(filepath, 'rb') as f:
64
+ value_buf = f.read()
65
+ return value_buf
66
+
67
+ def get_text(self, filepath):
68
+ filepath = str(filepath)
69
+ with open(filepath, 'r') as f:
70
+ value_buf = f.read()
71
+ return value_buf
72
+
73
+
74
+ class LmdbBackend(BaseStorageBackend):
75
+ """Lmdb storage backend.
76
+
77
+ Args:
78
+ db_paths (str | list[str]): Lmdb database paths.
79
+ client_keys (str | list[str]): Lmdb client keys. Default: 'default'.
80
+ readonly (bool, optional): Lmdb environment parameter. If True,
81
+ disallow any write operations. Default: True.
82
+ lock (bool, optional): Lmdb environment parameter. If False, when
83
+ concurrent access occurs, do not lock the database. Default: False.
84
+ readahead (bool, optional): Lmdb environment parameter. If False,
85
+ disable the OS filesystem readahead mechanism, which may improve
86
+ random read performance when a database is larger than RAM.
87
+ Default: False.
88
+
89
+ Attributes:
90
+ db_paths (list): Lmdb database path.
91
+ _client (list): A list of several lmdb envs.
92
+ """
93
+
94
+ def __init__(self, db_paths, client_keys='default', readonly=True, lock=False, readahead=False, **kwargs):
95
+ try:
96
+ import lmdb
97
+ except ImportError:
98
+ raise ImportError('Please install lmdb to enable LmdbBackend.')
99
+
100
+ if isinstance(client_keys, str):
101
+ client_keys = [client_keys]
102
+
103
+ if isinstance(db_paths, list):
104
+ self.db_paths = [str(v) for v in db_paths]
105
+ elif isinstance(db_paths, str):
106
+ self.db_paths = [str(db_paths)]
107
+ assert len(client_keys) == len(self.db_paths), ('client_keys and db_paths should have the same length, '
108
+ f'but received {len(client_keys)} and {len(self.db_paths)}.')
109
+
110
+ self._client = {}
111
+ for client, path in zip(client_keys, self.db_paths):
112
+ self._client[client] = lmdb.open(path, readonly=readonly, lock=lock, readahead=readahead, **kwargs)
113
+
114
+ def get(self, filepath, client_key):
115
+ """Get values according to the filepath from one lmdb named client_key.
116
+
117
+ Args:
118
+ filepath (str | obj:`Path`): Here, filepath is the lmdb key.
119
+ client_key (str): Used for distinguishing different lmdb envs.
120
+ """
121
+ filepath = str(filepath)
122
+ assert client_key in self._client, (f'client_key {client_key} is not in lmdb clients.')
123
+ client = self._client[client_key]
124
+ with client.begin(write=False) as txn:
125
+ value_buf = txn.get(filepath.encode('ascii'))
126
+ return value_buf
127
+
128
+ def get_text(self, filepath):
129
+ raise NotImplementedError
130
+
131
+
132
+ class FileClient(object):
133
+ """A general file client to access files in different backend.
134
+
135
+ The client loads a file or text in a specified backend from its path
136
+ and return it as a binary file. it can also register other backend
137
+ accessor with a given name and backend class.
138
+
139
+ Attributes:
140
+ backend (str): The storage backend type. Options are "disk",
141
+ "memcached" and "lmdb".
142
+ client (:obj:`BaseStorageBackend`): The backend object.
143
+ """
144
+
145
+ _backends = {
146
+ 'disk': HardDiskBackend,
147
+ 'memcached': MemcachedBackend,
148
+ 'lmdb': LmdbBackend,
149
+ }
150
+
151
+ def __init__(self, backend='disk', **kwargs):
152
+ if backend not in self._backends:
153
+ raise ValueError(f'Backend {backend} is not supported. Currently supported ones'
154
+ f' are {list(self._backends.keys())}')
155
+ self.backend = backend
156
+ self.client = self._backends[backend](**kwargs)
157
+
158
+ def get(self, filepath, client_key='default'):
159
+ # client_key is used only for lmdb, where different fileclients have
160
+ # different lmdb environments.
161
+ if self.backend == 'lmdb':
162
+ return self.client.get(filepath, client_key)
163
+ else:
164
+ return self.client.get(filepath)
165
+
166
+ def get_text(self, filepath):
167
+ return self.client.get_text(filepath)
official/bsr/utils/flow_util.py ADDED
@@ -0,0 +1,170 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Modified from https://github.com/open-mmlab/mmcv/blob/master/mmcv/video/optflow.py # noqa: E501
2
+ import cv2
3
+ import numpy as np
4
+ import os
5
+
6
+
7
+ def flowread(flow_path, quantize=False, concat_axis=0, *args, **kwargs):
8
+ """Read an optical flow map.
9
+
10
+ Args:
11
+ flow_path (ndarray or str): Flow path.
12
+ quantize (bool): whether to read quantized pair, if set to True,
13
+ remaining args will be passed to :func:`dequantize_flow`.
14
+ concat_axis (int): The axis that dx and dy are concatenated,
15
+ can be either 0 or 1. Ignored if quantize is False.
16
+
17
+ Returns:
18
+ ndarray: Optical flow represented as a (h, w, 2) numpy array
19
+ """
20
+ if quantize:
21
+ assert concat_axis in [0, 1]
22
+ cat_flow = cv2.imread(flow_path, cv2.IMREAD_UNCHANGED)
23
+ if cat_flow.ndim != 2:
24
+ raise IOError(f'{flow_path} is not a valid quantized flow file, its dimension is {cat_flow.ndim}.')
25
+ assert cat_flow.shape[concat_axis] % 2 == 0
26
+ dx, dy = np.split(cat_flow, 2, axis=concat_axis)
27
+ flow = dequantize_flow(dx, dy, *args, **kwargs)
28
+ else:
29
+ with open(flow_path, 'rb') as f:
30
+ try:
31
+ header = f.read(4).decode('utf-8')
32
+ except Exception:
33
+ raise IOError(f'Invalid flow file: {flow_path}')
34
+ else:
35
+ if header != 'PIEH':
36
+ raise IOError(f'Invalid flow file: {flow_path}, header does not contain PIEH')
37
+
38
+ w = np.fromfile(f, np.int32, 1).squeeze()
39
+ h = np.fromfile(f, np.int32, 1).squeeze()
40
+ flow = np.fromfile(f, np.float32, w * h * 2).reshape((h, w, 2))
41
+
42
+ return flow.astype(np.float32)
43
+
44
+
45
+ def flowwrite(flow, filename, quantize=False, concat_axis=0, *args, **kwargs):
46
+ """Write optical flow to file.
47
+
48
+ If the flow is not quantized, it will be saved as a .flo file losslessly,
49
+ otherwise a jpeg image which is lossy but of much smaller size. (dx and dy
50
+ will be concatenated horizontally into a single image if quantize is True.)
51
+
52
+ Args:
53
+ flow (ndarray): (h, w, 2) array of optical flow.
54
+ filename (str): Output filepath.
55
+ quantize (bool): Whether to quantize the flow and save it to 2 jpeg
56
+ images. If set to True, remaining args will be passed to
57
+ :func:`quantize_flow`.
58
+ concat_axis (int): The axis that dx and dy are concatenated,
59
+ can be either 0 or 1. Ignored if quantize is False.
60
+ """
61
+ if not quantize:
62
+ with open(filename, 'wb') as f:
63
+ f.write('PIEH'.encode('utf-8'))
64
+ np.array([flow.shape[1], flow.shape[0]], dtype=np.int32).tofile(f)
65
+ flow = flow.astype(np.float32)
66
+ flow.tofile(f)
67
+ f.flush()
68
+ else:
69
+ assert concat_axis in [0, 1]
70
+ dx, dy = quantize_flow(flow, *args, **kwargs)
71
+ dxdy = np.concatenate((dx, dy), axis=concat_axis)
72
+ os.makedirs(os.path.dirname(filename), exist_ok=True)
73
+ cv2.imwrite(filename, dxdy)
74
+
75
+
76
+ def quantize_flow(flow, max_val=0.02, norm=True):
77
+ """Quantize flow to [0, 255].
78
+
79
+ After this step, the size of flow will be much smaller, and can be
80
+ dumped as jpeg images.
81
+
82
+ Args:
83
+ flow (ndarray): (h, w, 2) array of optical flow.
84
+ max_val (float): Maximum value of flow, values beyond
85
+ [-max_val, max_val] will be truncated.
86
+ norm (bool): Whether to divide flow values by image width/height.
87
+
88
+ Returns:
89
+ tuple[ndarray]: Quantized dx and dy.
90
+ """
91
+ h, w, _ = flow.shape
92
+ dx = flow[..., 0]
93
+ dy = flow[..., 1]
94
+ if norm:
95
+ dx = dx / w # avoid inplace operations
96
+ dy = dy / h
97
+ # use 255 levels instead of 256 to make sure 0 is 0 after dequantization.
98
+ flow_comps = [quantize(d, -max_val, max_val, 255, np.uint8) for d in [dx, dy]]
99
+ return tuple(flow_comps)
100
+
101
+
102
+ def dequantize_flow(dx, dy, max_val=0.02, denorm=True):
103
+ """Recover from quantized flow.
104
+
105
+ Args:
106
+ dx (ndarray): Quantized dx.
107
+ dy (ndarray): Quantized dy.
108
+ max_val (float): Maximum value used when quantizing.
109
+ denorm (bool): Whether to multiply flow values with width/height.
110
+
111
+ Returns:
112
+ ndarray: Dequantized flow.
113
+ """
114
+ assert dx.shape == dy.shape
115
+ assert dx.ndim == 2 or (dx.ndim == 3 and dx.shape[-1] == 1)
116
+
117
+ dx, dy = [dequantize(d, -max_val, max_val, 255) for d in [dx, dy]]
118
+
119
+ if denorm:
120
+ dx *= dx.shape[1]
121
+ dy *= dx.shape[0]
122
+ flow = np.dstack((dx, dy))
123
+ return flow
124
+
125
+
126
+ def quantize(arr, min_val, max_val, levels, dtype=np.int64):
127
+ """Quantize an array of (-inf, inf) to [0, levels-1].
128
+
129
+ Args:
130
+ arr (ndarray): Input array.
131
+ min_val (scalar): Minimum value to be clipped.
132
+ max_val (scalar): Maximum value to be clipped.
133
+ levels (int): Quantization levels.
134
+ dtype (np.type): The type of the quantized array.
135
+
136
+ Returns:
137
+ tuple: Quantized array.
138
+ """
139
+ if not (isinstance(levels, int) and levels > 1):
140
+ raise ValueError(f'levels must be a positive integer, but got {levels}')
141
+ if min_val >= max_val:
142
+ raise ValueError(f'min_val ({min_val}) must be smaller than max_val ({max_val})')
143
+
144
+ arr = np.clip(arr, min_val, max_val) - min_val
145
+ quantized_arr = np.minimum(np.floor(levels * arr / (max_val - min_val)).astype(dtype), levels - 1)
146
+
147
+ return quantized_arr
148
+
149
+
150
+ def dequantize(arr, min_val, max_val, levels, dtype=np.float64):
151
+ """Dequantize an array.
152
+
153
+ Args:
154
+ arr (ndarray): Input array.
155
+ min_val (scalar): Minimum value to be clipped.
156
+ max_val (scalar): Maximum value to be clipped.
157
+ levels (int): Quantization levels.
158
+ dtype (np.type): The type of the dequantized array.
159
+
160
+ Returns:
161
+ tuple: Dequantized array.
162
+ """
163
+ if not (isinstance(levels, int) and levels > 1):
164
+ raise ValueError(f'levels must be a positive integer, but got {levels}')
165
+ if min_val >= max_val:
166
+ raise ValueError(f'min_val ({min_val}) must be smaller than max_val ({max_val})')
167
+
168
+ dequantized_arr = (arr + 0.5).astype(dtype) * (max_val - min_val) / levels + min_val
169
+
170
+ return dequantized_arr
official/bsr/utils/img_process_util.py ADDED
@@ -0,0 +1,83 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import numpy as np
3
+ import torch
4
+ from torch.nn import functional as F
5
+
6
+
7
+ def filter2D(img, kernel):
8
+ """PyTorch version of cv2.filter2D
9
+
10
+ Args:
11
+ img (Tensor): (b, c, h, w)
12
+ kernel (Tensor): (b, k, k)
13
+ """
14
+ k = kernel.size(-1)
15
+ b, c, h, w = img.size()
16
+ if k % 2 == 1:
17
+ img = F.pad(img, (k // 2, k // 2, k // 2, k // 2), mode='reflect')
18
+ else:
19
+ raise ValueError('Wrong kernel size')
20
+
21
+ ph, pw = img.size()[-2:]
22
+
23
+ if kernel.size(0) == 1:
24
+ # apply the same kernel to all batch images
25
+ img = img.view(b * c, 1, ph, pw)
26
+ kernel = kernel.view(1, 1, k, k)
27
+ return F.conv2d(img, kernel, padding=0).view(b, c, h, w)
28
+ else:
29
+ img = img.view(1, b * c, ph, pw)
30
+ kernel = kernel.view(b, 1, k, k).repeat(1, c, 1, 1).view(b * c, 1, k, k)
31
+ return F.conv2d(img, kernel, groups=b * c).view(b, c, h, w)
32
+
33
+
34
+ def usm_sharp(img, weight=0.5, radius=50, threshold=10):
35
+ """USM sharpening.
36
+
37
+ Input image: I; Blurry image: B.
38
+ 1. sharp = I + weight * (I - B)
39
+ 2. Mask = 1 if abs(I - B) > threshold, else: 0
40
+ 3. Blur mask:
41
+ 4. Out = Mask * sharp + (1 - Mask) * I
42
+
43
+
44
+ Args:
45
+ img (Numpy array): Input image, HWC, BGR; float32, [0, 1].
46
+ weight (float): Sharp weight. Default: 1.
47
+ radius (float): Kernel size of Gaussian blur. Default: 50.
48
+ threshold (int):
49
+ """
50
+ if radius % 2 == 0:
51
+ radius += 1
52
+ blur = cv2.GaussianBlur(img, (radius, radius), 0)
53
+ residual = img - blur
54
+ mask = np.abs(residual) * 255 > threshold
55
+ mask = mask.astype('float32')
56
+ soft_mask = cv2.GaussianBlur(mask, (radius, radius), 0)
57
+
58
+ sharp = img + weight * residual
59
+ sharp = np.clip(sharp, 0, 1)
60
+ return soft_mask * sharp + (1 - soft_mask) * img
61
+
62
+
63
+ class USMSharp(torch.nn.Module):
64
+
65
+ def __init__(self, radius=50, sigma=0):
66
+ super(USMSharp, self).__init__()
67
+ if radius % 2 == 0:
68
+ radius += 1
69
+ self.radius = radius
70
+ kernel = cv2.getGaussianKernel(radius, sigma)
71
+ kernel = torch.FloatTensor(np.dot(kernel, kernel.transpose())).unsqueeze_(0)
72
+ self.register_buffer('kernel', kernel)
73
+
74
+ def forward(self, img, weight=0.5, threshold=10):
75
+ blur = filter2D(img, self.kernel)
76
+ residual = img - blur
77
+
78
+ mask = torch.abs(residual) * 255 > threshold
79
+ mask = mask.float()
80
+ soft_mask = filter2D(mask, self.kernel)
81
+ sharp = img + weight * residual
82
+ sharp = torch.clip(sharp, 0, 1)
83
+ return soft_mask * sharp + (1 - soft_mask) * img
official/bsr/utils/img_util.py ADDED
@@ -0,0 +1,172 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import math
3
+ import numpy as np
4
+ import os
5
+ import torch
6
+ from torchvision.utils import make_grid
7
+
8
+
9
+ def img2tensor(imgs, bgr2rgb=True, float32=True):
10
+ """Numpy array to tensor.
11
+
12
+ Args:
13
+ imgs (list[ndarray] | ndarray): Input images.
14
+ bgr2rgb (bool): Whether to change bgr to rgb.
15
+ float32 (bool): Whether to change to float32.
16
+
17
+ Returns:
18
+ list[tensor] | tensor: Tensor images. If returned results only have
19
+ one element, just return tensor.
20
+ """
21
+
22
+ def _totensor(img, bgr2rgb, float32):
23
+ if img.shape[2] == 3 and bgr2rgb:
24
+ if img.dtype == 'float64':
25
+ img = img.astype('float32')
26
+ img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
27
+ img = torch.from_numpy(img.transpose(2, 0, 1))
28
+ if float32:
29
+ img = img.float()
30
+ return img
31
+
32
+ if isinstance(imgs, list):
33
+ return [_totensor(img, bgr2rgb, float32) for img in imgs]
34
+ else:
35
+ return _totensor(imgs, bgr2rgb, float32)
36
+
37
+
38
+ def tensor2img(tensor, rgb2bgr=True, out_type=np.uint8, min_max=(0, 1)):
39
+ """Convert torch Tensors into image numpy arrays.
40
+
41
+ After clamping to [min, max], values will be normalized to [0, 1].
42
+
43
+ Args:
44
+ tensor (Tensor or list[Tensor]): Accept shapes:
45
+ 1) 4D mini-batch Tensor of shape (B x 3/1 x H x W);
46
+ 2) 3D Tensor of shape (3/1 x H x W);
47
+ 3) 2D Tensor of shape (H x W).
48
+ Tensor channel should be in RGB order.
49
+ rgb2bgr (bool): Whether to change rgb to bgr.
50
+ out_type (numpy type): output types. If ``np.uint8``, transform outputs
51
+ to uint8 type with range [0, 255]; otherwise, float type with
52
+ range [0, 1]. Default: ``np.uint8``.
53
+ min_max (tuple[int]): min and max values for clamp.
54
+
55
+ Returns:
56
+ (Tensor or list): 3D ndarray of shape (H x W x C) OR 2D ndarray of
57
+ shape (H x W). The channel order is BGR.
58
+ """
59
+ if not (torch.is_tensor(tensor) or (isinstance(tensor, list) and all(torch.is_tensor(t) for t in tensor))):
60
+ raise TypeError(f'tensor or list of tensors expected, got {type(tensor)}')
61
+
62
+ if torch.is_tensor(tensor):
63
+ tensor = [tensor]
64
+ result = []
65
+ for _tensor in tensor:
66
+ _tensor = _tensor.squeeze(0).float().detach().cpu().clamp_(*min_max)
67
+ _tensor = (_tensor - min_max[0]) / (min_max[1] - min_max[0])
68
+
69
+ n_dim = _tensor.dim()
70
+ if n_dim == 4:
71
+ img_np = make_grid(_tensor, nrow=int(math.sqrt(_tensor.size(0))), normalize=False).numpy()
72
+ img_np = img_np.transpose(1, 2, 0)
73
+ if rgb2bgr:
74
+ img_np = cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR)
75
+ elif n_dim == 3:
76
+ img_np = _tensor.numpy()
77
+ img_np = img_np.transpose(1, 2, 0)
78
+ if img_np.shape[2] == 1: # gray image
79
+ img_np = np.squeeze(img_np, axis=2)
80
+ else:
81
+ if rgb2bgr:
82
+ img_np = cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR)
83
+ elif n_dim == 2:
84
+ img_np = _tensor.numpy()
85
+ else:
86
+ raise TypeError(f'Only support 4D, 3D or 2D tensor. But received with dimension: {n_dim}')
87
+ if out_type == np.uint8:
88
+ # Unlike MATLAB, numpy.unit8() WILL NOT round by default.
89
+ img_np = (img_np * 255.0).round()
90
+ img_np = img_np.astype(out_type)
91
+ result.append(img_np)
92
+ if len(result) == 1:
93
+ result = result[0]
94
+ return result
95
+
96
+
97
+ def tensor2img_fast(tensor, rgb2bgr=True, min_max=(0, 1)):
98
+ """This implementation is slightly faster than tensor2img.
99
+ It now only supports torch tensor with shape (1, c, h, w).
100
+
101
+ Args:
102
+ tensor (Tensor): Now only support torch tensor with (1, c, h, w).
103
+ rgb2bgr (bool): Whether to change rgb to bgr. Default: True.
104
+ min_max (tuple[int]): min and max values for clamp.
105
+ """
106
+ output = tensor.squeeze(0).detach().clamp_(*min_max).permute(1, 2, 0)
107
+ output = (output - min_max[0]) / (min_max[1] - min_max[0]) * 255
108
+ output = output.type(torch.uint8).cpu().numpy()
109
+ if rgb2bgr:
110
+ output = cv2.cvtColor(output, cv2.COLOR_RGB2BGR)
111
+ return output
112
+
113
+
114
+ def imfrombytes(content, flag='color', float32=False):
115
+ """Read an image from bytes.
116
+
117
+ Args:
118
+ content (bytes): Image bytes got from files or other streams.
119
+ flag (str): Flags specifying the color type of a loaded image,
120
+ candidates are `color`, `grayscale` and `unchanged`.
121
+ float32 (bool): Whether to change to float32., If True, will also norm
122
+ to [0, 1]. Default: False.
123
+
124
+ Returns:
125
+ ndarray: Loaded image array.
126
+ """
127
+ img_np = np.frombuffer(content, np.uint8)
128
+ imread_flags = {'color': cv2.IMREAD_COLOR, 'grayscale': cv2.IMREAD_GRAYSCALE, 'unchanged': cv2.IMREAD_UNCHANGED}
129
+ img = cv2.imdecode(img_np, imread_flags[flag])
130
+ if float32:
131
+ img = img.astype(np.float32) / 255.
132
+ return img
133
+
134
+
135
+ def imwrite(img, file_path, params=None, auto_mkdir=True):
136
+ """Write image to file.
137
+
138
+ Args:
139
+ img (ndarray): Image array to be written.
140
+ file_path (str): Image file path.
141
+ params (None or list): Same as opencv's :func:`imwrite` interface.
142
+ auto_mkdir (bool): If the parent folder of `file_path` does not exist,
143
+ whether to create it automatically.
144
+
145
+ Returns:
146
+ bool: Successful or not.
147
+ """
148
+ if auto_mkdir:
149
+ dir_name = os.path.abspath(os.path.dirname(file_path))
150
+ os.makedirs(dir_name, exist_ok=True)
151
+ ok = cv2.imwrite(file_path, img, params)
152
+ if not ok:
153
+ raise IOError('Failed in writing images.')
154
+
155
+
156
+ def crop_border(imgs, crop_border):
157
+ """Crop borders of images.
158
+
159
+ Args:
160
+ imgs (list[ndarray] | ndarray): Images with shape (h, w, c).
161
+ crop_border (int): Crop border for each end of height and weight.
162
+
163
+ Returns:
164
+ list[ndarray]: Cropped images.
165
+ """
166
+ if crop_border == 0:
167
+ return imgs
168
+ else:
169
+ if isinstance(imgs, list):
170
+ return [v[crop_border:-crop_border, crop_border:-crop_border, ...] for v in imgs]
171
+ else:
172
+ return imgs[crop_border:-crop_border, crop_border:-crop_border, ...]
official/bsr/utils/lmdb_util.py ADDED
@@ -0,0 +1,199 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import lmdb
3
+ import sys
4
+ from multiprocessing import Pool
5
+ from os import path as osp
6
+ from tqdm import tqdm
7
+
8
+
9
+ def make_lmdb_from_imgs(data_path,
10
+ lmdb_path,
11
+ img_path_list,
12
+ keys,
13
+ batch=5000,
14
+ compress_level=1,
15
+ multiprocessing_read=False,
16
+ n_thread=40,
17
+ map_size=None):
18
+ """Make lmdb from images.
19
+
20
+ Contents of lmdb. The file structure is:
21
+
22
+ ::
23
+
24
+ example.lmdb
25
+ ├── data.mdb
26
+ ├── lock.mdb
27
+ ├── meta_info.txt
28
+
29
+ The data.mdb and lock.mdb are standard lmdb files and you can refer to
30
+ https://lmdb.readthedocs.io/en/release/ for more details.
31
+
32
+ The meta_info.txt is a specified txt file to record the meta information
33
+ of our datasets. It will be automatically created when preparing
34
+ datasets by our provided dataset tools.
35
+ Each line in the txt file records 1)image name (with extension),
36
+ 2)image shape, and 3)compression level, separated by a white space.
37
+
38
+ For example, the meta information could be:
39
+ `000_00000000.png (720,1280,3) 1`, which means:
40
+ 1) image name (with extension): 000_00000000.png;
41
+ 2) image shape: (720,1280,3);
42
+ 3) compression level: 1
43
+
44
+ We use the image name without extension as the lmdb key.
45
+
46
+ If `multiprocessing_read` is True, it will read all the images to memory
47
+ using multiprocessing. Thus, your server needs to have enough memory.
48
+
49
+ Args:
50
+ data_path (str): Data path for reading images.
51
+ lmdb_path (str): Lmdb save path.
52
+ img_path_list (str): Image path list.
53
+ keys (str): Used for lmdb keys.
54
+ batch (int): After processing batch images, lmdb commits.
55
+ Default: 5000.
56
+ compress_level (int): Compress level when encoding images. Default: 1.
57
+ multiprocessing_read (bool): Whether use multiprocessing to read all
58
+ the images to memory. Default: False.
59
+ n_thread (int): For multiprocessing.
60
+ map_size (int | None): Map size for lmdb env. If None, use the
61
+ estimated size from images. Default: None
62
+ """
63
+
64
+ assert len(img_path_list) == len(keys), ('img_path_list and keys should have the same length, '
65
+ f'but got {len(img_path_list)} and {len(keys)}')
66
+ print(f'Create lmdb for {data_path}, save to {lmdb_path}...')
67
+ print(f'Totoal images: {len(img_path_list)}')
68
+ if not lmdb_path.endswith('.lmdb'):
69
+ raise ValueError("lmdb_path must end with '.lmdb'.")
70
+ if osp.exists(lmdb_path):
71
+ print(f'Folder {lmdb_path} already exists. Exit.')
72
+ sys.exit(1)
73
+
74
+ if multiprocessing_read:
75
+ # read all the images to memory (multiprocessing)
76
+ dataset = {} # use dict to keep the order for multiprocessing
77
+ shapes = {}
78
+ print(f'Read images with multiprocessing, #thread: {n_thread} ...')
79
+ pbar = tqdm(total=len(img_path_list), unit='image')
80
+
81
+ def callback(arg):
82
+ """get the image data and update pbar."""
83
+ key, dataset[key], shapes[key] = arg
84
+ pbar.update(1)
85
+ pbar.set_description(f'Read {key}')
86
+
87
+ pool = Pool(n_thread)
88
+ for path, key in zip(img_path_list, keys):
89
+ pool.apply_async(read_img_worker, args=(osp.join(data_path, path), key, compress_level), callback=callback)
90
+ pool.close()
91
+ pool.join()
92
+ pbar.close()
93
+ print(f'Finish reading {len(img_path_list)} images.')
94
+
95
+ # create lmdb environment
96
+ if map_size is None:
97
+ # obtain data size for one image
98
+ img = cv2.imread(osp.join(data_path, img_path_list[0]), cv2.IMREAD_UNCHANGED)
99
+ _, img_byte = cv2.imencode('.png', img, [cv2.IMWRITE_PNG_COMPRESSION, compress_level])
100
+ data_size_per_img = img_byte.nbytes
101
+ print('Data size per image is: ', data_size_per_img)
102
+ data_size = data_size_per_img * len(img_path_list)
103
+ map_size = data_size * 10
104
+
105
+ env = lmdb.open(lmdb_path, map_size=map_size)
106
+
107
+ # write data to lmdb
108
+ pbar = tqdm(total=len(img_path_list), unit='chunk')
109
+ txn = env.begin(write=True)
110
+ txt_file = open(osp.join(lmdb_path, 'meta_info.txt'), 'w')
111
+ for idx, (path, key) in enumerate(zip(img_path_list, keys)):
112
+ pbar.update(1)
113
+ pbar.set_description(f'Write {key}')
114
+ key_byte = key.encode('ascii')
115
+ if multiprocessing_read:
116
+ img_byte = dataset[key]
117
+ h, w, c = shapes[key]
118
+ else:
119
+ _, img_byte, img_shape = read_img_worker(osp.join(data_path, path), key, compress_level)
120
+ h, w, c = img_shape
121
+
122
+ txn.put(key_byte, img_byte)
123
+ # write meta information
124
+ txt_file.write(f'{key}.png ({h},{w},{c}) {compress_level}\n')
125
+ if idx % batch == 0:
126
+ txn.commit()
127
+ txn = env.begin(write=True)
128
+ pbar.close()
129
+ txn.commit()
130
+ env.close()
131
+ txt_file.close()
132
+ print('\nFinish writing lmdb.')
133
+
134
+
135
+ def read_img_worker(path, key, compress_level):
136
+ """Read image worker.
137
+
138
+ Args:
139
+ path (str): Image path.
140
+ key (str): Image key.
141
+ compress_level (int): Compress level when encoding images.
142
+
143
+ Returns:
144
+ str: Image key.
145
+ byte: Image byte.
146
+ tuple[int]: Image shape.
147
+ """
148
+
149
+ img = cv2.imread(path, cv2.IMREAD_UNCHANGED)
150
+ if img.ndim == 2:
151
+ h, w = img.shape
152
+ c = 1
153
+ else:
154
+ h, w, c = img.shape
155
+ _, img_byte = cv2.imencode('.png', img, [cv2.IMWRITE_PNG_COMPRESSION, compress_level])
156
+ return (key, img_byte, (h, w, c))
157
+
158
+
159
+ class LmdbMaker():
160
+ """LMDB Maker.
161
+
162
+ Args:
163
+ lmdb_path (str): Lmdb save path.
164
+ map_size (int): Map size for lmdb env. Default: 1024 ** 4, 1TB.
165
+ batch (int): After processing batch images, lmdb commits.
166
+ Default: 5000.
167
+ compress_level (int): Compress level when encoding images. Default: 1.
168
+ """
169
+
170
+ def __init__(self, lmdb_path, map_size=1024**4, batch=5000, compress_level=1):
171
+ if not lmdb_path.endswith('.lmdb'):
172
+ raise ValueError("lmdb_path must end with '.lmdb'.")
173
+ if osp.exists(lmdb_path):
174
+ print(f'Folder {lmdb_path} already exists. Exit.')
175
+ sys.exit(1)
176
+
177
+ self.lmdb_path = lmdb_path
178
+ self.batch = batch
179
+ self.compress_level = compress_level
180
+ self.env = lmdb.open(lmdb_path, map_size=map_size)
181
+ self.txn = self.env.begin(write=True)
182
+ self.txt_file = open(osp.join(lmdb_path, 'meta_info.txt'), 'w')
183
+ self.counter = 0
184
+
185
+ def put(self, img_byte, key, img_shape):
186
+ self.counter += 1
187
+ key_byte = key.encode('ascii')
188
+ self.txn.put(key_byte, img_byte)
189
+ # write meta information
190
+ h, w, c = img_shape
191
+ self.txt_file.write(f'{key}.png ({h},{w},{c}) {self.compress_level}\n')
192
+ if self.counter % self.batch == 0:
193
+ self.txn.commit()
194
+ self.txn = self.env.begin(write=True)
195
+
196
+ def close(self):
197
+ self.txn.commit()
198
+ self.env.close()
199
+ self.txt_file.close()
official/bsr/utils/logger.py ADDED
@@ -0,0 +1,213 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import datetime
2
+ import logging
3
+ import time
4
+
5
+ from .dist_util import get_dist_info, master_only
6
+
7
+ initialized_logger = {}
8
+
9
+
10
+ class AvgTimer():
11
+
12
+ def __init__(self, window=200):
13
+ self.window = window # average window
14
+ self.current_time = 0
15
+ self.total_time = 0
16
+ self.count = 0
17
+ self.avg_time = 0
18
+ self.start()
19
+
20
+ def start(self):
21
+ self.start_time = self.tic = time.time()
22
+
23
+ def record(self):
24
+ self.count += 1
25
+ self.toc = time.time()
26
+ self.current_time = self.toc - self.tic
27
+ self.total_time += self.current_time
28
+ # calculate average time
29
+ self.avg_time = self.total_time / self.count
30
+
31
+ # reset
32
+ if self.count > self.window:
33
+ self.count = 0
34
+ self.total_time = 0
35
+
36
+ self.tic = time.time()
37
+
38
+ def get_current_time(self):
39
+ return self.current_time
40
+
41
+ def get_avg_time(self):
42
+ return self.avg_time
43
+
44
+
45
+ class MessageLogger():
46
+ """Message logger for printing.
47
+
48
+ Args:
49
+ opt (dict): Config. It contains the following keys:
50
+ name (str): Exp name.
51
+ logger (dict): Contains 'print_freq' (str) for logger interval.
52
+ train (dict): Contains 'total_iter' (int) for total iters.
53
+ use_tb_logger (bool): Use tensorboard logger.
54
+ start_iter (int): Start iter. Default: 1.
55
+ tb_logger (obj:`tb_logger`): Tensorboard logger. Default: None.
56
+ """
57
+
58
+ def __init__(self, opt, start_iter=1, tb_logger=None):
59
+ self.exp_name = opt['name']
60
+ self.interval = opt['logger']['print_freq']
61
+ self.start_iter = start_iter
62
+ self.max_iters = opt['train']['total_iter']
63
+ self.use_tb_logger = opt['logger']['use_tb_logger']
64
+ self.tb_logger = tb_logger
65
+ self.start_time = time.time()
66
+ self.logger = get_root_logger()
67
+
68
+ def reset_start_time(self):
69
+ self.start_time = time.time()
70
+
71
+ @master_only
72
+ def __call__(self, log_vars):
73
+ """Format logging message.
74
+
75
+ Args:
76
+ log_vars (dict): It contains the following keys:
77
+ epoch (int): Epoch number.
78
+ iter (int): Current iter.
79
+ lrs (list): List for learning rates.
80
+
81
+ time (float): Iter time.
82
+ data_time (float): Data time for each iter.
83
+ """
84
+ # epoch, iter, learning rates
85
+ epoch = log_vars.pop('epoch')
86
+ current_iter = log_vars.pop('iter')
87
+ lrs = log_vars.pop('lrs')
88
+
89
+ message = (f'[{self.exp_name[:5]}..][epoch:{epoch:3d}, iter:{current_iter:8,d}, lr:(')
90
+ for v in lrs:
91
+ message += f'{v:.3e},'
92
+ message += ')] '
93
+
94
+ # time and estimated time
95
+ if 'time' in log_vars.keys():
96
+ iter_time = log_vars.pop('time')
97
+ data_time = log_vars.pop('data_time')
98
+
99
+ total_time = time.time() - self.start_time
100
+ time_sec_avg = total_time / (current_iter - self.start_iter + 1)
101
+ eta_sec = time_sec_avg * (self.max_iters - current_iter - 1)
102
+ eta_str = str(datetime.timedelta(seconds=int(eta_sec)))
103
+ message += f'[eta: {eta_str}, '
104
+ message += f'time (data): {iter_time:.3f} ({data_time:.3f})] '
105
+
106
+ # other items, especially losses
107
+ for k, v in log_vars.items():
108
+ message += f'{k}: {v:.4e} '
109
+ # tensorboard logger
110
+ if self.use_tb_logger and 'debug' not in self.exp_name:
111
+ if k.startswith('l_'):
112
+ self.tb_logger.add_scalar(f'losses/{k}', v, current_iter)
113
+ else:
114
+ self.tb_logger.add_scalar(k, v, current_iter)
115
+ self.logger.info(message)
116
+
117
+
118
+ @master_only
119
+ def init_tb_logger(log_dir):
120
+ from torch.utils.tensorboard import SummaryWriter
121
+ tb_logger = SummaryWriter(log_dir=log_dir)
122
+ return tb_logger
123
+
124
+
125
+ @master_only
126
+ def init_wandb_logger(opt):
127
+ """We now only use wandb to sync tensorboard log."""
128
+ import wandb
129
+ logger = get_root_logger()
130
+
131
+ project = opt['logger']['wandb']['project']
132
+ resume_id = opt['logger']['wandb'].get('resume_id')
133
+ if resume_id:
134
+ wandb_id = resume_id
135
+ resume = 'allow'
136
+ logger.warning(f'Resume wandb logger with id={wandb_id}.')
137
+ else:
138
+ wandb_id = wandb.util.generate_id()
139
+ resume = 'never'
140
+
141
+ wandb.init(id=wandb_id, resume=resume, name=opt['name'], config=opt, project=project, sync_tensorboard=True)
142
+
143
+ logger.info(f'Use wandb logger with id={wandb_id}; project={project}.')
144
+
145
+
146
+ def get_root_logger(logger_name='basicsr', log_level=logging.INFO, log_file=None):
147
+ """Get the root logger.
148
+
149
+ The logger will be initialized if it has not been initialized. By default a
150
+ StreamHandler will be added. If `log_file` is specified, a FileHandler will
151
+ also be added.
152
+
153
+ Args:
154
+ logger_name (str): root logger name. Default: 'basicsr'.
155
+ log_file (str | None): The log filename. If specified, a FileHandler
156
+ will be added to the root logger.
157
+ log_level (int): The root logger level. Note that only the process of
158
+ rank 0 is affected, while other processes will set the level to
159
+ "Error" and be silent most of the time.
160
+
161
+ Returns:
162
+ logging.Logger: The root logger.
163
+ """
164
+ logger = logging.getLogger(logger_name)
165
+ # if the logger has been initialized, just return it
166
+ if logger_name in initialized_logger:
167
+ return logger
168
+
169
+ format_str = '%(asctime)s %(levelname)s: %(message)s'
170
+ stream_handler = logging.StreamHandler()
171
+ stream_handler.setFormatter(logging.Formatter(format_str))
172
+ logger.addHandler(stream_handler)
173
+ logger.propagate = False
174
+ rank, _ = get_dist_info()
175
+ if rank != 0:
176
+ logger.setLevel('ERROR')
177
+ elif log_file is not None:
178
+ logger.setLevel(log_level)
179
+ # add file handler
180
+ file_handler = logging.FileHandler(log_file, 'w')
181
+ file_handler.setFormatter(logging.Formatter(format_str))
182
+ file_handler.setLevel(log_level)
183
+ logger.addHandler(file_handler)
184
+ initialized_logger[logger_name] = True
185
+ return logger
186
+
187
+
188
+ def get_env_info():
189
+ """Get environment information.
190
+
191
+ Currently, only log the software version.
192
+ """
193
+ import torch
194
+ import torchvision
195
+
196
+ from basicsr.version import __version__
197
+ msg = r"""
198
+ ____ _ _____ ____
199
+ / __ ) ____ _ _____ (_)_____/ ___/ / __ \
200
+ / __ |/ __ `// ___// // ___/\__ \ / /_/ /
201
+ / /_/ // /_/ /(__ )/ // /__ ___/ // _, _/
202
+ /_____/ \__,_//____//_/ \___//____//_/ |_|
203
+ ______ __ __ __ __
204
+ / ____/____ ____ ____/ / / / __ __ _____ / /__ / /
205
+ / / __ / __ \ / __ \ / __ / / / / / / // ___// //_/ / /
206
+ / /_/ // /_/ // /_/ // /_/ / / /___/ /_/ // /__ / /< /_/
207
+ \____/ \____/ \____/ \____/ /_____/\____/ \___//_/|_| (_)
208
+ """
209
+ msg += ('\nVersion Information: '
210
+ f'\n\tBasicSR: {__version__}'
211
+ f'\n\tPyTorch: {torch.__version__}'
212
+ f'\n\tTorchVision: {torchvision.__version__}')
213
+ return msg
official/bsr/utils/matlab_functions.py ADDED
@@ -0,0 +1,178 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ import numpy as np
3
+ import torch
4
+
5
+
6
+ def cubic(x):
7
+ """cubic function used for calculate_weights_indices."""
8
+ absx = torch.abs(x)
9
+ absx2 = absx**2
10
+ absx3 = absx**3
11
+ return (1.5 * absx3 - 2.5 * absx2 + 1) * (
12
+ (absx <= 1).type_as(absx)) + (-0.5 * absx3 + 2.5 * absx2 - 4 * absx + 2) * (((absx > 1) *
13
+ (absx <= 2)).type_as(absx))
14
+
15
+
16
+ def calculate_weights_indices(in_length, out_length, scale, kernel, kernel_width, antialiasing):
17
+ """Calculate weights and indices, used for imresize function.
18
+
19
+ Args:
20
+ in_length (int): Input length.
21
+ out_length (int): Output length.
22
+ scale (float): Scale factor.
23
+ kernel_width (int): Kernel width.
24
+ antialisaing (bool): Whether to apply anti-aliasing when downsampling.
25
+ """
26
+
27
+ if (scale < 1) and antialiasing:
28
+ # Use a modified kernel (larger kernel width) to simultaneously
29
+ # interpolate and antialias
30
+ kernel_width = kernel_width / scale
31
+
32
+ # Output-space coordinates
33
+ x = torch.linspace(1, out_length, out_length)
34
+
35
+ # Input-space coordinates. Calculate the inverse mapping such that 0.5
36
+ # in output space maps to 0.5 in input space, and 0.5 + scale in output
37
+ # space maps to 1.5 in input space.
38
+ u = x / scale + 0.5 * (1 - 1 / scale)
39
+
40
+ # What is the left-most pixel that can be involved in the computation?
41
+ left = torch.floor(u - kernel_width / 2)
42
+
43
+ # What is the maximum number of pixels that can be involved in the
44
+ # computation? Note: it's OK to use an extra pixel here; if the
45
+ # corresponding weights are all zero, it will be eliminated at the end
46
+ # of this function.
47
+ p = math.ceil(kernel_width) + 2
48
+
49
+ # The indices of the input pixels involved in computing the k-th output
50
+ # pixel are in row k of the indices matrix.
51
+ indices = left.view(out_length, 1).expand(out_length, p) + torch.linspace(0, p - 1, p).view(1, p).expand(
52
+ out_length, p)
53
+
54
+ # The weights used to compute the k-th output pixel are in row k of the
55
+ # weights matrix.
56
+ distance_to_center = u.view(out_length, 1).expand(out_length, p) - indices
57
+
58
+ # apply cubic kernel
59
+ if (scale < 1) and antialiasing:
60
+ weights = scale * cubic(distance_to_center * scale)
61
+ else:
62
+ weights = cubic(distance_to_center)
63
+
64
+ # Normalize the weights matrix so that each row sums to 1.
65
+ weights_sum = torch.sum(weights, 1).view(out_length, 1)
66
+ weights = weights / weights_sum.expand(out_length, p)
67
+
68
+ # If a column in weights is all zero, get rid of it. only consider the
69
+ # first and last column.
70
+ weights_zero_tmp = torch.sum((weights == 0), 0)
71
+ if not math.isclose(weights_zero_tmp[0], 0, rel_tol=1e-6):
72
+ indices = indices.narrow(1, 1, p - 2)
73
+ weights = weights.narrow(1, 1, p - 2)
74
+ if not math.isclose(weights_zero_tmp[-1], 0, rel_tol=1e-6):
75
+ indices = indices.narrow(1, 0, p - 2)
76
+ weights = weights.narrow(1, 0, p - 2)
77
+ weights = weights.contiguous()
78
+ indices = indices.contiguous()
79
+ sym_len_s = -indices.min() + 1
80
+ sym_len_e = indices.max() - in_length
81
+ indices = indices + sym_len_s - 1
82
+ return weights, indices, int(sym_len_s), int(sym_len_e)
83
+
84
+
85
+ @torch.no_grad()
86
+ def imresize(img, scale, antialiasing=True):
87
+ """imresize function same as MATLAB.
88
+
89
+ It now only supports bicubic.
90
+ The same scale applies for both height and width.
91
+
92
+ Args:
93
+ img (Tensor | Numpy array):
94
+ Tensor: Input image with shape (c, h, w), [0, 1] range.
95
+ Numpy: Input image with shape (h, w, c), [0, 1] range.
96
+ scale (float): Scale factor. The same scale applies for both height
97
+ and width.
98
+ antialisaing (bool): Whether to apply anti-aliasing when downsampling.
99
+ Default: True.
100
+
101
+ Returns:
102
+ Tensor: Output image with shape (c, h, w), [0, 1] range, w/o round.
103
+ """
104
+ squeeze_flag = False
105
+ if type(img).__module__ == np.__name__: # numpy type
106
+ numpy_type = True
107
+ if img.ndim == 2:
108
+ img = img[:, :, None]
109
+ squeeze_flag = True
110
+ img = torch.from_numpy(img.transpose(2, 0, 1)).float()
111
+ else:
112
+ numpy_type = False
113
+ if img.ndim == 2:
114
+ img = img.unsqueeze(0)
115
+ squeeze_flag = True
116
+
117
+ in_c, in_h, in_w = img.size()
118
+ out_h, out_w = math.ceil(in_h * scale), math.ceil(in_w * scale)
119
+ kernel_width = 4
120
+ kernel = 'cubic'
121
+
122
+ # get weights and indices
123
+ weights_h, indices_h, sym_len_hs, sym_len_he = calculate_weights_indices(in_h, out_h, scale, kernel, kernel_width,
124
+ antialiasing)
125
+ weights_w, indices_w, sym_len_ws, sym_len_we = calculate_weights_indices(in_w, out_w, scale, kernel, kernel_width,
126
+ antialiasing)
127
+ # process H dimension
128
+ # symmetric copying
129
+ img_aug = torch.FloatTensor(in_c, in_h + sym_len_hs + sym_len_he, in_w)
130
+ img_aug.narrow(1, sym_len_hs, in_h).copy_(img)
131
+
132
+ sym_patch = img[:, :sym_len_hs, :]
133
+ inv_idx = torch.arange(sym_patch.size(1) - 1, -1, -1).long()
134
+ sym_patch_inv = sym_patch.index_select(1, inv_idx)
135
+ img_aug.narrow(1, 0, sym_len_hs).copy_(sym_patch_inv)
136
+
137
+ sym_patch = img[:, -sym_len_he:, :]
138
+ inv_idx = torch.arange(sym_patch.size(1) - 1, -1, -1).long()
139
+ sym_patch_inv = sym_patch.index_select(1, inv_idx)
140
+ img_aug.narrow(1, sym_len_hs + in_h, sym_len_he).copy_(sym_patch_inv)
141
+
142
+ out_1 = torch.FloatTensor(in_c, out_h, in_w)
143
+ kernel_width = weights_h.size(1)
144
+ for i in range(out_h):
145
+ idx = int(indices_h[i][0])
146
+ for j in range(in_c):
147
+ out_1[j, i, :] = img_aug[j, idx:idx + kernel_width, :].transpose(0, 1).mv(weights_h[i])
148
+
149
+ # process W dimension
150
+ # symmetric copying
151
+ out_1_aug = torch.FloatTensor(in_c, out_h, in_w + sym_len_ws + sym_len_we)
152
+ out_1_aug.narrow(2, sym_len_ws, in_w).copy_(out_1)
153
+
154
+ sym_patch = out_1[:, :, :sym_len_ws]
155
+ inv_idx = torch.arange(sym_patch.size(2) - 1, -1, -1).long()
156
+ sym_patch_inv = sym_patch.index_select(2, inv_idx)
157
+ out_1_aug.narrow(2, 0, sym_len_ws).copy_(sym_patch_inv)
158
+
159
+ sym_patch = out_1[:, :, -sym_len_we:]
160
+ inv_idx = torch.arange(sym_patch.size(2) - 1, -1, -1).long()
161
+ sym_patch_inv = sym_patch.index_select(2, inv_idx)
162
+ out_1_aug.narrow(2, sym_len_ws + in_w, sym_len_we).copy_(sym_patch_inv)
163
+
164
+ out_2 = torch.FloatTensor(in_c, out_h, out_w)
165
+ kernel_width = weights_w.size(1)
166
+ for i in range(out_w):
167
+ idx = int(indices_w[i][0])
168
+ for j in range(in_c):
169
+ out_2[j, :, i] = out_1_aug[j, :, idx:idx + kernel_width].mv(weights_w[i])
170
+
171
+ if squeeze_flag:
172
+ out_2 = out_2.squeeze(0)
173
+ if numpy_type:
174
+ out_2 = out_2.numpy()
175
+ if not squeeze_flag:
176
+ out_2 = out_2.transpose(1, 2, 0)
177
+
178
+ return out_2
official/bsr/utils/misc.py ADDED
@@ -0,0 +1,141 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import os
3
+ import random
4
+ import time
5
+ import torch
6
+ from os import path as osp
7
+
8
+ from .dist_util import master_only
9
+
10
+
11
+ def set_random_seed(seed):
12
+ """Set random seeds."""
13
+ random.seed(seed)
14
+ np.random.seed(seed)
15
+ torch.manual_seed(seed)
16
+ torch.cuda.manual_seed(seed)
17
+ torch.cuda.manual_seed_all(seed)
18
+
19
+
20
+ def get_time_str():
21
+ return time.strftime('%Y%m%d_%H%M%S', time.localtime())
22
+
23
+
24
+ def mkdir_and_rename(path):
25
+ """mkdirs. If path exists, rename it with timestamp and create a new one.
26
+
27
+ Args:
28
+ path (str): Folder path.
29
+ """
30
+ if osp.exists(path):
31
+ new_name = path + '_archived_' + get_time_str()
32
+ print(f'Path already exists. Rename it to {new_name}', flush=True)
33
+ os.rename(path, new_name)
34
+ os.makedirs(path, exist_ok=True)
35
+
36
+
37
+ @master_only
38
+ def make_exp_dirs(opt):
39
+ """Make dirs for experiments."""
40
+ path_opt = opt['path'].copy()
41
+ if opt['is_train']:
42
+ mkdir_and_rename(path_opt.pop('experiments_root'))
43
+ else:
44
+ mkdir_and_rename(path_opt.pop('results_root'))
45
+ for key, path in path_opt.items():
46
+ if ('strict_load' in key) or ('pretrain_network' in key) or ('resume' in key) or ('param_key' in key):
47
+ continue
48
+ else:
49
+ os.makedirs(path, exist_ok=True)
50
+
51
+
52
+ def scandir(dir_path, suffix=None, recursive=False, full_path=False):
53
+ """Scan a directory to find the interested files.
54
+
55
+ Args:
56
+ dir_path (str): Path of the directory.
57
+ suffix (str | tuple(str), optional): File suffix that we are
58
+ interested in. Default: None.
59
+ recursive (bool, optional): If set to True, recursively scan the
60
+ directory. Default: False.
61
+ full_path (bool, optional): If set to True, include the dir_path.
62
+ Default: False.
63
+
64
+ Returns:
65
+ A generator for all the interested files with relative paths.
66
+ """
67
+
68
+ if (suffix is not None) and not isinstance(suffix, (str, tuple)):
69
+ raise TypeError('"suffix" must be a string or tuple of strings')
70
+
71
+ root = dir_path
72
+
73
+ def _scandir(dir_path, suffix, recursive):
74
+ for entry in os.scandir(dir_path):
75
+ if not entry.name.startswith('.') and entry.is_file():
76
+ if full_path:
77
+ return_path = entry.path
78
+ else:
79
+ return_path = osp.relpath(entry.path, root)
80
+
81
+ if suffix is None:
82
+ yield return_path
83
+ elif return_path.endswith(suffix):
84
+ yield return_path
85
+ else:
86
+ if recursive:
87
+ yield from _scandir(entry.path, suffix=suffix, recursive=recursive)
88
+ else:
89
+ continue
90
+
91
+ return _scandir(dir_path, suffix=suffix, recursive=recursive)
92
+
93
+
94
+ def check_resume(opt, resume_iter):
95
+ """Check resume states and pretrain_network paths.
96
+
97
+ Args:
98
+ opt (dict): Options.
99
+ resume_iter (int): Resume iteration.
100
+ """
101
+ if opt['path']['resume_state']:
102
+ # get all the networks
103
+ networks = [key for key in opt.keys() if key.startswith('network_')]
104
+ flag_pretrain = False
105
+ for network in networks:
106
+ if opt['path'].get(f'pretrain_{network}') is not None:
107
+ flag_pretrain = True
108
+ if flag_pretrain:
109
+ print('pretrain_network path will be ignored during resuming.')
110
+ # set pretrained model paths
111
+ for network in networks:
112
+ name = f'pretrain_{network}'
113
+ basename = network.replace('network_', '')
114
+ if opt['path'].get('ignore_resume_networks') is None or (network
115
+ not in opt['path']['ignore_resume_networks']):
116
+ opt['path'][name] = osp.join(opt['path']['models'], f'net_{basename}_{resume_iter}.pth')
117
+ print(f"Set {name} to {opt['path'][name]}")
118
+
119
+ # change param_key to params in resume
120
+ param_keys = [key for key in opt['path'].keys() if key.startswith('param_key')]
121
+ for param_key in param_keys:
122
+ if opt['path'][param_key] == 'params_ema':
123
+ opt['path'][param_key] = 'params'
124
+ print(f'Set {param_key} to params')
125
+
126
+
127
+ def sizeof_fmt(size, suffix='B'):
128
+ """Get human readable file size.
129
+
130
+ Args:
131
+ size (int): File size.
132
+ suffix (str): Suffix. Default: 'B'.
133
+
134
+ Return:
135
+ str: Formatted file size.
136
+ """
137
+ for unit in ['', 'K', 'M', 'G', 'T', 'P', 'E', 'Z']:
138
+ if abs(size) < 1024.0:
139
+ return f'{size:3.1f} {unit}{suffix}'
140
+ size /= 1024.0
141
+ return f'{size:3.1f} Y{suffix}'
official/bsr/utils/options.py ADDED
@@ -0,0 +1,218 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import os
3
+ import random
4
+ import torch
5
+ import yaml
6
+ from collections import OrderedDict
7
+ from os import path as osp
8
+
9
+ from bsr.utils import set_random_seed
10
+ from bsr.utils.dist_util import get_dist_info, init_dist, master_only
11
+
12
+
13
+ def ordered_yaml():
14
+ """Support OrderedDict for yaml.
15
+
16
+ Returns:
17
+ tuple: yaml Loader and Dumper.
18
+ """
19
+ try:
20
+ from yaml import CDumper as Dumper
21
+ from yaml import CLoader as Loader
22
+ except ImportError:
23
+ from yaml import Dumper, Loader
24
+
25
+ _mapping_tag = yaml.resolver.BaseResolver.DEFAULT_MAPPING_TAG
26
+
27
+ def dict_representer(dumper, data):
28
+ return dumper.represent_dict(data.items())
29
+
30
+ def dict_constructor(loader, node):
31
+ return OrderedDict(loader.construct_pairs(node))
32
+
33
+ Dumper.add_representer(OrderedDict, dict_representer)
34
+ Loader.add_constructor(_mapping_tag, dict_constructor)
35
+ return Loader, Dumper
36
+
37
+
38
+ def yaml_load(f):
39
+ """Load yaml file or string.
40
+
41
+ Args:
42
+ f (str): File path or a python string.
43
+
44
+ Returns:
45
+ dict: Loaded dict.
46
+ """
47
+ if os.path.isfile(f):
48
+ with open(f, 'r') as f:
49
+ return yaml.load(f, Loader=ordered_yaml()[0])
50
+ else:
51
+ return yaml.load(f, Loader=ordered_yaml()[0])
52
+
53
+
54
+ def dict2str(opt, indent_level=1):
55
+ """dict to string for printing options.
56
+
57
+ Args:
58
+ opt (dict): Option dict.
59
+ indent_level (int): Indent level. Default: 1.
60
+
61
+ Return:
62
+ (str): Option string for printing.
63
+ """
64
+ msg = '\n'
65
+ for k, v in opt.items():
66
+ if isinstance(v, dict):
67
+ msg += ' ' * (indent_level * 2) + k + ':['
68
+ msg += dict2str(v, indent_level + 1)
69
+ msg += ' ' * (indent_level * 2) + ']\n'
70
+ else:
71
+ msg += ' ' * (indent_level * 2) + k + ': ' + str(v) + '\n'
72
+ return msg
73
+
74
+
75
+ def _postprocess_yml_value(value):
76
+ # None
77
+ if value == '~' or value.lower() == 'none':
78
+ return None
79
+ # bool
80
+ if value.lower() == 'true':
81
+ return True
82
+ elif value.lower() == 'false':
83
+ return False
84
+ # !!float number
85
+ if value.startswith('!!float'):
86
+ return float(value.replace('!!float', ''))
87
+ # number
88
+ if value.isdigit():
89
+ return int(value)
90
+ elif value.replace('.', '', 1).isdigit() and value.count('.') < 2:
91
+ return float(value)
92
+ # list
93
+ if value.startswith('['):
94
+ return eval(value)
95
+ # str
96
+ return value
97
+
98
+
99
+ def parse_options(root_path, is_train=True):
100
+ parser = argparse.ArgumentParser()
101
+ parser.add_argument('-opt', type=str, required=True, help='Path to option YAML file.')
102
+ parser.add_argument('--launcher', choices=['none', 'pytorch', 'slurm'], default='none', help='job launcher')
103
+ parser.add_argument('--auto_resume', action='store_true')
104
+ parser.add_argument('--debug', action='store_true')
105
+ parser.add_argument('--local_rank', type=int, default=0)
106
+ parser.add_argument(
107
+ '--force_yml', nargs='+', default=None, help='Force to update yml files. Examples: train:ema_decay=0.999')
108
+ args = parser.parse_args()
109
+
110
+ # parse yml to dict
111
+ opt = yaml_load(args.opt)
112
+
113
+ # distributed settings
114
+ if args.launcher == 'none':
115
+ opt['dist'] = False
116
+ print('Disable distributed.', flush=True)
117
+ else:
118
+ opt['dist'] = True
119
+ if args.launcher == 'slurm' and 'dist_params' in opt:
120
+ init_dist(args.launcher, **opt['dist_params'])
121
+ else:
122
+ init_dist(args.launcher)
123
+ opt['rank'], opt['world_size'] = get_dist_info()
124
+
125
+ # random seed
126
+ seed = opt.get('manual_seed')
127
+ if seed is None:
128
+ seed = random.randint(1, 10000)
129
+ opt['manual_seed'] = seed
130
+ set_random_seed(seed + opt['rank'])
131
+
132
+ # force to update yml options
133
+ if args.force_yml is not None:
134
+ for entry in args.force_yml:
135
+ # now do not support creating new keys
136
+ keys, value = entry.split('=')
137
+ keys, value = keys.strip(), value.strip()
138
+ value = _postprocess_yml_value(value)
139
+ eval_str = 'opt'
140
+ for key in keys.split(':'):
141
+ eval_str += f'["{key}"]'
142
+ eval_str += '=value'
143
+ # using exec function
144
+ exec(eval_str)
145
+
146
+ opt['auto_resume'] = args.auto_resume
147
+ opt['is_train'] = is_train
148
+
149
+ # debug setting
150
+ if args.debug and not opt['name'].startswith('debug'):
151
+ opt['name'] = 'debug_' + opt['name']
152
+
153
+ if opt['num_gpu'] == 'auto':
154
+ opt['num_gpu'] = torch.cuda.device_count()
155
+
156
+ # datasets
157
+ for phase, dataset in opt['datasets'].items():
158
+ # for multiple datasets, e.g., val_1, val_2; test_1, test_2
159
+ phase = phase.split('_')[0]
160
+ dataset['phase'] = phase
161
+ if 'scale' in opt:
162
+ dataset['scale'] = opt['scale']
163
+ if dataset.get('dataroot_gt') is not None:
164
+ dataset['dataroot_gt'] = osp.expanduser(dataset['dataroot_gt'])
165
+ if dataset.get('dataroot_lq') is not None:
166
+ dataset['dataroot_lq'] = osp.expanduser(dataset['dataroot_lq'])
167
+
168
+ # paths
169
+ for key, val in opt['path'].items():
170
+ if (val is not None) and ('resume_state' in key or 'pretrain_network' in key):
171
+ opt['path'][key] = osp.expanduser(val)
172
+
173
+ if is_train:
174
+ experiments_root = opt['path'].get('experiments_root')
175
+ if experiments_root is None:
176
+ experiments_root = osp.join(root_path, 'experiments')
177
+ experiments_root = osp.join(experiments_root, opt['name'])
178
+
179
+ opt['path']['experiments_root'] = experiments_root
180
+ opt['path']['models'] = osp.join(experiments_root, 'models')
181
+ opt['path']['training_states'] = osp.join(experiments_root, 'training_states')
182
+ opt['path']['log'] = experiments_root
183
+ opt['path']['visualization'] = osp.join(experiments_root, 'visualization')
184
+
185
+ # change some options for debug mode
186
+ if 'debug' in opt['name']:
187
+ if 'val' in opt:
188
+ opt['val']['val_freq'] = 8
189
+ opt['logger']['print_freq'] = 1
190
+ opt['logger']['save_checkpoint_freq'] = 8
191
+ else: # test
192
+ results_root = opt['path'].get('results_root')
193
+ if results_root is None:
194
+ results_root = osp.join(root_path, 'results')
195
+ results_root = osp.join(results_root, opt['name'])
196
+
197
+ opt['path']['results_root'] = results_root
198
+ opt['path']['log'] = results_root
199
+ opt['path']['visualization'] = osp.join(results_root, 'visualization')
200
+
201
+ return opt, args
202
+
203
+
204
+ @master_only
205
+ def copy_opt_file(opt_file, experiments_root):
206
+ # copy the yml file to the experiment root
207
+ import sys
208
+ import time
209
+ from shutil import copyfile
210
+ cmd = ' '.join(sys.argv)
211
+ filename = osp.join(experiments_root, osp.basename(opt_file))
212
+ copyfile(opt_file, filename)
213
+
214
+ with open(filename, 'r+') as f:
215
+ lines = f.readlines()
216
+ lines.insert(0, f'# GENERATE TIME: {time.asctime()}\n# CMD:\n# {cmd}\n\n')
217
+ f.seek(0)
218
+ f.writelines(lines)
official/bsr/utils/plot_util.py ADDED
@@ -0,0 +1,83 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import re
2
+
3
+
4
+ def read_data_from_tensorboard(log_path, tag):
5
+ """Get raw data (steps and values) from tensorboard events.
6
+
7
+ Args:
8
+ log_path (str): Path to the tensorboard log.
9
+ tag (str): tag to be read.
10
+ """
11
+ from tensorboard.backend.event_processing.event_accumulator import EventAccumulator
12
+
13
+ # tensorboard event
14
+ event_acc = EventAccumulator(log_path)
15
+ event_acc.Reload()
16
+ scalar_list = event_acc.Tags()['scalars']
17
+ print('tag list: ', scalar_list)
18
+ steps = [int(s.step) for s in event_acc.Scalars(tag)]
19
+ values = [s.value for s in event_acc.Scalars(tag)]
20
+ return steps, values
21
+
22
+
23
+ def read_data_from_txt_2v(path, pattern, step_one=False):
24
+ """Read data from txt with 2 returned values (usually [step, value]).
25
+
26
+ Args:
27
+ path (str): path to the txt file.
28
+ pattern (str): re (regular expression) pattern.
29
+ step_one (bool): add 1 to steps. Default: False.
30
+ """
31
+ with open(path) as f:
32
+ lines = f.readlines()
33
+ lines = [line.strip() for line in lines]
34
+ steps = []
35
+ values = []
36
+
37
+ pattern = re.compile(pattern)
38
+ for line in lines:
39
+ match = pattern.match(line)
40
+ if match:
41
+ steps.append(int(match.group(1)))
42
+ values.append(float(match.group(2)))
43
+ if step_one:
44
+ steps = [v + 1 for v in steps]
45
+ return steps, values
46
+
47
+
48
+ def read_data_from_txt_1v(path, pattern):
49
+ """Read data from txt with 1 returned values.
50
+
51
+ Args:
52
+ path (str): path to the txt file.
53
+ pattern (str): re (regular expression) pattern.
54
+ """
55
+ with open(path) as f:
56
+ lines = f.readlines()
57
+ lines = [line.strip() for line in lines]
58
+ data = []
59
+
60
+ pattern = re.compile(pattern)
61
+ for line in lines:
62
+ match = pattern.match(line)
63
+ if match:
64
+ data.append(float(match.group(1)))
65
+ return data
66
+
67
+
68
+ def smooth_data(values, smooth_weight):
69
+ """ Smooth data using 1st-order IIR low-pass filter (what tensorflow does).
70
+
71
+ Reference: https://github.com/tensorflow/tensorboard/blob/f801ebf1f9fbfe2baee1ddd65714d0bccc640fb1/tensorboard/plugins/scalar/vz_line_chart/vz-line-chart.ts#L704 # noqa: E501
72
+
73
+ Args:
74
+ values (list): A list of values to be smoothed.
75
+ smooth_weight (float): Smooth weight.
76
+ """
77
+ values_sm = []
78
+ last_sm_value = values[0]
79
+ for value in values:
80
+ value_sm = last_sm_value * smooth_weight + (1 - smooth_weight) * value
81
+ values_sm.append(value_sm)
82
+ last_sm_value = value_sm
83
+ return values_sm
official/bsr/utils/registry.py ADDED
@@ -0,0 +1,88 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Modified from: https://github.com/facebookresearch/fvcore/blob/master/fvcore/common/registry.py # noqa: E501
2
+
3
+
4
+ class Registry():
5
+ """
6
+ The registry that provides name -> object mapping, to support third-party
7
+ users' custom modules.
8
+
9
+ To create a registry (e.g. a backbone registry):
10
+
11
+ .. code-block:: python
12
+
13
+ BACKBONE_REGISTRY = Registry('BACKBONE')
14
+
15
+ To register an object:
16
+
17
+ .. code-block:: python
18
+
19
+ @BACKBONE_REGISTRY.register()
20
+ class MyBackbone():
21
+ ...
22
+
23
+ Or:
24
+
25
+ .. code-block:: python
26
+
27
+ BACKBONE_REGISTRY.register(MyBackbone)
28
+ """
29
+
30
+ def __init__(self, name):
31
+ """
32
+ Args:
33
+ name (str): the name of this registry
34
+ """
35
+ self._name = name
36
+ self._obj_map = {}
37
+
38
+ def _do_register(self, name, obj, suffix=None):
39
+ if isinstance(suffix, str):
40
+ name = name + '_' + suffix
41
+
42
+ assert (name not in self._obj_map), (f"An object named '{name}' was already registered "
43
+ f"in '{self._name}' registry!")
44
+ self._obj_map[name] = obj
45
+
46
+ def register(self, obj=None, suffix=None):
47
+ """
48
+ Register the given object under the the name `obj.__name__`.
49
+ Can be used as either a decorator or not.
50
+ See docstring of this class for usage.
51
+ """
52
+ if obj is None:
53
+ # used as a decorator
54
+ def deco(func_or_class):
55
+ name = func_or_class.__name__
56
+ self._do_register(name, func_or_class, suffix)
57
+ return func_or_class
58
+
59
+ return deco
60
+
61
+ # used as a function call
62
+ name = obj.__name__
63
+ self._do_register(name, obj, suffix)
64
+
65
+ def get(self, name, suffix='basicsr'):
66
+ ret = self._obj_map.get(name)
67
+ if ret is None:
68
+ ret = self._obj_map.get(name + '_' + suffix)
69
+ print(f'Name {name} is not found, use name: {name}_{suffix}!')
70
+ if ret is None:
71
+ raise KeyError(f"No object named '{name}' found in '{self._name}' registry!")
72
+ return ret
73
+
74
+ def __contains__(self, name):
75
+ return name in self._obj_map
76
+
77
+ def __iter__(self):
78
+ return iter(self._obj_map.items())
79
+
80
+ def keys(self):
81
+ return self._obj_map.keys()
82
+
83
+
84
+ DATASET_REGISTRY = Registry('dataset')
85
+ ARCH_REGISTRY = Registry('arch')
86
+ MODEL_REGISTRY = Registry('model')
87
+ LOSS_REGISTRY = Registry('loss')
88
+ METRIC_REGISTRY = Registry('metric')
official/dataset.py ADDED
@@ -0,0 +1,280 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, random, cv2, os, math, glob
2
+ import torch.nn.functional as F
3
+ import numpy as np
4
+ from bsr.degradations import circular_lowpass_kernel, random_mixed_kernels, random_add_gaussian_noise_pt, random_add_poisson_noise_pt
5
+ from bsr.transforms import augment, paired_random_crop
6
+ from bsr.utils import FileClient, imfrombytes, img2tensor, DiffJPEG
7
+ from bsr.utils.img_process_util import filter2D
8
+
9
+ class RealESRGANDataset(torch.utils.data.Dataset):
10
+ def __init__(self, opt, bsz):
11
+ super(RealESRGANDataset, self).__init__()
12
+ self.opt = opt
13
+ self.file_client = FileClient("disk")
14
+ self.gt_folder = opt["dataroot_gt"]
15
+ self.len = bsz * opt["iter_num"]
16
+ self.paths = glob.glob(os.path.join(self.gt_folder, "**/*"), recursive=True)
17
+
18
+ # blur settings for the first degradation
19
+ self.blur_kernel_size = opt["blur_kernel_size"]
20
+ self.kernel_list = opt["kernel_list"]
21
+ self.kernel_prob = opt["kernel_prob"] # a list for each kernel probability
22
+ self.blur_sigma = opt["blur_sigma"]
23
+ self.betag_range = opt["betag_range"] # betag used in generalized Gaussian blur kernels
24
+ self.betap_range = opt["betap_range"] # betap used in plateau blur kernels
25
+ self.sinc_prob = opt["sinc_prob"] # the probability for sinc filters
26
+
27
+ # blur settings for the second degradation
28
+ self.blur_kernel_size2 = opt["blur_kernel_size2"]
29
+ self.kernel_list2 = opt["kernel_list2"]
30
+ self.kernel_prob2 = opt["kernel_prob2"]
31
+ self.blur_sigma2 = opt["blur_sigma2"]
32
+ self.betag_range2 = opt["betag_range2"]
33
+ self.betap_range2 = opt["betap_range2"]
34
+ self.sinc_prob2 = opt["sinc_prob2"]
35
+
36
+ # a final sinc filter
37
+ self.final_sinc_prob = opt["final_sinc_prob"]
38
+
39
+ self.kernel_range = [2 * v + 1 for v in range(3, 11)] # kernel size ranges from 7 to 21
40
+ # TODO: kernel range is now hard-coded, should be in the configure file
41
+ self.pulse_tensor = torch.zeros(21, 21).float() # convolving with pulse tensor brings no blurry effect
42
+ self.pulse_tensor[10, 10] = 1
43
+
44
+ def __getitem__(self, index):
45
+ index = random.randint(0, len(self.paths) - 1)
46
+ gt_path = self.paths[index]
47
+ img_gt = imfrombytes(self.file_client.get(gt_path, "gt"), float32=True)
48
+ img_gt = augment(img_gt, self.opt["use_hflip"], self.opt["use_rot"])
49
+ h, w = img_gt.shape[0:2]
50
+ crop_pad_size = self.opt.gt_size
51
+ if h < crop_pad_size or w < crop_pad_size:
52
+ pad_h = max(0, crop_pad_size - h)
53
+ pad_w = max(0, crop_pad_size - w)
54
+ img_gt = cv2.copyMakeBorder(img_gt, 0, pad_h, 0, pad_w, cv2.BORDER_REFLECT_101)
55
+ if img_gt.shape[0] > crop_pad_size or img_gt.shape[1] > crop_pad_size:
56
+ h, w = img_gt.shape[0:2]
57
+ top = random.randint(0, h - crop_pad_size)
58
+ left = random.randint(0, w - crop_pad_size)
59
+ img_gt = img_gt[top:top + crop_pad_size, left:left + crop_pad_size, ...]
60
+
61
+ # ------------------------ Generate kernels (used in the first degradation) ------------------------ #
62
+ kernel_size = random.choice(self.kernel_range)
63
+ if np.random.uniform() < self.opt["sinc_prob"]:
64
+ # this sinc filter setting is for kernels ranging from [7, 21]
65
+ if kernel_size < 13:
66
+ omega_c = np.random.uniform(np.pi / 3, np.pi)
67
+ else:
68
+ omega_c = np.random.uniform(np.pi / 5, np.pi)
69
+ kernel = circular_lowpass_kernel(omega_c, kernel_size, pad_to=False)
70
+ else:
71
+ kernel = random_mixed_kernels(
72
+ self.kernel_list,
73
+ self.kernel_prob,
74
+ kernel_size,
75
+ self.blur_sigma,
76
+ self.blur_sigma, [-math.pi, math.pi],
77
+ self.betag_range,
78
+ self.betap_range,
79
+ noise_range=None)
80
+ # pad kernel
81
+ pad_size = (21 - kernel_size) // 2
82
+ kernel = np.pad(kernel, ((pad_size, pad_size), (pad_size, pad_size)))
83
+
84
+ # ------------------------ Generate kernels (used in the second degradation) ------------------------ #
85
+ kernel_size = random.choice(self.kernel_range)
86
+ if np.random.uniform() < self.opt["sinc_prob2"]:
87
+ if kernel_size < 13:
88
+ omega_c = np.random.uniform(np.pi / 3, np.pi)
89
+ else:
90
+ omega_c = np.random.uniform(np.pi / 5, np.pi)
91
+ kernel2 = circular_lowpass_kernel(omega_c, kernel_size, pad_to=False)
92
+ else:
93
+ kernel2 = random_mixed_kernels(
94
+ self.kernel_list2,
95
+ self.kernel_prob2,
96
+ kernel_size,
97
+ self.blur_sigma2,
98
+ self.blur_sigma2, [-math.pi, math.pi],
99
+ self.betag_range2,
100
+ self.betap_range2,
101
+ noise_range=None)
102
+
103
+ # pad kernel
104
+ pad_size = (21 - kernel_size) // 2
105
+ kernel2 = np.pad(kernel2, ((pad_size, pad_size), (pad_size, pad_size)))
106
+
107
+ # ------------------------------------- the final sinc kernel ------------------------------------- #
108
+ if np.random.uniform() < self.opt["final_sinc_prob"]:
109
+ kernel_size = random.choice(self.kernel_range)
110
+ omega_c = np.random.uniform(np.pi / 3, np.pi)
111
+ sinc_kernel = circular_lowpass_kernel(omega_c, kernel_size, pad_to=21)
112
+ sinc_kernel = torch.FloatTensor(sinc_kernel)
113
+ else:
114
+ sinc_kernel = self.pulse_tensor
115
+
116
+ # BGR to RGB, HWC to CHW, numpy to tensor
117
+ img_gt = img2tensor([img_gt], bgr2rgb=True, float32=True)[0]
118
+ kernel = torch.FloatTensor(kernel)
119
+ kernel2 = torch.FloatTensor(kernel2)
120
+
121
+ return_d = {"gt": img_gt, "kernel1": kernel, "kernel2": kernel2, "sinc_kernel": sinc_kernel, "gt_path": gt_path}
122
+ return return_d
123
+
124
+ def __len__(self):
125
+ return self.len
126
+
127
+ class RealESRGANDegrader:
128
+ def __init__(self, opt, device):
129
+ self.opt = opt
130
+ self.device = device
131
+ self.jpeger = DiffJPEG(differentiable=False).to(device) # simulate JPEG compression artifacts
132
+ self.queue_size = 1200
133
+
134
+ @torch.no_grad()
135
+ def _dequeue_and_enqueue(self):
136
+ """It is the training pair pool for increasing the diversity in a batch.
137
+
138
+ Batch processing limits the diversity of synthetic degradations in a batch. For example, samples in a
139
+ batch could not have different resize scaling factors. Therefore, we employ this training pair pool
140
+ to increase the degradation diversity in a batch.
141
+ """
142
+ # initialize
143
+ b, c, h, w = self.lq.size()
144
+ if not hasattr(self, "queue_lr"):
145
+ assert self.queue_size % b == 0, f"queue size {self.queue_size} should be divisible by batch size {b}"
146
+ self.queue_lr = torch.zeros(self.queue_size, c, h, w).to(self.device)
147
+ _, c, h, w = self.gt.size()
148
+ self.queue_gt = torch.zeros(self.queue_size, c, h, w).to(self.device)
149
+ self.queue_ptr = 0
150
+ if self.queue_ptr == self.queue_size: # the pool is full
151
+ # do dequeue and enqueue
152
+ # shuffle
153
+ idx = torch.randperm(self.queue_size)
154
+ self.queue_lr = self.queue_lr[idx]
155
+ self.queue_gt = self.queue_gt[idx]
156
+ # get first b samples
157
+ lq_dequeue = self.queue_lr[0:b, :, :, :].clone()
158
+ gt_dequeue = self.queue_gt[0:b, :, :, :].clone()
159
+ # update the queue
160
+ self.queue_lr[0:b, :, :, :] = self.lq.clone()
161
+ self.queue_gt[0:b, :, :, :] = self.gt.clone()
162
+
163
+ self.lq = lq_dequeue
164
+ self.gt = gt_dequeue
165
+ else:
166
+ # only do enqueue
167
+ self.queue_lr[self.queue_ptr:self.queue_ptr + b, :, :, :] = self.lq.clone()
168
+ self.queue_gt[self.queue_ptr:self.queue_ptr + b, :, :, :] = self.gt.clone()
169
+ self.queue_ptr = self.queue_ptr + b
170
+
171
+ @torch.no_grad()
172
+ def degrade(self, data):
173
+ """Accept data from dataloader, and then add two-order degradations to obtain LQ images.
174
+ """
175
+ # training data synthesis
176
+ self.gt = data["gt"].to(self.device)
177
+
178
+ self.kernel1 = data["kernel1"].to(self.device)
179
+ self.kernel2 = data["kernel2"].to(self.device)
180
+ self.sinc_kernel = data["sinc_kernel"].to(self.device)
181
+
182
+ ori_h, ori_w = self.gt.size()[2:4]
183
+
184
+ # ----------------------- The first degradation process ----------------------- #
185
+ # blur
186
+ out = filter2D(self.gt, self.kernel1)
187
+ # random resize
188
+ updown_type = random.choices(["up", "down", "keep"], self.opt["resize_prob"])[0]
189
+ if updown_type == "up":
190
+ scale = np.random.uniform(1, self.opt["resize_range"][1])
191
+ elif updown_type == "down":
192
+ scale = np.random.uniform(self.opt["resize_range"][0], 1)
193
+ else:
194
+ scale = 1
195
+ mode = random.choice(["area", "bilinear", "bicubic"])
196
+ out = F.interpolate(out, scale_factor=scale, mode=mode)
197
+ # add noise
198
+ gray_noise_prob = self.opt["gray_noise_prob"]
199
+ if np.random.uniform() < self.opt["gaussian_noise_prob"]:
200
+ out = random_add_gaussian_noise_pt(
201
+ out, sigma_range=self.opt["noise_range"], clip=True, rounds=False, gray_prob=gray_noise_prob)
202
+ else:
203
+ out = random_add_poisson_noise_pt(
204
+ out,
205
+ scale_range=self.opt["poisson_scale_range"],
206
+ gray_prob=gray_noise_prob,
207
+ clip=True,
208
+ rounds=False)
209
+ # JPEG compression
210
+ jpeg_p = out.new_zeros(out.size(0)).uniform_(*self.opt["jpeg_range"])
211
+ out = torch.clamp(out, 0, 1) # clamp to [0, 1], otherwise JPEGer will result in unpleasant artifacts
212
+ out = self.jpeger(out, quality=jpeg_p)
213
+
214
+ # ----------------------- The second degradation process ----------------------- #
215
+ # blur
216
+ if np.random.uniform() < self.opt["second_blur_prob"]:
217
+ out = filter2D(out, self.kernel2)
218
+ # random resize
219
+ updown_type = random.choices(["up", "down", "keep"], self.opt["resize_prob2"])[0]
220
+ if updown_type == "up":
221
+ scale = np.random.uniform(1, self.opt["resize_range2"][1])
222
+ elif updown_type == "down":
223
+ scale = np.random.uniform(self.opt["resize_range2"][0], 1)
224
+ else:
225
+ scale = 1
226
+ mode = random.choice(["area", "bilinear", "bicubic"])
227
+ out = F.interpolate(
228
+ out, size=(int(ori_h / self.opt["scale"] * scale), int(ori_w / self.opt["scale"] * scale)), mode=mode)
229
+ # add noise
230
+ gray_noise_prob = self.opt["gray_noise_prob2"]
231
+ if np.random.uniform() < self.opt["gaussian_noise_prob2"]:
232
+ out = random_add_gaussian_noise_pt(
233
+ out, sigma_range=self.opt["noise_range2"], clip=True, rounds=False, gray_prob=gray_noise_prob)
234
+ else:
235
+ out = random_add_poisson_noise_pt(
236
+ out,
237
+ scale_range=self.opt["poisson_scale_range2"],
238
+ gray_prob=gray_noise_prob,
239
+ clip=True,
240
+ rounds=False)
241
+
242
+ # JPEG compression + the final sinc filter
243
+ # We also need to resize images to desired sizes. We group [resize back + sinc filter] together
244
+ # as one operation.
245
+ # We consider two orders:
246
+ # 1. [resize back + sinc filter] + JPEG compression
247
+ # 2. JPEG compression + [resize back + sinc filter]
248
+ # Empirically, we find other combinations (sinc + JPEG + Resize) will introduce twisted lines.
249
+ if np.random.uniform() < 0.5:
250
+ # resize back + the final sinc filter
251
+ mode = random.choice(["area", "bilinear", "bicubic"])
252
+ out = F.interpolate(out, size=(ori_h // self.opt["scale"], ori_w // self.opt["scale"]), mode=mode)
253
+ out = filter2D(out, self.sinc_kernel)
254
+ # JPEG compression
255
+ jpeg_p = out.new_zeros(out.size(0)).uniform_(*self.opt["jpeg_range2"])
256
+ out = torch.clamp(out, 0, 1)
257
+ out = self.jpeger(out, quality=jpeg_p)
258
+ else:
259
+ # JPEG compression
260
+ jpeg_p = out.new_zeros(out.size(0)).uniform_(*self.opt["jpeg_range2"])
261
+ out = torch.clamp(out, 0, 1)
262
+ out = self.jpeger(out, quality=jpeg_p)
263
+ # resize back + the final sinc filter
264
+ mode = random.choice(["area", "bilinear", "bicubic"])
265
+ out = F.interpolate(out, size=(ori_h // self.opt["scale"], ori_w // self.opt["scale"]), mode=mode)
266
+ out = filter2D(out, self.sinc_kernel)
267
+
268
+ # clamp and round
269
+ self.lq = torch.clamp((out * 255.0).round(), 0, 255) / 255.
270
+
271
+ # random crop
272
+ gt_size = self.opt["gt_size"]
273
+ self.gt, self.lq = paired_random_crop(self.gt, self.lq, gt_size, self.opt["scale"])
274
+
275
+ # training pair pool
276
+ self._dequeue_and_enqueue()
277
+ # sharpen self.gt again, as we have changed the self.gt with self._dequeue_and_enqueue
278
+ self.lq = self.lq.contiguous() # for the warning: grad and param do not obey the gradient layout contract
279
+
280
+ return self.lq, self.gt
official/evaluate.py ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, os, glob, pyiqa
2
+ from argparse import ArgumentParser
3
+ import numpy as np
4
+ from PIL import Image
5
+ from tqdm import tqdm
6
+ from torchvision import transforms
7
+
8
+ parser = ArgumentParser()
9
+ parser.add_argument("--HR_dir", type=str, default="testset/RealSR/HR")
10
+ parser.add_argument("--SR_dir", type=str, default="result/RealSR")
11
+ args = parser.parse_args()
12
+
13
+ device = torch.device("cuda")
14
+
15
+ psnr = pyiqa.create_metric("psnr", test_y_channel=True, color_space="ycbcr", device=device)
16
+ ssim = pyiqa.create_metric("ssim", test_y_channel=True, color_space="ycbcr", device=device)
17
+ lpips = pyiqa.create_metric("lpips", device=device)
18
+ dists = pyiqa.create_metric("dists", device=device)
19
+ fid = pyiqa.create_metric("fid", device=device)
20
+ niqe = pyiqa.create_metric("niqe", device=device)
21
+ maniqa = pyiqa.create_metric("maniqa-pipal", device=device)
22
+ clipiqa = pyiqa.create_metric("clipiqa", device=device)
23
+ musiq = pyiqa.create_metric("musiq", device=device)
24
+
25
+ test_SR_paths = list(sorted(glob.glob(os.path.join(args.SR_dir, "*"))))
26
+ test_HR_paths = list(sorted(glob.glob(os.path.join(args.HR_dir, "*"))))
27
+
28
+ metrics = {"psnr": [], "ssim": [], "lpips": [], "dists": [], "niqe": [], "maniqa": [], "musiq": [], "clipiqa": []}
29
+
30
+ for i, (SR_path, HR_path) in tqdm(enumerate(zip(test_SR_paths, test_HR_paths))):
31
+ SR = Image.open(SR_path).convert("RGB")
32
+ SR = transforms.ToTensor()(SR).to(device).unsqueeze(0)
33
+ HR = Image.open(HR_path).convert("RGB")
34
+ HR = transforms.ToTensor()(HR).to(device).unsqueeze(0)
35
+ metrics["psnr"].append(psnr(SR, HR).item())
36
+ metrics["ssim"].append(ssim(SR, HR).item())
37
+ metrics["lpips"].append(lpips(SR, HR).item())
38
+ metrics["dists"].append(dists(SR, HR).item())
39
+ metrics["niqe"].append(niqe(SR).item())
40
+ metrics["maniqa"].append(maniqa(SR).item())
41
+ metrics["clipiqa"].append(clipiqa(SR).item())
42
+ metrics["musiq"].append(musiq(SR).item())
43
+
44
+ for k in metrics.keys():
45
+ metrics[k] = np.mean(metrics[k])
46
+
47
+ metrics["fid"] = fid(args.SR_dir, args.HR_dir)
48
+
49
+ for k, v in metrics.items():
50
+ if k == "niqe":
51
+ print(k, f"{v:.3g}")
52
+ elif k == "fid":
53
+ print(k, f"{v:.5g}")
54
+ else:
55
+ print(k, f"{v:.4g}")
official/forward.py ADDED
@@ -0,0 +1,67 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+
3
+ def MyUNet2DConditionModel_SD_forward(self, x):
4
+ global skip
5
+ x = self.conv_in(x)
6
+ skip = [x]
7
+ x = self.body(x)
8
+ return x
9
+
10
+ def MyCrossAttnDownBlock2D_SD_forward(self, x):
11
+ for i in range(2):
12
+ x = self.resnets[i](x)
13
+ x = self.attentions[i](x)
14
+ skip.append(x)
15
+ if self.downsamplers is not None:
16
+ x = self.downsamplers[0](x)
17
+ skip.append(x)
18
+ return x
19
+
20
+ def MyCrossAttnUpBlock2D_SD_forward(self, x):
21
+ for i in range(3):
22
+ x = self.resnets[i](torch.cat([x, skip.pop()], dim=1))
23
+ x = self.attentions[i](x)
24
+ if self.upsamplers is not None:
25
+ x = self.upsamplers[0](x)
26
+ return x
27
+
28
+ def MyDownBlock2D_SD_forward(self, x):
29
+ for i in range(2):
30
+ x = self.resnets[i](x)
31
+ skip.append(x)
32
+ return x
33
+
34
+ def MyUNetMidBlock2DCrossAttn_SD_forward(self, x):
35
+ x = self.resnets[0](x)
36
+ x = self.attentions[0](x)
37
+ x = self.resnets[1](x)
38
+ return x
39
+
40
+ def MyUpBlock2D_SD_forward(self, x):
41
+ for i in range(3):
42
+ x = self.resnets[i](torch.cat([x, skip.pop()], dim=1))
43
+ x = self.upsamplers[0](x)
44
+ return x
45
+
46
+ def MyResnetBlock2D_SD_forward(self, x_in):
47
+ x = self.norm1(x_in)
48
+ x = self.nonlinearity(x)
49
+ x = self.conv1(x)
50
+ x = self.norm2(x)
51
+ x = self.nonlinearity(x)
52
+ x = self.conv2(x)
53
+ if self.in_channels == self.out_channels:
54
+ return x + x_in
55
+ return x + self.conv_shortcut(x_in)
56
+
57
+ def MyTransformer2DModel_SD_forward(self, x_in):
58
+ b, c, h, w = x_in.shape
59
+ x = self.norm(x_in)
60
+ x = x.permute(0, 2, 3, 1).reshape(b, h * w, c).contiguous()
61
+ x = self.proj_in(x)
62
+ for block in self.transformer_blocks:
63
+ x = x + block.attn1(block.norm1(x))
64
+ x = x + block.ff(block.norm3(x))
65
+ x = self.proj_out(x)
66
+ x = x.reshape(b, h, w, c).permute(0, 3, 1, 2).contiguous()
67
+ return x + x_in
official/model.py ADDED
@@ -0,0 +1,152 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, types, copy
2
+ from torch import nn
3
+ import torch.nn.functional as F
4
+ from diffusers.models.unets.unet_2d_blocks import CrossAttnDownBlock2D, \
5
+ CrossAttnUpBlock2D, \
6
+ DownBlock2D, \
7
+ UpBlock2D, \
8
+ UNetMidBlock2DCrossAttn
9
+ from diffusers.models.resnet import ResnetBlock2D
10
+ from diffusers.models.transformers.transformer_2d import Transformer2DModel
11
+ from diffusers.models.attention import BasicTransformerBlock
12
+ from diffusers.models.downsampling import Downsample2D
13
+ from diffusers.models.upsampling import Upsample2D
14
+ from forward import MyUNet2DConditionModel_SD_forward, \
15
+ MyCrossAttnDownBlock2D_SD_forward, \
16
+ MyDownBlock2D_SD_forward, \
17
+ MyUNetMidBlock2DCrossAttn_SD_forward, \
18
+ MyCrossAttnUpBlock2D_SD_forward, \
19
+ MyUpBlock2D_SD_forward, \
20
+ MyResnetBlock2D_SD_forward, \
21
+ MyTransformer2DModel_SD_forward
22
+
23
+ def find_parent(model, module_name):
24
+ components = module_name.split(".")
25
+ parent = model
26
+ for comp in components[:-1]:
27
+ parent = getattr(parent, comp)
28
+ return parent, components[-1]
29
+
30
+ def halve_channels(model):
31
+ for name, module in model.named_modules():
32
+ if hasattr(module, "pruned"):
33
+ continue
34
+ if isinstance(module, nn.Conv2d):
35
+ in_channels = int(module.in_channels * 0.75)
36
+ out_channels = int(module.out_channels * 0.75)
37
+ new_conv = nn.Conv2d(in_channels=in_channels,
38
+ out_channels=out_channels,
39
+ kernel_size=module.kernel_size,
40
+ stride=module.stride,
41
+ padding=module.padding,
42
+ dilation=module.dilation,
43
+ groups=module.groups,
44
+ bias=module.bias is not None)
45
+ with torch.no_grad():
46
+ new_conv.weight.copy_(module.weight[:out_channels, :in_channels])
47
+ if module.bias is not None:
48
+ new_conv.bias.copy_(module.bias[:out_channels])
49
+ parent, last_name = find_parent(model, name)
50
+ setattr(parent, last_name, new_conv)
51
+ new_conv.pruned = True
52
+ elif isinstance(module, nn.Linear):
53
+ in_features = int(module.in_features * 0.75)
54
+ out_features = int(module.out_features * 0.75)
55
+ new_linear = nn.Linear(in_features=in_features,
56
+ out_features=out_features,
57
+ bias=module.bias is not None)
58
+ with torch.no_grad():
59
+ new_linear.weight.copy_(module.weight[:out_features, :in_features])
60
+ if module.bias is not None:
61
+ new_linear.bias.copy_(module.bias[:out_features])
62
+ parent, last_name = find_parent(model, name)
63
+ setattr(parent, last_name, new_linear)
64
+ new_linear.pruned = True
65
+ elif isinstance(module, nn.GroupNorm):
66
+ num_channels = int(module.num_channels * 0.75)
67
+ for num_groups in [32, 24, 16, 12, 8, 6, 4, 2, 1]:
68
+ if num_channels % num_groups == 0:
69
+ break
70
+ new_gn = nn.GroupNorm(num_groups=num_groups,
71
+ num_channels=num_channels,
72
+ eps=module.eps,
73
+ affine=module.affine)
74
+ with torch.no_grad():
75
+ new_gn.weight.copy_(module.weight[:num_channels])
76
+ new_gn.bias.copy_(module.bias[:num_channels])
77
+ parent, last_name = find_parent(model, name)
78
+ setattr(parent, last_name, new_gn)
79
+ new_gn.pruned = True
80
+ elif isinstance(module, nn.LayerNorm):
81
+ normalized_shape = int(module.normalized_shape[0] * 0.75)
82
+ new_ln = nn.LayerNorm(normalized_shape,
83
+ eps=module.eps,
84
+ elementwise_affine=module.elementwise_affine)
85
+ with torch.no_grad():
86
+ new_ln.weight.copy_(module.weight[:normalized_shape])
87
+ new_ln.bias.copy_(module.bias[:normalized_shape])
88
+ parent, last_name = find_parent(model, name)
89
+ setattr(parent, last_name, new_ln)
90
+ new_ln.pruned = True
91
+ elif isinstance(module, Downsample2D) or isinstance(module, Upsample2D):
92
+ module.channels = int(module.channels * 0.75)
93
+
94
+ class Net(nn.Module):
95
+ def __init__(self, unet, decoder):
96
+ super().__init__()
97
+ del unet.time_embedding
98
+ new_conv_in = nn.Conv2d(16, 320, 3, padding=1)
99
+ new_conv_in.weight.data = unet.conv_in.weight.data.repeat(1, 4, 1, 1)
100
+ new_conv_in.bias.data = unet.conv_in.bias.data
101
+ unet.conv_in = new_conv_in
102
+ new_conv_out = nn.Conv2d(320, 342, 3, padding=1)
103
+ new_conv_out.weight.data = unet.conv_out.weight.data.repeat(86, 1, 1, 1)[:342]
104
+ new_conv_out.bias.data = unet.conv_out.bias.data.repeat(86,)[:342]
105
+ unet.conv_out = new_conv_out
106
+ def ResnetBlock2D_remove_time_emb_proj(module):
107
+ if isinstance(module, ResnetBlock2D):
108
+ del module.time_emb_proj
109
+ unet.apply(ResnetBlock2D_remove_time_emb_proj)
110
+ def BasicTransformerBlock_remove_cross_attn(module):
111
+ if isinstance(module, BasicTransformerBlock):
112
+ del module.attn2, module.norm2
113
+ unet.apply(BasicTransformerBlock_remove_cross_attn)
114
+ def set_inplace_to_true(module):
115
+ if isinstance(module, nn.Dropout) or isinstance(module, nn.SiLU):
116
+ module.inplace = True
117
+ unet.apply(set_inplace_to_true)
118
+ def replace_forward_methods(module):
119
+ if isinstance(module, CrossAttnDownBlock2D):
120
+ module.forward = types.MethodType(MyCrossAttnDownBlock2D_SD_forward, module)
121
+ elif isinstance(module, DownBlock2D):
122
+ module.forward = types.MethodType(MyDownBlock2D_SD_forward, module)
123
+ elif isinstance(module, UNetMidBlock2DCrossAttn):
124
+ module.forward = types.MethodType(MyUNetMidBlock2DCrossAttn_SD_forward, module)
125
+ elif isinstance(module, UpBlock2D):
126
+ module.forward = types.MethodType(MyUpBlock2D_SD_forward, module)
127
+ elif isinstance(module, CrossAttnUpBlock2D):
128
+ module.forward = types.MethodType(MyCrossAttnUpBlock2D_SD_forward, module)
129
+ elif isinstance(module, ResnetBlock2D):
130
+ module.forward = types.MethodType(MyResnetBlock2D_SD_forward, module)
131
+ elif isinstance(module, Transformer2DModel):
132
+ module.forward = types.MethodType(MyTransformer2DModel_SD_forward, module)
133
+ unet.apply(replace_forward_methods)
134
+ unet.forward = types.MethodType(MyUNet2DConditionModel_SD_forward, unet)
135
+ halve_channels(unet)
136
+ unet.body = nn.Sequential(
137
+ *unet.down_blocks,
138
+ unet.mid_block,
139
+ *unet.up_blocks,
140
+ unet.conv_norm_out,
141
+ unet.conv_act,
142
+ unet.conv_out,
143
+ )
144
+ del decoder.conv_in, decoder.up_blocks, decoder.conv_norm_out, decoder.conv_act, decoder.conv_out
145
+ self.body = nn.Sequential(
146
+ nn.PixelUnshuffle(2),
147
+ unet,
148
+ decoder.mid_block,
149
+ )
150
+
151
+ def forward(self, x):
152
+ return self.body(x)
official/test.py ADDED
@@ -0,0 +1,70 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, os, glob, copy
2
+ import torch.nn.functional as F
3
+ import numpy as np
4
+ from PIL import Image
5
+ from argparse import ArgumentParser
6
+ from torchvision import transforms
7
+ from model import Net
8
+
9
+ parser = ArgumentParser()
10
+ parser.add_argument("--epoch", type=int, default=200)
11
+ parser.add_argument("--model_dir", type=str, default="weight")
12
+ parser.add_argument("--LR_dir", type=str, default="testset/RealSR/LR")
13
+ parser.add_argument("--HR_dir", type=str, default="testset/RealSR/HR")
14
+ parser.add_argument("--SR_dir", type=str, default="result/RealSR")
15
+ args = parser.parse_args()
16
+
17
+ device = torch.device("cuda")
18
+
19
+ from diffusers import StableDiffusionPipeline
20
+ model_id = "stabilityai/stable-diffusion-2-1-base"
21
+ pipe = StableDiffusionPipeline.from_pretrained(model_id).to(device)
22
+
23
+ vae = pipe.vae
24
+ tokenizer = pipe.tokenizer
25
+ unet = pipe.unet
26
+ noise_scheduler = pipe.scheduler
27
+ text_encoder = pipe.text_encoder
28
+
29
+ from diffusers.models.autoencoders.vae import Decoder
30
+ ckpt_halfdecoder = torch.load("./weight/pretrained/halfDecoder.ckpt", weights_only=False)
31
+ decoder = Decoder(in_channels=4,
32
+ out_channels=3,
33
+ up_block_types=["UpDecoderBlock2D" for _ in range(4)],
34
+ block_out_channels=[64, 128, 256, 256],
35
+ layers_per_block=2,
36
+ norm_num_groups=32,
37
+ act_fn="silu",
38
+ norm_type="group",
39
+ mid_block_add_attention=True).to(device)
40
+ decoder_ckpt = {}
41
+ for k,v in ckpt_halfdecoder["state_dict"].items():
42
+ if "decoder" in k:
43
+ new_k = k.replace("decoder.", "")
44
+ decoder_ckpt[new_k] = v
45
+ decoder.load_state_dict(decoder_ckpt, strict=True)
46
+
47
+ model = torch.nn.DataParallel(Net(unet, copy.deepcopy(decoder)))
48
+ model.load_state_dict(torch.load("./%s/net_params_%d.pkl" % (args.model_dir, args.epoch), weights_only=False))
49
+ model = torch.nn.Sequential(
50
+ model.module,
51
+ *decoder.up_blocks,
52
+ decoder.conv_norm_out,
53
+ decoder.conv_act,
54
+ decoder.conv_out,
55
+ ).to(device)
56
+
57
+ test_LR_paths = list(sorted(glob.glob(os.path.join(args.LR_dir, "*.png"))))
58
+ test_HR_paths = list(sorted(glob.glob(os.path.join(args.HR_dir, "*.png"))))
59
+
60
+ os.makedirs(args.SR_dir, exist_ok=True)
61
+
62
+ with torch.no_grad():
63
+ for i, path in enumerate(test_LR_paths):
64
+ LR = Image.open(path).convert("RGB")
65
+ LR = transforms.ToTensor()(LR).to(device).unsqueeze(0) * 2 - 1
66
+ SR = model(LR)
67
+ SR = (SR - SR.mean(dim=[2,3],keepdim=True)) / SR.std(dim=[2,3],keepdim=True) \
68
+ * LR.std(dim=[2,3],keepdim=True) + LR.mean(dim=[2,3],keepdim=True)
69
+ SR = transforms.ToPILImage()((SR[0] / 2 + 0.5).clamp(0, 1).cpu())
70
+ SR.save(os.path.join(args.SR_dir, os.path.basename(path)))
official/train.py ADDED
@@ -0,0 +1,225 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, os, glob, random, copy
2
+ import torch.nn.functional as F
3
+ from torch.utils.data import DataLoader
4
+ import torch.distributed as dist
5
+ from torch.nn.parallel import DistributedDataParallel as DDP
6
+ import numpy as np
7
+ from argparse import ArgumentParser
8
+ from time import time
9
+ from tqdm import tqdm
10
+ from omegaconf import OmegaConf
11
+ from dataset import RealESRGANDataset, RealESRGANDegrader
12
+ from model import Net
13
+ from ram.models.ram_lora import ram
14
+ from torchvision import transforms
15
+ from utils import add_lora_to_unet
16
+
17
+ dist.init_process_group(backend="nccl", init_method="env://")
18
+ rank = dist.get_rank()
19
+ world_size = dist.get_world_size()
20
+
21
+ parser = ArgumentParser()
22
+ parser.add_argument("--epoch", type=int, default=200)
23
+ parser.add_argument("--batch_size", type=int, default=12)
24
+ parser.add_argument("--learning_rate", type=float, default=1e-4)
25
+ parser.add_argument("--model_dir", type=str, default="weight")
26
+ parser.add_argument("--log_dir", type=str, default="log")
27
+ parser.add_argument("--save_interval", type=int, default=10)
28
+
29
+ args = parser.parse_args()
30
+
31
+ # fixed seed for reproduction
32
+ seed = rank
33
+ random.seed(seed)
34
+ np.random.seed(seed)
35
+ torch.manual_seed(seed)
36
+ torch.cuda.manual_seed_all(seed)
37
+
38
+ config = OmegaConf.load("config.yml")
39
+
40
+ epoch = args.epoch
41
+ learning_rate = args.learning_rate
42
+ bsz = args.batch_size
43
+
44
+ device = torch.device(f"cuda:{rank}" if torch.cuda.is_available() else "cpu")
45
+ torch.backends.cudnn.allow_tf32 = True
46
+ torch.backends.cuda.matmul.allow_tf32 = True
47
+
48
+ if rank == 0:
49
+ print("batch size per gpu =", bsz)
50
+
51
+ from diffusers import StableDiffusionPipeline
52
+ model_id = "stabilityai/stable-diffusion-2-1-base"
53
+ pipe = StableDiffusionPipeline.from_pretrained(model_id).to(device)
54
+
55
+ vae = pipe.vae
56
+ tokenizer = pipe.tokenizer
57
+ unet = pipe.unet
58
+ text_encoder = pipe.text_encoder
59
+
60
+ unet_D = copy.deepcopy(unet)
61
+ new_conv_in = torch.nn.Conv2d(256, 320, 3, padding=1).to(device)
62
+ new_conv_in.weight.data = unet_D.conv_in.weight.data.repeat(1, 64, 1, 1) / 64
63
+ new_conv_in.bias.data = unet_D.conv_in.bias.data
64
+ unet_D.conv_in = new_conv_in
65
+ unet_D = add_lora_to_unet(unet_D)
66
+ unet_D.set_adapters(["default_encoder", "default_decoder", "default_others"])
67
+
68
+ vae_teacher = copy.deepcopy(vae)
69
+ unet_teacher = copy.deepcopy(unet)
70
+
71
+ osediff = torch.load("./weight/pretrained/osediff.pkl", weights_only=False)
72
+ vae_teacher.load_state_dict(osediff["vae"])
73
+ unet_teacher.load_state_dict(osediff["unet"])
74
+
75
+ from diffusers.models.autoencoders.vae import Decoder
76
+ ckpt_halfdecoder = torch.load("./weight/pretrained/halfDecoder.ckpt", weights_only=False)
77
+ decoder = Decoder(in_channels=4,
78
+ out_channels=3,
79
+ up_block_types=["UpDecoderBlock2D" for _ in range(4)],
80
+ block_out_channels=[64, 128, 256, 256],
81
+ layers_per_block=2,
82
+ norm_num_groups=32,
83
+ act_fn="silu",
84
+ norm_type="group",
85
+ mid_block_add_attention=True).to(device)
86
+ decoder_ckpt = {}
87
+ for k, v in ckpt_halfdecoder["state_dict"].items():
88
+ if "decoder" in k:
89
+ new_k = k.replace("decoder.", "")
90
+ decoder_ckpt[new_k] = v
91
+ decoder.load_state_dict(decoder_ckpt, strict=True)
92
+
93
+ ram_transforms = transforms.Compose([
94
+ transforms.Resize((384, 384)),
95
+ transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
96
+ ])
97
+
98
+ DAPE = ram(pretrained="./weight/pretrained/ram_swin_large_14m.pth",
99
+ pretrained_condition="./weight/pretrained/DAPE.pth",
100
+ image_size=384,
101
+ vit="swin_l").eval().to(device)
102
+
103
+ vae.requires_grad_(False)
104
+ unet.requires_grad_(False)
105
+ text_encoder.requires_grad_(False)
106
+ vae_teacher.requires_grad_(False)
107
+ unet_teacher.requires_grad_(False)
108
+ decoder.requires_grad_(False)
109
+ DAPE.requires_grad_(False)
110
+
111
+ model = DDP(Net(unet, copy.deepcopy(decoder)).to(device), device_ids=[rank])
112
+ model_D = DDP(unet_D.to(device), device_ids=[rank])
113
+ model.requires_grad_(True)
114
+ model_D.requires_grad_(False)
115
+ params_to_opt = []
116
+ for n, p in model_D.named_parameters():
117
+ if "lora" in n or "conv_in" in n:
118
+ p.requires_grad = True
119
+ params_to_opt.append(p)
120
+
121
+ if rank == 0:
122
+ param_cnt = sum(p.numel() for p in model.parameters() if p.requires_grad)
123
+ print("#Param.", param_cnt/1e6, "M")
124
+
125
+ dataset = RealESRGANDataset(config, bsz)
126
+ degrader = RealESRGANDegrader(config, device)
127
+ dataloader = DataLoader(dataset, batch_size=bsz, num_workers=8)
128
+ optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)
129
+ optimizer_D = torch.optim.Adam(params_to_opt, lr=1e-6)
130
+ scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, milestones=[100,], gamma=0.5)
131
+ scaler = torch.cuda.amp.GradScaler()
132
+
133
+ model_dir = "./%s" % (args.model_dir,)
134
+ log_path = "./%s/log.txt" % (args.log_dir,)
135
+ os.makedirs(model_dir, exist_ok=True)
136
+ os.makedirs(args.log_dir, exist_ok=True)
137
+
138
+ print("start training...")
139
+ timesteps = torch.tensor([999], device=device).long().expand(bsz,)
140
+ alpha = pipe.scheduler.alphas_cumprod[999]
141
+ for epoch_i in range(1, epoch + 1):
142
+ start_time = time()
143
+ loss_avg = 0.0
144
+ loss_distil_avg = 0.0
145
+ loss_adv_avg = 0.0
146
+ loss_D_avg = 0.0
147
+ iter_num = 0
148
+ dist.barrier()
149
+ for batch in tqdm(dataloader):
150
+ with torch.cuda.amp.autocast(enabled=True):
151
+ with torch.no_grad():
152
+ LR, HR = degrader.degrade(batch)
153
+ text_input = tokenizer(DAPE.generate_tag(ram_transforms(LR))[0],
154
+ max_length=tokenizer.model_max_length,
155
+ padding="max_length", return_tensors="pt").to(device)
156
+ encoder_hidden_states = text_encoder(text_input.input_ids, return_dict=False)[0]
157
+ LR, HR = LR * 2 - 1, HR * 2 - 1
158
+ LR_ = F.interpolate(LR, scale_factor=4, mode="bicubic")
159
+ LR_latents = vae_teacher.encode(LR_).latent_dist.mean * vae_teacher.config.scaling_factor
160
+ HR_latents = vae.encode(HR).latent_dist.mean
161
+ pred_teacher = unet_teacher(
162
+ LR_latents,
163
+ timesteps,
164
+ encoder_hidden_states=encoder_hidden_states,
165
+ return_dict=False,
166
+ )[0]
167
+ z0_teacher = (LR_latents-((1-alpha)**0.5)*pred_teacher)/(alpha**0.5)
168
+ z0_teacher = vae_teacher.post_quant_conv(z0_teacher / vae_teacher.config.scaling_factor)
169
+ z0_teacher = decoder.conv_in(z0_teacher)
170
+ z0_teacher = decoder.mid_block(z0_teacher)
171
+ z0_gt = vae.post_quant_conv(HR_latents)
172
+ z0_gt = decoder.conv_in(z0_gt)
173
+ z0_gt = decoder.mid_block(z0_gt)
174
+ z0_student = model(LR)
175
+ loss_distil = (z0_student - z0_teacher).abs().mean()
176
+ loss_adv = F.softplus(-model_D(
177
+ z0_student,
178
+ timesteps,
179
+ encoder_hidden_states=encoder_hidden_states,
180
+ return_dict=False,
181
+ )[0]).mean()
182
+ loss = loss_distil + loss_adv
183
+ optimizer.zero_grad(set_to_none=True)
184
+ scaler.scale(loss).backward()
185
+ scaler.step(optimizer)
186
+ scaler.update()
187
+ with torch.cuda.amp.autocast(enabled=True):
188
+ pred_real = model_D(
189
+ z0_gt.detach(),
190
+ timesteps,
191
+ encoder_hidden_states=encoder_hidden_states,
192
+ return_dict=False,
193
+ )[0]
194
+ pred_fake = model_D(
195
+ z0_student.detach(),
196
+ timesteps,
197
+ encoder_hidden_states=encoder_hidden_states,
198
+ return_dict=False,
199
+ )[0]
200
+ loss_D = F.softplus(pred_fake).mean() + F.softplus(-pred_real).mean()
201
+ optimizer_D.zero_grad(set_to_none=True)
202
+ scaler.scale(loss_D).backward()
203
+ scaler.step(optimizer_D)
204
+ scaler.update()
205
+ loss_avg += loss.item()
206
+ loss_distil_avg += loss_distil.item()
207
+ loss_adv_avg += loss_adv.item()
208
+ loss_D_avg += loss_D.item()
209
+ iter_num += 1
210
+ # print("loss", loss.item())
211
+ # print("loss_distil", loss_distil.item())
212
+ # print("loss_adv", loss_adv.item())
213
+ # print("loss_D", loss_D.item())
214
+ scheduler.step()
215
+ loss_avg /= iter_num
216
+ loss_distil_avg /= iter_num
217
+ loss_adv_avg /= iter_num
218
+ loss_D_avg /= iter_num
219
+ 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])
220
+ if rank == 0:
221
+ print(log_data)
222
+ with open(log_path, "a") as log_file:
223
+ log_file.write(log_data + "\n")
224
+ if epoch_i % args.save_interval == 0:
225
+ torch.save(model.state_dict(), "./%s/net_params_%d.pkl" % (model_dir, epoch_i))
official/utils.py ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from peft import LoraConfig
3
+
4
+ def add_lora_to_unet(unet, rank=4):
5
+ l_target_modules_encoder, l_target_modules_decoder, l_modules_others = [], [], []
6
+ 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"]
7
+ for n, p in unet.named_parameters():
8
+ check_flag = 0
9
+ if "bias" in n or "norm" in n:
10
+ continue
11
+ for pattern in l_grep:
12
+ if pattern in n and ("down_blocks" in n or "conv_in" in n):
13
+ l_target_modules_encoder.append(n.replace(".weight",""))
14
+ break
15
+ elif pattern in n and ("up_blocks" in n or "conv_out" in n):
16
+ l_target_modules_decoder.append(n.replace(".weight",""))
17
+ break
18
+ elif pattern in n:
19
+ l_modules_others.append(n.replace(".weight",""))
20
+ break
21
+ unet.add_adapter(LoraConfig(r=rank,init_lora_weights="gaussian",target_modules=l_target_modules_encoder), adapter_name="default_encoder")
22
+ unet.add_adapter(LoraConfig(r=rank,init_lora_weights="gaussian",target_modules=l_target_modules_decoder), adapter_name="default_decoder")
23
+ unet.add_adapter(LoraConfig(r=rank,init_lora_weights="gaussian",target_modules=l_modules_others), adapter_name="default_others")
24
+ return unet
requirements.txt ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # A100 训练环境依赖 (Python 3.10/3.11, CUDA 12.x)
2
+ torch>=2.1
3
+ torchvision>=0.16
4
+ diffusers>=0.27
5
+ accelerate>=0.28
6
+ omegaconf>=2.3
7
+ einops
8
+ lpips
9
+ pyiqa
10
+ opencv-python
11
+ scipy
12
+ numpy
13
+ pillow
14
+ safetensors
15
+ huggingface_hub
16
+ pyyaml
17
+ requests
scripts/archive_run.py ADDED
@@ -0,0 +1,31 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """归档: 把 logs/<run> 与权重副本整理到 archive/<date>_<name>/, 生成 summary.txt
3
+ 用法: python scripts/archive_run.py --name s2_best --src_log logs/s2 --weights weight/s2/net_params_*.pkl
4
+ """
5
+ import argparse, os, shutil, json, glob, datetime
6
+
7
+ def main():
8
+ ap = argparse.ArgumentParser()
9
+ ap.add_argument("--name", required=True)
10
+ ap.add_argument("--src_log", default="logs")
11
+ ap.add_argument("--weights", nargs="*", default=[])
12
+ ap.add_argument("--eval_json", default="", help="eval_val.py 输出")
13
+ args = ap.parse_args()
14
+ dst = os.path.join("archive", f"{datetime.date.today().isoformat()}_{args.name}")
15
+ os.makedirs(dst, exist_ok=True)
16
+ if os.path.isdir(args.src_log):
17
+ shutil.copytree(args.src_log, os.path.join(dst, "log"), dirs_exist_ok=True)
18
+ wd = os.path.join(dst, "weights"); os.makedirs(wd, exist_ok=True)
19
+ for pat in args.weights:
20
+ for f in glob.glob(pat):
21
+ shutil.copy2(f, os.path.join(wd, os.path.basename(f)))
22
+ 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)]}
23
+ if args.eval_json and os.path.exists(args.eval_json):
24
+ with open(args.eval_json, encoding="utf-8") as fh:
25
+ summary["eval_mean"] = json.load(fh).get("mean")
26
+ with open(os.path.join(dst, "summary.txt"), "w", encoding="utf-8") as fh:
27
+ fh.write(json.dumps(summary, ensure_ascii=False, indent=1))
28
+ print("archived to", dst)
29
+
30
+ if __name__ == "__main__":
31
+ main()
scripts/download_weights.sh ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ # Plan A: 服务器有外网时下载权重(在项目根执行)
3
+ # 双方案: 无外网时在本地跑同脚本并把 weight/ 与 models/ 传上来(见 pack_upload)
4
+ set -euo pipefail
5
+ cd "$(dirname "$0")/.."
6
+ export HF_ENDPOINT=https://hf-mirror.com
7
+ mkdir -p weight/pretrained models/stable-diffusion-2-1-base weight/gdpo
8
+ HF=https://huggingface.co/Guaishou74851/AdcSR/resolve/main
9
+ for f in weight/net_params_200.pkl weight/pretrained/DAPE.pth weight/pretrained/halfDecoder.ckpt \
10
+ weight/pretrained/osediff.pkl weight/pretrained/ram_swin_large_14m.pth; do
11
+ echo "==> $f"; curl -fL --retry 5 --retry-delay 3 -o "$f" "$HF/$f"
12
+ done
13
+ MS=https://modelscope.cn/models/AI-ModelScope/stable-diffusion-2-1-base/resolve/master
14
+ for f in model_index.json feature_extractor/preprocessor_config.json scheduler/scheduler_config.json \
15
+ tokenizer/merges.txt tokenizer/special_tokens_map.json tokenizer/tokenizer_config.json \
16
+ tokenizer/vocab.json text_encoder/config.json text_encoder/model.fp16.safetensors \
17
+ unet/config.json unet/diffusion_pytorch_model.fp16.safetensors \
18
+ vae/config.json vae/diffusion_pytorch_model.fp16.safetensors; do
19
+ echo "==> $f"; curl -fL --retry 5 --retry-delay 3 --create-dirs -o "models/stable-diffusion-2-1-base/$f" "$MS/$f"
20
+ done
21
+ # GDPO 教师 (S2 前下载; 失败可 --teacher osediff)
22
+ echo "==> GDPO (Joypop/GDPO)"
23
+ python - <<'EOF'
24
+ import os
25
+ try:
26
+ from huggingface_hub import snapshot_download
27
+ snapshot_download(repo_id="Joypop/GDPO", repo_type="model", local_dir="weight/gdpo", allow_patterns=["ckp/*", "ram/*"])
28
+ print("GDPO downloaded")
29
+ except Exception as e:
30
+ print("GDPO download failed (可回退 osediff):", e)
31
+ EOF
32
+ du -sh weight models
scripts/env_check.py ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """环境自检: GPU/CUDA/pyiqa/dwt/官方import 冒烟(不加载大权重)"""
3
+ import os, sys
4
+ sys.path.insert(0, "official"); sys.path.insert(0, "src")
5
+ import torch
6
+ print("torch", torch.__version__, "cuda", torch.cuda.is_available())
7
+ if torch.cuda.is_available():
8
+ print("gpu", torch.cuda.get_device_name(0), "mem_GB", round(torch.cuda.get_device_properties(0).total_memory/1e9,1))
9
+ try:
10
+ import pyiqa; print("pyiqa ok")
11
+ except Exception as e: print("pyiqa FAIL", e)
12
+ try:
13
+ import pywt; print("pywt ok")
14
+ except Exception as e: print("pywt missing(可选)", e)
15
+ try:
16
+ import diffusers, transformers, peft, omegaconf
17
+ print("diffusers/transformers/peft ok")
18
+ except Exception as e: print("libs FAIL", e)
19
+ # 官方模块可 import
20
+ try:
21
+ import model, dataset, utils, forward
22
+ print("official imports ok")
23
+ except Exception as e:
24
+ print("official imports FAIL", e)
25
+ print("env_check done")
scripts/probe_gdpo.py ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """GDPO 权重 probe: 打印键结构, 供 src/common.load_gdpo_teacher 排错。
3
+ 用法: python scripts/probe_gdpo.py weight/gdpo/ckp/diffusion_pytorch_model.safetensors
4
+ """
5
+ import sys
6
+ from pathlib import Path
7
+ sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
8
+ from common import probe_gdpo
9
+
10
+ if __name__ == "__main__":
11
+ if len(sys.argv) < 2:
12
+ print("用法: probe_gdpo.py <gdpo权重路径>"); sys.exit(1)
13
+ probe_gdpo(sys.argv[1])
scripts/run_smoke.sh ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ # S0 冒烟(每个训练配置 1 迭代)
3
+ set -euo pipefail
4
+ cd "$(dirname "$0")/.."
5
+ source "$(conda info --base)/etc/profile.d/conda.sh"; conda activate AdcSR
6
+ CUDA_VISIBLE_DEVICES=0 python src/train_lora.py --config configs/config_smoke.yml \
7
+ --manifest data/manifest_train.json --real_prob 0.0 --steps 2 --batch_size 1 --grad_accum 1 \
8
+ --init_net weight/net_params_200.pkl --out weight/smoke --log_dir logs/smoke_lora || true
9
+ CUDA_VISIBLE_DEVICES=0 python src/train_stage2.py --config configs/config_smoke.yml \
10
+ --steps 2 --batch_size 1 --grad_accum 1 --teacher osediff --skip_ram \
11
+ --init_net weight/net_params_200.pkl --out weight/smoke --log_dir logs/smoke_s2 || true
12
+ echo smoke done
scripts/run_stage1.sh ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ # S1 LoRA 像素域适配 (A100 40G)
3
+ set -euo pipefail
4
+ cd "$(dirname "$0")/.."
5
+ source "$(conda info --base)/etc/profile.d/conda.sh"; conda activate AdcSR
6
+ CUDA_VISIBLE_DEVICES=0 python src/train_lora.py \
7
+ --config configs/config_s1_lora.yml \
8
+ --manifest data/manifest_train.json --real_prob 0.35 \
9
+ --model_id models/stable-diffusion-2-1-base \
10
+ --half_decoder weight/pretrained/halfDecoder.ckpt \
11
+ --init_net weight/net_params_200.pkl \
12
+ --out weight/s1 --log_dir logs/s1 \
13
+ --steps 20000 --batch_size 8 --grad_accum 2 --lr 5e-5 --lora_rank 64 \
14
+ --save_every 2000 --w_l1 1.0 --w_lpips 1.0 --w_dists 0.3 --w_wave 0.5 --w_color 0.2
scripts/run_stage2.sh ADDED
@@ -0,0 +1,18 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ # S2 GDPO 教师蒸馏(先 probe; GDPO 不可用回退 osediff)
3
+ set -euo pipefail
4
+ cd "$(dirname "$0")/.."
5
+ source "$(conda info --base)/etc/profile.d/conda.sh"; conda activate AdcSR
6
+ TEACHER="${1:-gdpo}"
7
+ if [ "$TEACHER" = "gdpo" ] && [ ! -d weight/gdpo ]; then
8
+ echo "未找到 weight/gdpo, 回退 osediff"; TEACHER=osediff
9
+ fi
10
+ CUDA_VISIBLE_DEVICES=0 python src/train_stage2.py \
11
+ --config configs/config_s2_distill.yml \
12
+ --model_id models/stable-diffusion-2-1-base \
13
+ --half_decoder weight/pretrained/halfDecoder.ckpt \
14
+ --teacher "$TEACHER" --gdpo_dir weight/gdpo \
15
+ --init_net weight/s1/net_params_BEST.pkl \
16
+ --out weight/s2 --log_dir logs/s2 \
17
+ --steps 30000 --batch_size 8 --grad_accum 2 --lr 5e-5 --lr_D 1e-6 \
18
+ --save_every 2000 --w_wave 0.5 --w_lpips 0.5 --w_dists 0.3 --w_color 0.2
scripts/run_stage3.sh ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ # S3 域微调(S2 最优权重, 小 LR)
3
+ set -euo pipefail
4
+ cd "$(dirname "$0")/.."
5
+ source "$(conda info --base)/etc/profile.d/conda.sh"; conda activate AdcSR
6
+ CUDA_VISIBLE_DEVICES=0 python src/train_stage2.py \
7
+ --config configs/config_s3_ft.yml \
8
+ --model_id models/stable-diffusion-2-1-base \
9
+ --half_decoder weight/pretrained/halfDecoder.ckpt \
10
+ --teacher gdpo --gdpo_dir weight/gdpo \
11
+ --init_net weight/s2/net_params_BEST.pkl \
12
+ --out weight/s3 --log_dir logs/s3 \
13
+ --steps 20000 --batch_size 8 --grad_accum 2 --lr 2e-5 --lr_D 5e-7 \
14
+ --save_every 2000 --w_wave 0.5 --w_lpips 0.5 --w_dists 0.3 --w_color 0.2
scripts/setup_env.sh ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ # aspire2a 环境准备(A100 40G 单卡; python 3.10)
3
+ set -euo pipefail
4
+ cd "$(dirname "$0")/.."
5
+ if ! command -v conda >/dev/null 2>&1; then echo "需要 conda"; exit 1; fi
6
+ conda create -n AdcSR python=3.10 -y || true
7
+ source "$(conda info --base)/etc/profile.d/conda.sh"
8
+ conda activate AdcSR
9
+ pip install --upgrade pip
10
+ pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121
11
+ pip install -r official/requirements.txt
12
+ pip install diffusers transformers accelerate peft omegaconf tqdm pyyaml pyiqa lpips scikit-image imageio pillow numpy
13
+ pip install huggingface_hub hf_transfer safetensors gdown
14
+ pip install "dwt-pytorch" 2>/dev/null || pip install pywavelets || true
15
+ python - <<'EOF'
16
+ try:
17
+ import pywt; print("pywt ok")
18
+ except Exception as e:
19
+ print("pywt missing:", e)
20
+ EOF
21
+ echo "setup done"
src/check_submission.py ADDED
@@ -0,0 +1,48 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """提交包检查: 命名/数量/大小/结构/可加载/确定性。
3
+ 用法: python src/check_submission.py --zip xxx.zip [--names data/test_names.txt] [--test_runner]
4
+ """
5
+ import argparse, json, os, sys, tempfile, zipfile
6
+ import torch
7
+
8
+ def main():
9
+ ap = argparse.ArgumentParser()
10
+ ap.add_argument("--zip", required=True)
11
+ ap.add_argument("--names", default="", help="每行一个期望文件名")
12
+ ap.add_argument("--test_runner", action="store_true", help="解压并实际加载 jit 跑 1 次")
13
+ args = ap.parse_args()
14
+ size = os.path.getsize(args.zip)
15
+ checks = {"exists": True, "size_gb": round(size / 1e9, 2), "le_10gb": size <= 10e9}
16
+ with zipfile.ZipFile(args.zip) as z:
17
+ names = z.namelist()
18
+ outs = [n for n in names if "/output_dir/" in n or n.startswith("output_dir/")]
19
+ jpgs = [n for n in outs if n.lower().endswith(".jpg")]
20
+ mods = [n for n in names if "/model_dir/" in n or n.startswith("model_dir/")]
21
+ checks["output_jpg_count"] = len(jpgs)
22
+ checks["has_model"] = any(n.endswith(".pt") for n in mods)
23
+ checks["has_runner"] = any(n.endswith("runner.py") for n in mods)
24
+ if args.names:
25
+ expect = [l.strip() for l in open(args.names, encoding="utf-8") if l.strip()]
26
+ got = {os.path.basename(n) for n in jpgs}
27
+ checks["missing"] = [e for e in expect if e not in got]
28
+ checks["extra"] = sorted(got - set(expect))[:10]
29
+ if args.test_runner:
30
+ tmp = tempfile.mkdtemp()
31
+ z.extractall(tmp)
32
+ pt = [os.path.join(tmp, n) for n in mods if n.endswith(".pt")][0]
33
+ m = torch.jit.load(pt, map_location="cpu")
34
+ m.eval()
35
+ x = torch.randn(1, 3, 512, 512).half() * 0.5
36
+ with torch.no_grad():
37
+ o1 = m(x); o2 = m(x)
38
+ checks["jit_load_ok"] = True
39
+ checks["out_shape"] = list(o1.shape)
40
+ checks["deterministic_maxdiff"] = float((o1 - o2).abs().max())
41
+ for k, v in checks.items():
42
+ print(f"{k}: {v}")
43
+ bad = [k for k, v in checks.items() if v is False] or \
44
+ (checks.get("missing") if checks.get("missing") else [])
45
+ print("PASS" if not bad else f"FAIL items: {bad}")
46
+
47
+ if __name__ == "__main__":
48
+ main()
src/common.py ADDED
@@ -0,0 +1,213 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """AdcSR 工程公共工具:路径、模型装配、手工 LoRA 注入、教师加载、GDPO probe。
3
+ 训练脚本统一从这里 import,禁止各自重复实现装配逻辑。
4
+ """
5
+ import os, sys, copy, json, types
6
+ from pathlib import Path
7
+
8
+ REPO = Path(__file__).resolve().parents[1]
9
+ OFFICIAL = REPO / "official"
10
+
11
+ def ensure_official():
12
+ if str(OFFICIAL) not in sys.path:
13
+ sys.path.insert(0, str(OFFICIAL))
14
+
15
+ ensure_official()
16
+
17
+ import torch
18
+ import torch.nn as nn
19
+ import torch.nn.functional as F
20
+
21
+ # ---------------------------------------------------------------------------
22
+ # 模型装配(与 official/test.py 全链一致)
23
+ # ---------------------------------------------------------------------------
24
+ def load_diffusers_sd(model_id, dtype=torch.float32, device="cpu"):
25
+ from diffusers import StableDiffusionPipeline
26
+ pipe = StableDiffusionPipeline.from_pretrained(model_id, torch_dtype=dtype).to(device)
27
+ return pipe.vae, pipe.unet, pipe.text_encoder, pipe.tokenizer
28
+
29
+ def load_pruned_decoder(half_decoder_ckpt, device="cpu", dtype=torch.float32):
30
+ from diffusers.models.autoencoders.vae import Decoder
31
+ decoder = Decoder(in_channels=4, out_channels=3,
32
+ up_block_types=["UpDecoderBlock2D"] * 4,
33
+ block_out_channels=[64, 128, 256, 256], layers_per_block=2,
34
+ norm_num_groups=32, act_fn="silu", norm_type="group",
35
+ mid_block_add_attention=True).to(device=device, dtype=dtype)
36
+ ckpt = torch.load(half_decoder_ckpt, map_location="cpu", weights_only=False)
37
+ sd = {k.replace("decoder.", ""): v for k, v in ckpt["state_dict"].items() if k.startswith("decoder.")}
38
+ decoder.load_state_dict(sd, strict=True)
39
+ return decoder
40
+
41
+ def build_net(unet, decoder):
42
+ from model import Net # official
43
+ return Net(unet, copy.deepcopy(decoder))
44
+
45
+ def assemble_full_student(unet, decoder, net_weights=None, device="cuda", dtype=torch.float32):
46
+ """???? 512 ???: Net(unet,decoder) + decoder ??(up_blocks...conv_out)?"""
47
+ net = build_net(unet, decoder)
48
+ if net_weights is not None:
49
+ sd = torch.load(net_weights, map_location="cpu", weights_only=False)
50
+ if any(k.startswith("module.") for k in sd):
51
+ sd = {k.replace("module.", "", 1): v for k, v in sd.items()}
52
+ net.load_state_dict(sd, strict=True)
53
+ net.to(device=device, dtype=dtype)
54
+ tail_mods = [decoder.up_blocks, decoder.conv_norm_out, decoder.conv_act, decoder.conv_out]
55
+ for m in tail_mods:
56
+ m.to(device=device, dtype=dtype)
57
+ full = nn.Sequential(net, *tail_mods)
58
+ return full
59
+
60
+
61
+ def build_discriminator_unet(unet_copy, rank=4, dtype=torch.float32, device="cuda"):
62
+ """official 判别器: conv_in 4->256 + LoRA(unet)。"""
63
+ from utils import add_lora_to_unet
64
+ unet_D = copy.deepcopy(unet_copy).to(device=device, dtype=dtype)
65
+ cin = unet_D.conv_in
66
+ new_conv_in = nn.Conv2d(256, cin.out_channels, 3, padding=1).to(device=device, dtype=dtype)
67
+ new_conv_in.weight.data = cin.weight.data.repeat(1, 64, 1, 1) / 64
68
+ new_conv_in.bias.data = cin.bias.data
69
+ unet_D.conv_in = new_conv_in
70
+ unet_D = add_lora_to_unet(unet_D, rank=rank)
71
+ unet_D.set_adapters(["default_encoder", "default_decoder", "default_others"])
72
+ return unet_D
73
+
74
+ # ---------------------------------------------------------------------------
75
+ # 教师加载
76
+ # ---------------------------------------------------------------------------
77
+ def load_osediff_teacher(osediff_pkl, device="cuda", dtype=torch.float32):
78
+ ckpt = torch.load(osediff_pkl, map_location="cpu", weights_only=False)
79
+ return ckpt # {"vae":..., "unet":...}
80
+
81
+ def load_gdpo_teacher(gdpo_dir, device="cuda", dtype=torch.float32):
82
+ """GDPO ???????????????????????? probe_gdpo?
83
+ ?? diffusers UNet2DConditionModel?state dict ????? dict?"""
84
+ from diffusers import UNet2DConditionModel
85
+ if os.path.isdir(os.path.join(gdpo_dir, "unet")) and os.path.exists(
86
+ os.path.join(gdpo_dir, "unet", "diffusion_pytorch_model.safetensors")):
87
+ return UNet2DConditionModel.from_pretrained(os.path.join(gdpo_dir, "unet"),
88
+ torch_dtype=dtype).to(device)
89
+ if os.path.isdir(gdpo_dir):
90
+ if os.path.exists(os.path.join(gdpo_dir, "diffusion_pytorch_model.safetensors")):
91
+ p = os.path.join(gdpo_dir, "diffusion_pytorch_model.safetensors")
92
+ else:
93
+ p = os.path.join(gdpo_dir, "ckp", "diffusion_pytorch_model.safetensors")
94
+ if os.path.exists(p):
95
+ try:
96
+ return UNet2DConditionModel.from_pretrained(os.path.dirname(p),
97
+ torch_dtype=dtype).to(device)
98
+ except Exception:
99
+ return _load_raw(p)
100
+ if os.path.isfile(gdpo_dir):
101
+ return _load_raw(gdpo_dir)
102
+ raise RuntimeError("GDPO ??????????? python -m src.common --probe_gdpo <path> ?????"
103
+ "??? --teacher osediff")
104
+
105
+ def _load_raw(p):
106
+ if p.endswith(".safetensors"):
107
+ from safetensors.torch import load_file
108
+ return load_file(p)
109
+ return torch.load(p, map_location="cpu", weights_only=False)
110
+
111
+ def probe_gdpo(gdpo_path):
112
+ """打印权重键结构与前缀,帮助实现 GDPO->diffusers UNet 映射。"""
113
+ if gdpo_path.endswith(".safetensors"):
114
+ from safetensors.torch import load_file
115
+ sd = load_file(gdpo_path)
116
+ else:
117
+ sd = torch.load(gdpo_path, map_location="cpu", weights_only=False)
118
+ if isinstance(sd, dict) and "state_dict" in sd:
119
+ sd = sd["state_dict"]
120
+ keys = list(sd.keys())
121
+ print("num keys:", len(keys))
122
+ for k in keys[:40]:
123
+ print(k, tuple(sd[k].shape) if hasattr(sd[k], "shape") else type(sd[k]))
124
+ # 判断是否为完整 UNet(含 down_blocks)或 LoRA 或 Pipeline
125
+ has_unet = any("down_blocks" in k for k in keys)
126
+ has_lora = any("lora" in k.lower() for k in keys)
127
+ print("has_unet_blocks:", has_unet, "| has_lora:", has_lora)
128
+
129
+ # ---------------------------------------------------------------------------
130
+ # 手工 LoRA(对任意 Conv2d/Linear 注入,规避 peft 在剪枝/删模块后的解析问题)
131
+ # ---------------------------------------------------------------------------
132
+ class LoRAConv2d(nn.Module):
133
+ def __init__(self, conv: nn.Conv2d, r: int, alpha: float = 1.0):
134
+ super().__init__()
135
+ self.conv = conv
136
+ self.r = max(1, r)
137
+ self.alpha = alpha
138
+ cin, cout = conv.in_channels, conv.out_channels
139
+ self.lora_a = nn.Parameter(torch.zeros(cin, self.r))
140
+ self.lora_b = nn.Parameter(torch.zeros(self.r, cout))
141
+ nn.init.kaiming_uniform_(self.lora_a, a=5 ** 0.5)
142
+ nn.init.zeros_(self.lora_b)
143
+ self.requires_grad_(False)
144
+ self.lora_a.requires_grad_(True)
145
+ self.lora_b.requires_grad_(True)
146
+
147
+ def forward(self, x):
148
+ y = self.conv(x)
149
+ if self.training or True:
150
+ # 1x1 conv low-rank: 输入cin->r->cout, 保持空间尺寸
151
+ z = F.conv2d(x, self.lora_a.t().view(self.r, cin, 1, 1))
152
+ z = F.conv2d(z, self.lora_b.t().view(cout, self.r, 1, 1))
153
+ return y + self.alpha * z
154
+ return y
155
+
156
+ class LoRALinear(nn.Module):
157
+ def __init__(self, lin: nn.Linear, r: int, alpha: float = 1.0):
158
+ super().__init__()
159
+ self.lin = lin
160
+ self.r = max(1, r)
161
+ self.alpha = alpha
162
+ cin, cout = lin.in_features, lin.out_features
163
+ self.lora_a = nn.Parameter(torch.zeros(cin, self.r))
164
+ self.lora_b = nn.Parameter(torch.zeros(self.r, cout))
165
+ nn.init.kaiming_uniform_(self.lora_a, a=5 ** 0.5)
166
+ nn.init.zeros_(self.lora_b)
167
+ self.requires_grad_(False)
168
+ self.lora_a.requires_grad_(True)
169
+ self.lora_b.requires_grad_(True)
170
+
171
+ def forward(self, x):
172
+ y = self.lin(x)
173
+ z = F.linear(x, self.lora_a.t())
174
+ z = F.linear(z, self.lora_b.t())
175
+ return y + self.alpha * z
176
+
177
+ def _names(model):
178
+ for n, m in model.named_modules():
179
+ if isinstance(m, (nn.Conv2d, nn.Linear)):
180
+ yield n, m
181
+
182
+ def inject_lora(model, rank=64, alpha=1.0, skip_bias_norm=True, include=("conv", "to_q", "to_k", "to_v", "proj", "ff", "linear")):
183
+ """替换模型内所有 Conv2d/Linear 为 LoRA 包装(原始权重冻结,仅训练 lora_a/b)。
184
+ include: 子串过滤,None=全部。"""
185
+ for n, m in list(_names(model)):
186
+ if include is not None and not any(s in n for s in include):
187
+ continue
188
+ parent, attr = _find_parent(model, n)
189
+ if isinstance(m, nn.Conv2d):
190
+ setattr(parent, attr, LoRAConv2d(m, rank, alpha))
191
+ elif isinstance(m, nn.Linear):
192
+ setattr(parent, attr, LoRALinear(m, rank, alpha))
193
+ return model
194
+
195
+ def _find_parent(model, name):
196
+ parts = name.split(".")
197
+ node = model
198
+ for p in parts[:-1]:
199
+ node = getattr(node, p)
200
+ return node, parts[-1]
201
+
202
+ def lora_params(model):
203
+ for p in model.parameters():
204
+ if p.requires_grad:
205
+ yield p
206
+
207
+ def count_params(model, only_trainable=False):
208
+ if only_trainable:
209
+ return sum(p.numel() for p in model.parameters() if p.requires_grad)
210
+ return sum(p.numel() for p in model.parameters())
211
+
212
+ if __name__ == "__main__":
213
+ print("common module OK; official dir:", OFFICIAL)
src/eval_val.py ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """Val 评测 + 代理分: 在隔离 val(真实对/官方对) 上计算 8 指标 + proxy + 保真护栏。
3
+ 用法:
4
+ python src/eval_val.py --weights weight/s2/net_params_X.pkl --pairs_json data/manifest_val.json
5
+ python src/eval_val.py --weights ... --official_lq <val lq dir> --official_gt <val gt dir>
6
+ """
7
+ import argparse, json, os, sys
8
+ from pathlib import Path
9
+ REPO = Path(__file__).resolve().parents[1]
10
+ sys.path.insert(0, str(REPO)); sys.path.insert(0, str(REPO / "src")); sys.path.insert(0, str(REPO / "official"))
11
+ import numpy as np, torch
12
+ from PIL import Image
13
+ from inference_4k import build_model, infer_4k_single
14
+
15
+ def load_metrics(device):
16
+ import pyiqa
17
+ return {k: pyiqa.create_metric(k, device=device) for k in
18
+ ["psnr", "ssim", "lpips", "dists", "niqe", "maniqa", "musiq", "clipiqa"]}
19
+
20
+ def resize_pair(sr, hr, max_side=1024):
21
+ sr = sr.convert("RGB"); hr = hr.convert("RGB")
22
+ for im in (sr, hr):
23
+ w, h = im.size
24
+ if max(w, h) > max_side:
25
+ s = max_side / max(w, h)
26
+ im.thumbnail((int(w * s), int(h * s)), Image.LANCZOS)
27
+ return sr, hr
28
+
29
+ def metric_dict(metrics, sr, hr, device):
30
+ import torchvision.transforms as T
31
+ def t(im):
32
+ return T.ToTensor()(im).to(device).unsqueeze(0)
33
+ a, b = t(sr), t(hr)
34
+ out = {"psnr": metrics["psnr"](a, b).item(), "ssim": metrics["ssim"](a, b).item(),
35
+ "lpips": metrics["lpips"](a, b).item(), "dists": metrics["dists"](a, b).item(),
36
+ "niqe": metrics["niqe"](a).item(), "maniqa": metrics["maniqa"](a).item(),
37
+ "musiq": metrics["musiq"](a).item(), "clipiqa": metrics["clipiqa"](a).item()}
38
+ return out
39
+
40
+ def proxy(m):
41
+ return (0.25 * m["clipiqa"] + 0.2 * (m["musiq"] / 100.0) + 0.15 * m["maniqa"]
42
+ + 0.15 * (1 - m["niqe"] / 10.0) + 0.25 * (1 - m["lpips"]))
43
+
44
+ def main():
45
+ ap = argparse.ArgumentParser()
46
+ ap.add_argument("--weights", required=True)
47
+ ap.add_argument("--half_decoder", default="weight/pretrained/halfDecoder.ckpt")
48
+ ap.add_argument("--model_id", default="models/stable-diffusion-2-1-base")
49
+ ap.add_argument("--pairs_json", default="", help='[{lr,hr}] 或 manifest 结构')
50
+ ap.add_argument("--official_lq", default="", help="同分辨率官方对(整图滑窗)")
51
+ ap.add_argument("--official_gt", default="")
52
+ ap.add_argument("--out", default="logs/eval_result.json")
53
+ ap.add_argument("--max_side", type=int, default=1024)
54
+ args = ap.parse_args()
55
+ device = "cuda" if torch.cuda.is_available() else "cpu"
56
+ net, tail = build_model(args.weights, args.half_decoder, args.model_id, device, bf16=True)
57
+ metrics = load_metrics(device)
58
+ rows = []
59
+ if args.pairs_json:
60
+ with open(args.pairs_json, encoding="utf-8") as fh:
61
+ data = json.load(fh)
62
+ pairs = data.get("real_pairs", data if isinstance(data, list) else [])
63
+ for i, p in enumerate(pairs):
64
+ with torch.no_grad():
65
+ lr = Image.open(p["lr"]).convert("RGB")
66
+ hr = Image.open(p["hr"]).convert("RGB")
67
+ lr128 = lr.resize((max(1, lr.width // 4), max(1, lr.height // 4)), Image.LANCZOS)
68
+ t = torch.from_numpy(np.asarray(lr128, dtype=np.float32).transpose(2, 0, 1) / 255.0 * 2 - 1)[None].to(device)
69
+ with torch.autocast("cuda", dtype=torch.bfloat16):
70
+ z = net(t); sr_arr = tail(z)
71
+ sr = Image.fromarray(((sr_arr[0].float().cpu().numpy().transpose(1, 2, 0) + 1) / 2 * 255).clip(0, 255).astype(np.uint8))
72
+ sr, hr = resize_pair(sr, hr, args.max_side)
73
+ m = metric_dict(metrics, sr, hr, device)
74
+ m["name"] = os.path.basename(p["lr"]); m["proxy"] = proxy(m)
75
+ rows.append(m)
76
+ print(f" [{i+1}] {m['name']} proxy {m['proxy']:.4f}", flush=True)
77
+ if args.official_lq and args.official_gt:
78
+ lq_files = sorted(os.listdir(args.official_lq))
79
+ for f in lq_files:
80
+ lq_p = os.path.join(args.official_lq, f)
81
+ gt_p = os.path.join(args.official_gt, f.replace("_lq", "_gt").replace("lq.jpg", "gt.png"))
82
+ if not os.path.exists(gt_p):
83
+ gt_p = os.path.join(args.official_gt, f.replace("_lq.jpg", "_gt.jpg"))
84
+ if not os.path.exists(gt_p):
85
+ continue
86
+ sr = infer_4k_single(lq_p, net, tail, device)
87
+ gt = Image.open(gt_p).convert("RGB")
88
+ sr, gt = resize_pair(sr, gt, args.max_side)
89
+ m = metric_dict(metrics, sr, gt, device)
90
+ m["name"] = f; m["proxy"] = proxy(m)
91
+ rows.append(m)
92
+ print(f" [official] {f} proxy {m['proxy']:.4f}", flush=True)
93
+ if not rows:
94
+ print("无评测样本"); return
95
+ keys = ["psnr", "ssim", "lpips", "dists", "niqe", "maniqa", "musiq", "clipiqa", "proxy"]
96
+ agg = {k: float(np.mean([r[k] for r in rows])) for k in keys}
97
+ result = {"per_image": rows, "mean": agg}
98
+ os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True)
99
+ with open(args.out, "w", encoding="utf-8") as fh:
100
+ json.dump(result, fh, ensure_ascii=False, indent=1)
101
+ print(json.dumps(agg, indent=1))
102
+
103
+ if __name__ == "__main__":
104
+ main()
src/export_jit.py ADDED
@@ -0,0 +1,62 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """导出 torch.jit(fp16) 512x512 模型: 输入 LR[1,3,512,512][-1,1] -> 输出同尺寸。
3
+ 内部: bicubic 512->128 -> 官方 4x 学生全链 -> 512。 (AdaIN 后处理不进测速模型)
4
+ 用法: python src/export_jit.py --net weight/s2/net_params_X.pkl --out model_dir
5
+ """
6
+ import argparse, os, sys
7
+ from pathlib import Path
8
+ REPO = Path(__file__).resolve().parents[1]
9
+ sys.path.insert(0, str(REPO)); sys.path.insert(0, str(REPO / "src")); sys.path.insert(0, str(REPO / "official"))
10
+ import torch, torch.nn as nn, torch.nn.functional as F
11
+ from common import load_diffusers_sd, load_pruned_decoder, build_net
12
+
13
+ class SR512(nn.Module):
14
+ def __init__(self, net, tail):
15
+ super().__init__()
16
+ self.net = net
17
+ self.tail = tail
18
+
19
+ def forward(self, x512):
20
+ x128 = F.interpolate(x512, size=(128, 128), mode="bicubic", align_corners=False)
21
+ z = self.net(x128)
22
+ return self.tail(z)
23
+
24
+ def main():
25
+ ap = argparse.ArgumentParser()
26
+ ap.add_argument("--net", required=True)
27
+ ap.add_argument("--out", default="model_dir")
28
+ ap.add_argument("--half_decoder", default="weight/pretrained/halfDecoder.ckpt")
29
+ ap.add_argument("--model_id", default="models/stable-diffusion-2-1-base")
30
+ ap.add_argument("--name", default="your_model.pt")
31
+ args = ap.parse_args()
32
+ os.makedirs(args.out, exist_ok=True)
33
+ device = "cuda" if torch.cuda.is_available() else "cpu"
34
+ vae, unet, _, _ = load_diffusers_sd(args.model_id, dtype=torch.float32, device="cpu")
35
+ del vae
36
+ decoder = load_pruned_decoder(args.half_decoder, device="cpu", dtype=torch.float32)
37
+ net = build_net(unet, decoder)
38
+ sd = torch.load(args.net, map_location="cpu", weights_only=False)
39
+ if any(k.startswith("module.") for k in sd):
40
+ sd = {k.replace("module.", "", 1): v for k, v in sd.items()}
41
+ net.load_state_dict(sd, strict=True)
42
+ net.eval()
43
+ tail = nn.Sequential(*decoder.up_blocks, decoder.conv_norm_out,
44
+ decoder.conv_act, decoder.conv_out).eval()
45
+ model = SR512(net, tail).to(device).half().eval()
46
+ # trace on fixed 512 fp16
47
+ dummy = torch.randn(1, 3, 512, 512, device=device).half() * 0.5
48
+ with torch.no_grad():
49
+ traced = torch.jit.trace(model, dummy, check_trace=False)
50
+ traced = torch.jit.freeze(traced)
51
+ out_path = os.path.join(args.out, args.name)
52
+ traced.save(out_path)
53
+ # 自检: 两次前向一致性 + 形状
54
+ with torch.no_grad():
55
+ o1 = traced(dummy); o2 = traced(dummy)
56
+ assert o1.shape == dummy.shape, o1.shape
57
+ err = (o1 - o2).abs().max().item()
58
+ print("saved", out_path, "| deterministic max-diff:", err)
59
+ print("output range sample:", float(o1.min()), float(o1.max()))
60
+
61
+ if __name__ == "__main__":
62
+ main()
src/inference_4k.py ADDED
@@ -0,0 +1,105 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """4K 整图增强: 512 窗 + overlap 128 线性融合 + AdaIN 对齐 LR(禁止 zero-padding, 边缘 reflect)。
3
+ x1 语义: 对每个 512 窗先缩到 128, 走官方 4x 学生模型, 回 512。
4
+ 用法: python src/inference_4k.py --lr_dir <dir> --net weight/s2/net_params_X.pkl --out output_dir
5
+ """
6
+ import argparse, copy, os, sys
7
+ from pathlib import Path
8
+ REPO = Path(__file__).resolve().parents[1]
9
+ sys.path.insert(0, str(REPO)); sys.path.insert(0, str(REPO / "src")); sys.path.insert(0, str(REPO / "official"))
10
+ import torch, torch.nn.functional as F
11
+ import numpy as np
12
+ from PIL import Image
13
+ from common import load_diffusers_sd, load_pruned_decoder, build_net
14
+
15
+ def build_model(net_pkl, half_decoder, model_id, device, bf16=True):
16
+ dtype = torch.bfloat16 if (bf16 and device == "cuda") else torch.float32
17
+ vae, unet, _, _ = load_diffusers_sd(model_id, dtype=torch.float32, device="cpu")
18
+ del vae
19
+ decoder = load_pruned_decoder(half_decoder, device="cpu", dtype=torch.float32)
20
+ net = build_net(unet, decoder)
21
+ net.to(device=device, dtype=dtype)
22
+ sd = torch.load(net_pkl, map_location="cpu", weights_only=False)
23
+ if any(k.startswith("module.") for k in sd):
24
+ sd = {k.replace("module.", "", 1): v for k, v in sd.items()}
25
+ net.load_state_dict(sd, strict=True)
26
+ net.eval()
27
+ tail = torch.nn.Sequential(*decoder.up_blocks, decoder.conv_norm_out,
28
+ decoder.conv_act, decoder.conv_out).to(device=device, dtype=dtype).eval()
29
+ return net, tail
30
+
31
+ def infer_tile(net, tail, lr128, lr_stats):
32
+ """lr128: [1,3,128,128] [-1,1]; 返回 [1,3,512,512] [-1,1] 并 AdaIN 对齐 lr_stats"""
33
+ with torch.no_grad():
34
+ z = net(lr128)
35
+ out = tail(z)
36
+ out = out.float()
37
+ mu, std = out.mean(dim=(2, 3), keepdim=True), out.std(dim=(2, 3), keepdim=True)
38
+ lmu, lstd = lr_stats
39
+ out = (out - mu) / (std + 1e-6) * lstd + lmu
40
+ return out.clamp(-1, 1)
41
+
42
+ def infer_4k_single(lr_path, net, tail, device, tile=512, overlap=128, bf16=True):
43
+ im = Image.open(lr_path).convert("RGB")
44
+ w, h = im.size
45
+ a = np.asarray(im, dtype=np.float32) / 255.0 # [H,W,3]
46
+ # reflect pad so tile 整除
47
+ stride = tile - overlap
48
+ pad_r = (stride - w % stride) % stride + (tile - stride) if (w % stride) else max(0, tile - stride)
49
+ pad_b = (stride - h % stride) % stride + (tile - stride) if (h % stride) else max(0, tile - stride)
50
+ a = np.pad(a, ((0, pad_b), (0, pad_r), (0, 0)), mode="reflect")
51
+ H, W = a.shape[:2]
52
+ acc = np.zeros((H, W, 3), dtype=np.float64)
53
+ wsum = np.zeros((H, W, 1), dtype=np.float64)
54
+ rows = list(range(0, H - tile + 1, stride)) or [0]
55
+ cols = list(range(0, W - tile + 1, stride)) or [0]
56
+ if rows[-1] + tile < H: rows.append(H - tile)
57
+ if cols[-1] + tile < W: cols.append(W - tile)
58
+ # 线性窗权重
59
+ ramp = np.minimum(np.arange(tile), np.arange(tile)[::-1]) / (tile // 2)
60
+ w2d = (ramp[None, :] * ramp[:, None])[..., None].astype(np.float64)
61
+ for y in rows:
62
+ for x in cols:
63
+ crop = a[y:y + tile, x:x + tile]
64
+ lr128 = np.asarray(Image.fromarray((crop * 255).astype(np.uint8)).resize((128, 128), Image.BICUBIC),
65
+ dtype=np.float32) / 255.0
66
+ lr_t = torch.from_numpy(lr128.transpose(2, 0, 1))[None].to(device) * 2 - 1
67
+ stats = (torch.tensor(crop.mean(axis=(0, 1))[None, :, None, None], device=device) * 2 - 1,
68
+ torch.tensor(crop.std(axis=(0, 1))[None, :, None, None], device=device))
69
+ with torch.autocast("cuda", dtype=torch.bfloat16, enabled=bf16):
70
+ out = infer_tile(net, tail, lr_t, stats)
71
+ o = out[0].float().cpu().numpy().transpose(1, 2, 0) # [512,512,3] [-1,1]
72
+ o = (o + 1) / 2.0
73
+ acc[y:y + tile, x:x + tile] += o.astype(np.float64) * w2d
74
+ wsum[y:y + tile, x:x + tile] += w2d
75
+ res = acc / (wsum + 1e-9)
76
+ res = res[:h, :w]
77
+ return Image.fromarray((np.clip(res, 0, 1) * 255).astype(np.uint8))
78
+
79
+ def main():
80
+ ap = argparse.ArgumentParser()
81
+ ap.add_argument("--lr_dir", required=True)
82
+ ap.add_argument("--net", required=True)
83
+ ap.add_argument("--out", required=True)
84
+ ap.add_argument("--half_decoder", default="weight/pretrained/halfDecoder.ckpt")
85
+ ap.add_argument("--model_id", default="models/stable-diffusion-2-1-base")
86
+ ap.add_argument("--tile", type=int, default=512)
87
+ ap.add_argument("--overlap", type=int, default=128)
88
+ ap.add_argument("--no_bf16", action="store_true")
89
+ args = ap.parse_args()
90
+ device = "cuda" if torch.cuda.is_available() else "cpu"
91
+ os.makedirs(args.out, exist_ok=True)
92
+ net, tail = build_model(args.net, args.half_decoder, args.model_id, device, bf16=not args.no_bf16)
93
+ files = sorted([f for f in os.listdir(args.lr_dir) if f.lower().endswith((".jpg", ".jpeg", ".png"))])
94
+ import time
95
+ t0 = time.time()
96
+ for i, f in enumerate(files):
97
+ out_im = infer_4k_single(os.path.join(args.lr_dir, f), net, tail, device,
98
+ tile=args.tile, overlap=args.overlap, bf16=not args.no_bf16)
99
+ out_im.save(os.path.join(args.out, f), quality=95)
100
+ if (i + 1) % 10 == 0:
101
+ print(f" {i+1}/{len(files)} elapsed {time.time()-t0:.1f}s", flush=True)
102
+ print(f"done {len(files)} imgs in {time.time()-t0:.1f}s")
103
+
104
+ if __name__ == "__main__":
105
+ main()
src/latency_test.py ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """512x512 fp16 时延测速: 加载 torch.jit 模型, warmup + N 次前向。
3
+ 用法: python src/latency_test.py --model model_dir/your_model.pt [--osediff_latency 0.168]
4
+ """
5
+ import argparse, time, statistics
6
+ import torch
7
+
8
+ def main():
9
+ ap = argparse.ArgumentParser()
10
+ ap.add_argument("--model", required=True)
11
+ ap.add_argument("--n", type=int, default=100)
12
+ ap.add_argument("--warmup", type=int, default=10)
13
+ ap.add_argument("--osediff_latency", type=float, default=0.0, help="秒; 提供则打印加速比")
14
+ args = ap.parse_args()
15
+ device = "cuda" if torch.cuda.is_available() else "cpu"
16
+ m = torch.jit.load(args.model, map_location=device)
17
+ m.eval()
18
+ x = torch.randn(1, 3, 512, 512, device=device).half()
19
+ with torch.no_grad():
20
+ for _ in range(args.warmup):
21
+ m(x)
22
+ torch.cuda.synchronize() if device == "cuda" else None
23
+ times = []
24
+ for _ in range(args.n):
25
+ if device == "cuda":
26
+ torch.cuda.synchronize()
27
+ t0 = time.perf_counter()
28
+ with torch.no_grad():
29
+ m(x)
30
+ if device == "cuda":
31
+ torch.cuda.synchronize()
32
+ times.append(time.perf_counter() - t0)
33
+ mean = statistics.mean(times)
34
+ med = statistics.median(times)
35
+ print(f"mean {mean*1000:.3f} ms | median {med*1000:.3f} ms | n={args.n}")
36
+ if args.osediff_latency > 0:
37
+ print(f"speedup vs OSEDiff({args.osediff_latency*1000:.1f}ms): {args.osediff_latency/mean:.2f}x")
38
+
39
+ if __name__ == "__main__":
40
+ main()