CSIGv3 AdcSR train scripts + A100 runbook
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- NOTES.md +107 -0
- README.md +36 -0
- RUNBOOK_A100.md +51 -0
- configs/config_base.yml +35 -0
- configs/config_s1_lora.yml +37 -0
- configs/config_s2_distill.yml +37 -0
- configs/config_s3_ft.yml +37 -0
- configs/config_smoke.yml +37 -0
- official/bsr/__pycache__/degradations.cpython-312.pyc +0 -0
- official/bsr/degradations.py +764 -0
- official/bsr/transforms.py +179 -0
- official/bsr/utils/__init__.py +47 -0
- official/bsr/utils/color_util.py +208 -0
- official/bsr/utils/diffjpeg.py +515 -0
- official/bsr/utils/dist_util.py +82 -0
- official/bsr/utils/download_util.py +98 -0
- official/bsr/utils/file_client.py +167 -0
- official/bsr/utils/flow_util.py +170 -0
- official/bsr/utils/img_process_util.py +83 -0
- official/bsr/utils/img_util.py +172 -0
- official/bsr/utils/lmdb_util.py +199 -0
- official/bsr/utils/logger.py +213 -0
- official/bsr/utils/matlab_functions.py +178 -0
- official/bsr/utils/misc.py +141 -0
- official/bsr/utils/options.py +218 -0
- official/bsr/utils/plot_util.py +83 -0
- official/bsr/utils/registry.py +88 -0
- official/dataset.py +280 -0
- official/evaluate.py +55 -0
- official/forward.py +67 -0
- official/model.py +152 -0
- official/test.py +70 -0
- official/train.py +225 -0
- official/utils.py +24 -0
- requirements.txt +17 -0
- scripts/archive_run.py +31 -0
- scripts/download_weights.sh +32 -0
- scripts/env_check.py +25 -0
- scripts/probe_gdpo.py +13 -0
- scripts/run_smoke.sh +12 -0
- scripts/run_stage1.sh +14 -0
- scripts/run_stage2.sh +18 -0
- scripts/run_stage3.sh +14 -0
- scripts/setup_env.sh +21 -0
- src/check_submission.py +48 -0
- src/common.py +213 -0
- src/eval_val.py +104 -0
- src/export_jit.py +62 -0
- src/inference_4k.py +105 -0
- 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()
|