Colorizer round5: baseline checkpoint and source before GPU work
Browse files- experiments/round5-20260927/SHA256SUMS.json +36 -3
- experiments/round5-20260927/best/README.md +16 -0
- experiments/round5-20260927/best/config.json +956 -0
- experiments/round5-20260927/best/model.safetensors +3 -0
- experiments/round5-20260927/selection.json +5 -0
- experiments/round5-20260927/source/PROTOCOL.md +36 -0
- experiments/round5-20260927/source/data.py +116 -0
- experiments/round5-20260927/source/inference.py +64 -0
- experiments/round5-20260927/source/metrics.py +35 -0
- experiments/round5-20260927/source/model.py +122 -0
- experiments/round5-20260927/source/persistence.py +42 -0
- experiments/round5-20260927/source/previous_manifest.json +0 -0
- experiments/round5-20260927/source/spatial.py +34 -0
- experiments/round5-20260927/source/train.py +259 -0
- experiments/round5-20260927/source/vendor_ddcolor/LICENSE +201 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/__init__.py +16 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/__init__.py +41 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/__init__.py +0 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/convnext.py +206 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/position_encoding.py +52 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/transformer.py +368 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/transformer_utils.py +192 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/unet.py +208 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/util.py +63 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/__init__.py +37 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/diffjpeg.py +515 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/dist_util.py +82 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/file_client.py +167 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/img_process_util.py +83 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/img_util.py +227 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/logger.py +209 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/misc.py +141 -0
- experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/registry.py +82 -0
- experiments/round5-20260927/source/vendor_ddcolor/ddcolor/__init__.py +9 -0
- experiments/round5-20260927/source/vendor_ddcolor/ddcolor/model.py +278 -0
- experiments/round5-20260927/source/vendor_ddcolor/ddcolor/pipeline.py +127 -0
- experiments/round5-20260927/status.json +3 -0
experiments/round5-20260927/SHA256SUMS.json
CHANGED
|
@@ -1,5 +1,38 @@
|
|
| 1 |
{
|
| 2 |
-
"
|
| 3 |
-
"
|
| 4 |
-
"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5 |
}
|
|
|
|
| 1 |
{
|
| 2 |
+
"best/README.md": "fb835cbd31ee2df2d88d83f69354f0d82fc6213f24efbc6b3e92e1238a191493",
|
| 3 |
+
"best/config.json": "90299ee25ee3c3e645b362ddbd5a859b9e739c2a0a1c8a5e9ea79df171288bff",
|
| 4 |
+
"best/model.safetensors": "0e4c417375684a044860f8af3ac3a2fb44e1a5729254ca33abda758ea71aea6e",
|
| 5 |
+
"selection.json": "a714182d0a1985452e9ad7bb7a44729325382bde322edcc34b55f6680e2d9004",
|
| 6 |
+
"source/PROTOCOL.md": "cb34624b5f431b859c31f583f96f5ab531f0248dea2f33271947523c5ec4aa07",
|
| 7 |
+
"source/data.py": "eb43b886fe45435841439c36646985f914386a48faa712fb6806accac2ad48e0",
|
| 8 |
+
"source/inference.py": "e2b98bd98a15eb0c5beb4e9a2e64afb1cc4f381567f3ef213119f9bb21401a32",
|
| 9 |
+
"source/metrics.py": "af0ac996f6be0a351a55e372fab637a0d06740afe4a744cb80d0eae56ec97f66",
|
| 10 |
+
"source/model.py": "4cc57f82ffd6378bdf23a088ce7a1ed56c09e5673de3ee5232b8ee1cacc9be0a",
|
| 11 |
+
"source/persistence.py": "9ccfc8d5908cef98078422049e8caaaf558ac8293792cd77e8b483f92ad908be",
|
| 12 |
+
"source/previous_manifest.json": "24bc081bf48f1109d9dfa217ae0d8593c0c06af6785a5cbfc5dc5d9e25461f5c",
|
| 13 |
+
"source/spatial.py": "d13abe399ef049e21a6459a7003461afaae0c00e0c560262c6cf375de4c9884a",
|
| 14 |
+
"source/train.py": "03e11298fdc8668042c61caee5cf5789ce4675cacc881cb214d5d9bbb64ad056",
|
| 15 |
+
"source/vendor_ddcolor/LICENSE": "43070e2d4e532684de521b885f385d0841030efa2b1a20bafb76133a5e1379c1",
|
| 16 |
+
"source/vendor_ddcolor/basicsr/__init__.py": "376dd0503a4853c5f9831a15ce29a821eb668689efaf1c1a6a87d09f7696a9db",
|
| 17 |
+
"source/vendor_ddcolor/basicsr/archs/__init__.py": "6b0ddbe90c089837eeca316425bb708e30d75c89efa2015f5aba41dc820cf1c9",
|
| 18 |
+
"source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/__init__.py": "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855",
|
| 19 |
+
"source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/convnext.py": "f03e339488570aa79e0690719546317a375ca85bd8a9e1e5f2c797cda033db45",
|
| 20 |
+
"source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/position_encoding.py": "beb5b3b52f2cc4f2dfc9f312cdc5d712468ff1d2985be7471a94a0c37a8c01d6",
|
| 21 |
+
"source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/transformer.py": "95548deb6125bb42801e9be5569007256cc50f5a19f7b6591348e4aea9ee1a5d",
|
| 22 |
+
"source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/transformer_utils.py": "61d99be451ef0717eaa1bf84ecfd29f9c4224ac1e5a35d68a963c4b97f4dbd0b",
|
| 23 |
+
"source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/unet.py": "2b72fd6ef60d73f033cd3b2a3b4737d5371d4190b6410883e4d987d8417df1b2",
|
| 24 |
+
"source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/util.py": "a771c34a643735107ee0788ceab15fd5baa39e4e6417149194c7255ad39d2169",
|
| 25 |
+
"source/vendor_ddcolor/basicsr/utils/__init__.py": "d06ae0685849d5090788f2d5ebaf93b2c8a67e1ef0284c98e6efeb70dbe3ec77",
|
| 26 |
+
"source/vendor_ddcolor/basicsr/utils/diffjpeg.py": "ac20cc17c29684b27cf08a0d0c12ad6aee3811a9333ccfc3a45dd96d29b8bbbf",
|
| 27 |
+
"source/vendor_ddcolor/basicsr/utils/dist_util.py": "e6a6d7fced5146ff2e1f16cb0ce1c922171fe7e8b28c011178695ebed4f477b8",
|
| 28 |
+
"source/vendor_ddcolor/basicsr/utils/file_client.py": "0a48e9073e2813651303003dc212579098764e62a2e58fcf889832ddf9c00a25",
|
| 29 |
+
"source/vendor_ddcolor/basicsr/utils/img_process_util.py": "e92f3ee102bca3b7dd8f2c2b196a40c15c7cd181c782455ec4e71f4331ccd625",
|
| 30 |
+
"source/vendor_ddcolor/basicsr/utils/img_util.py": "ca32b89d12de4494640d5ed0160f827aa1cfc55a3c0b5aea4b6e72118311199e",
|
| 31 |
+
"source/vendor_ddcolor/basicsr/utils/logger.py": "5361d5d9bcb92ed9b8a81278c6acc4d2dcb57bef29da92f16c0693023edb4abe",
|
| 32 |
+
"source/vendor_ddcolor/basicsr/utils/misc.py": "4e94ff742db4938d636cfcc096dec4707a036ef1aa5312925944b21d0e306dff",
|
| 33 |
+
"source/vendor_ddcolor/basicsr/utils/registry.py": "21c23bd3bf727eb2decd99035120fc47316db92486d73116cca18e9fa8b6c854",
|
| 34 |
+
"source/vendor_ddcolor/ddcolor/__init__.py": "2606a3377189bda8beb1a2075c7bd789961163bca3e1e53b01bbe325dc561c30",
|
| 35 |
+
"source/vendor_ddcolor/ddcolor/model.py": "e8e30115aa65a9558db33641c999001d1735c444402890898b47b2cc52b37c25",
|
| 36 |
+
"source/vendor_ddcolor/ddcolor/pipeline.py": "61fb3a7642309d1dc5069fb8089dcdff13ab5680c78269013b579bfa654ce005",
|
| 37 |
+
"status.json": "756d10f0138e6fd95b56a37010d0ea857f12f9d53ca695523ea11f8b75e821b6"
|
| 38 |
}
|
experiments/round5-20260927/best/README.md
ADDED
|
@@ -0,0 +1,16 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: apache-2.0
|
| 3 |
+
pipeline_tag: image-to-image
|
| 4 |
+
tags:
|
| 5 |
+
- classification
|
| 6 |
+
- colorization
|
| 7 |
+
- image-to-image
|
| 8 |
+
- model_hub_mixin
|
| 9 |
+
- pytorch_model_hub_mixin
|
| 10 |
+
- unet
|
| 11 |
+
---
|
| 12 |
+
|
| 13 |
+
This model has been pushed to the Hub using the [PytorchModelHubMixin](https://huggingface.co/docs/huggingface_hub/package_reference/mixins#huggingface_hub.PyTorchModelHubMixin) integration:
|
| 14 |
+
- Code: [More Information Needed]
|
| 15 |
+
- Paper: [More Information Needed]
|
| 16 |
+
- Docs: [More Information Needed]
|
experiments/round5-20260927/best/config.json
ADDED
|
@@ -0,0 +1,956 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"base": 44,
|
| 3 |
+
"bin_centers": [
|
| 4 |
+
[
|
| 5 |
+
-75.0,
|
| 6 |
+
25.0
|
| 7 |
+
],
|
| 8 |
+
[
|
| 9 |
+
-75.0,
|
| 10 |
+
35.0
|
| 11 |
+
],
|
| 12 |
+
[
|
| 13 |
+
-75.0,
|
| 14 |
+
45.0
|
| 15 |
+
],
|
| 16 |
+
[
|
| 17 |
+
-75.0,
|
| 18 |
+
55.0
|
| 19 |
+
],
|
| 20 |
+
[
|
| 21 |
+
-75.0,
|
| 22 |
+
65.0
|
| 23 |
+
],
|
| 24 |
+
[
|
| 25 |
+
-75.0,
|
| 26 |
+
75.0
|
| 27 |
+
],
|
| 28 |
+
[
|
| 29 |
+
-65.0,
|
| 30 |
+
25.0
|
| 31 |
+
],
|
| 32 |
+
[
|
| 33 |
+
-65.0,
|
| 34 |
+
35.0
|
| 35 |
+
],
|
| 36 |
+
[
|
| 37 |
+
-65.0,
|
| 38 |
+
45.0
|
| 39 |
+
],
|
| 40 |
+
[
|
| 41 |
+
-65.0,
|
| 42 |
+
55.0
|
| 43 |
+
],
|
| 44 |
+
[
|
| 45 |
+
-65.0,
|
| 46 |
+
65.0
|
| 47 |
+
],
|
| 48 |
+
[
|
| 49 |
+
-65.0,
|
| 50 |
+
75.0
|
| 51 |
+
],
|
| 52 |
+
[
|
| 53 |
+
-55.0,
|
| 54 |
+
5.0
|
| 55 |
+
],
|
| 56 |
+
[
|
| 57 |
+
-55.0,
|
| 58 |
+
15.0
|
| 59 |
+
],
|
| 60 |
+
[
|
| 61 |
+
-55.0,
|
| 62 |
+
25.0
|
| 63 |
+
],
|
| 64 |
+
[
|
| 65 |
+
-55.0,
|
| 66 |
+
35.0
|
| 67 |
+
],
|
| 68 |
+
[
|
| 69 |
+
-55.0,
|
| 70 |
+
45.0
|
| 71 |
+
],
|
| 72 |
+
[
|
| 73 |
+
-55.0,
|
| 74 |
+
55.0
|
| 75 |
+
],
|
| 76 |
+
[
|
| 77 |
+
-55.0,
|
| 78 |
+
65.0
|
| 79 |
+
],
|
| 80 |
+
[
|
| 81 |
+
-55.0,
|
| 82 |
+
75.0
|
| 83 |
+
],
|
| 84 |
+
[
|
| 85 |
+
-45.0,
|
| 86 |
+
-15.0
|
| 87 |
+
],
|
| 88 |
+
[
|
| 89 |
+
-45.0,
|
| 90 |
+
-5.0
|
| 91 |
+
],
|
| 92 |
+
[
|
| 93 |
+
-45.0,
|
| 94 |
+
5.0
|
| 95 |
+
],
|
| 96 |
+
[
|
| 97 |
+
-45.0,
|
| 98 |
+
15.0
|
| 99 |
+
],
|
| 100 |
+
[
|
| 101 |
+
-45.0,
|
| 102 |
+
25.0
|
| 103 |
+
],
|
| 104 |
+
[
|
| 105 |
+
-45.0,
|
| 106 |
+
35.0
|
| 107 |
+
],
|
| 108 |
+
[
|
| 109 |
+
-45.0,
|
| 110 |
+
45.0
|
| 111 |
+
],
|
| 112 |
+
[
|
| 113 |
+
-45.0,
|
| 114 |
+
55.0
|
| 115 |
+
],
|
| 116 |
+
[
|
| 117 |
+
-45.0,
|
| 118 |
+
65.0
|
| 119 |
+
],
|
| 120 |
+
[
|
| 121 |
+
-45.0,
|
| 122 |
+
75.0
|
| 123 |
+
],
|
| 124 |
+
[
|
| 125 |
+
-45.0,
|
| 126 |
+
85.0
|
| 127 |
+
],
|
| 128 |
+
[
|
| 129 |
+
-35.0,
|
| 130 |
+
-35.0
|
| 131 |
+
],
|
| 132 |
+
[
|
| 133 |
+
-35.0,
|
| 134 |
+
-25.0
|
| 135 |
+
],
|
| 136 |
+
[
|
| 137 |
+
-35.0,
|
| 138 |
+
-15.0
|
| 139 |
+
],
|
| 140 |
+
[
|
| 141 |
+
-35.0,
|
| 142 |
+
-5.0
|
| 143 |
+
],
|
| 144 |
+
[
|
| 145 |
+
-35.0,
|
| 146 |
+
5.0
|
| 147 |
+
],
|
| 148 |
+
[
|
| 149 |
+
-35.0,
|
| 150 |
+
15.0
|
| 151 |
+
],
|
| 152 |
+
[
|
| 153 |
+
-35.0,
|
| 154 |
+
25.0
|
| 155 |
+
],
|
| 156 |
+
[
|
| 157 |
+
-35.0,
|
| 158 |
+
35.0
|
| 159 |
+
],
|
| 160 |
+
[
|
| 161 |
+
-35.0,
|
| 162 |
+
45.0
|
| 163 |
+
],
|
| 164 |
+
[
|
| 165 |
+
-35.0,
|
| 166 |
+
55.0
|
| 167 |
+
],
|
| 168 |
+
[
|
| 169 |
+
-35.0,
|
| 170 |
+
65.0
|
| 171 |
+
],
|
| 172 |
+
[
|
| 173 |
+
-35.0,
|
| 174 |
+
75.0
|
| 175 |
+
],
|
| 176 |
+
[
|
| 177 |
+
-35.0,
|
| 178 |
+
85.0
|
| 179 |
+
],
|
| 180 |
+
[
|
| 181 |
+
-35.0,
|
| 182 |
+
95.0
|
| 183 |
+
],
|
| 184 |
+
[
|
| 185 |
+
-25.0,
|
| 186 |
+
-35.0
|
| 187 |
+
],
|
| 188 |
+
[
|
| 189 |
+
-25.0,
|
| 190 |
+
-25.0
|
| 191 |
+
],
|
| 192 |
+
[
|
| 193 |
+
-25.0,
|
| 194 |
+
-15.0
|
| 195 |
+
],
|
| 196 |
+
[
|
| 197 |
+
-25.0,
|
| 198 |
+
-5.0
|
| 199 |
+
],
|
| 200 |
+
[
|
| 201 |
+
-25.0,
|
| 202 |
+
5.0
|
| 203 |
+
],
|
| 204 |
+
[
|
| 205 |
+
-25.0,
|
| 206 |
+
15.0
|
| 207 |
+
],
|
| 208 |
+
[
|
| 209 |
+
-25.0,
|
| 210 |
+
25.0
|
| 211 |
+
],
|
| 212 |
+
[
|
| 213 |
+
-25.0,
|
| 214 |
+
35.0
|
| 215 |
+
],
|
| 216 |
+
[
|
| 217 |
+
-25.0,
|
| 218 |
+
45.0
|
| 219 |
+
],
|
| 220 |
+
[
|
| 221 |
+
-25.0,
|
| 222 |
+
55.0
|
| 223 |
+
],
|
| 224 |
+
[
|
| 225 |
+
-25.0,
|
| 226 |
+
65.0
|
| 227 |
+
],
|
| 228 |
+
[
|
| 229 |
+
-25.0,
|
| 230 |
+
75.0
|
| 231 |
+
],
|
| 232 |
+
[
|
| 233 |
+
-25.0,
|
| 234 |
+
85.0
|
| 235 |
+
],
|
| 236 |
+
[
|
| 237 |
+
-25.0,
|
| 238 |
+
95.0
|
| 239 |
+
],
|
| 240 |
+
[
|
| 241 |
+
-15.0,
|
| 242 |
+
-45.0
|
| 243 |
+
],
|
| 244 |
+
[
|
| 245 |
+
-15.0,
|
| 246 |
+
-35.0
|
| 247 |
+
],
|
| 248 |
+
[
|
| 249 |
+
-15.0,
|
| 250 |
+
-25.0
|
| 251 |
+
],
|
| 252 |
+
[
|
| 253 |
+
-15.0,
|
| 254 |
+
-15.0
|
| 255 |
+
],
|
| 256 |
+
[
|
| 257 |
+
-15.0,
|
| 258 |
+
-5.0
|
| 259 |
+
],
|
| 260 |
+
[
|
| 261 |
+
-15.0,
|
| 262 |
+
5.0
|
| 263 |
+
],
|
| 264 |
+
[
|
| 265 |
+
-15.0,
|
| 266 |
+
15.0
|
| 267 |
+
],
|
| 268 |
+
[
|
| 269 |
+
-15.0,
|
| 270 |
+
25.0
|
| 271 |
+
],
|
| 272 |
+
[
|
| 273 |
+
-15.0,
|
| 274 |
+
35.0
|
| 275 |
+
],
|
| 276 |
+
[
|
| 277 |
+
-15.0,
|
| 278 |
+
45.0
|
| 279 |
+
],
|
| 280 |
+
[
|
| 281 |
+
-15.0,
|
| 282 |
+
55.0
|
| 283 |
+
],
|
| 284 |
+
[
|
| 285 |
+
-15.0,
|
| 286 |
+
65.0
|
| 287 |
+
],
|
| 288 |
+
[
|
| 289 |
+
-15.0,
|
| 290 |
+
75.0
|
| 291 |
+
],
|
| 292 |
+
[
|
| 293 |
+
-15.0,
|
| 294 |
+
85.0
|
| 295 |
+
],
|
| 296 |
+
[
|
| 297 |
+
-15.0,
|
| 298 |
+
95.0
|
| 299 |
+
],
|
| 300 |
+
[
|
| 301 |
+
-5.0,
|
| 302 |
+
-55.0
|
| 303 |
+
],
|
| 304 |
+
[
|
| 305 |
+
-5.0,
|
| 306 |
+
-45.0
|
| 307 |
+
],
|
| 308 |
+
[
|
| 309 |
+
-5.0,
|
| 310 |
+
-35.0
|
| 311 |
+
],
|
| 312 |
+
[
|
| 313 |
+
-5.0,
|
| 314 |
+
-25.0
|
| 315 |
+
],
|
| 316 |
+
[
|
| 317 |
+
-5.0,
|
| 318 |
+
-15.0
|
| 319 |
+
],
|
| 320 |
+
[
|
| 321 |
+
-5.0,
|
| 322 |
+
-5.0
|
| 323 |
+
],
|
| 324 |
+
[
|
| 325 |
+
-5.0,
|
| 326 |
+
5.0
|
| 327 |
+
],
|
| 328 |
+
[
|
| 329 |
+
-5.0,
|
| 330 |
+
15.0
|
| 331 |
+
],
|
| 332 |
+
[
|
| 333 |
+
-5.0,
|
| 334 |
+
25.0
|
| 335 |
+
],
|
| 336 |
+
[
|
| 337 |
+
-5.0,
|
| 338 |
+
35.0
|
| 339 |
+
],
|
| 340 |
+
[
|
| 341 |
+
-5.0,
|
| 342 |
+
45.0
|
| 343 |
+
],
|
| 344 |
+
[
|
| 345 |
+
-5.0,
|
| 346 |
+
55.0
|
| 347 |
+
],
|
| 348 |
+
[
|
| 349 |
+
-5.0,
|
| 350 |
+
65.0
|
| 351 |
+
],
|
| 352 |
+
[
|
| 353 |
+
-5.0,
|
| 354 |
+
75.0
|
| 355 |
+
],
|
| 356 |
+
[
|
| 357 |
+
-5.0,
|
| 358 |
+
85.0
|
| 359 |
+
],
|
| 360 |
+
[
|
| 361 |
+
5.0,
|
| 362 |
+
-65.0
|
| 363 |
+
],
|
| 364 |
+
[
|
| 365 |
+
5.0,
|
| 366 |
+
-55.0
|
| 367 |
+
],
|
| 368 |
+
[
|
| 369 |
+
5.0,
|
| 370 |
+
-45.0
|
| 371 |
+
],
|
| 372 |
+
[
|
| 373 |
+
5.0,
|
| 374 |
+
-35.0
|
| 375 |
+
],
|
| 376 |
+
[
|
| 377 |
+
5.0,
|
| 378 |
+
-25.0
|
| 379 |
+
],
|
| 380 |
+
[
|
| 381 |
+
5.0,
|
| 382 |
+
-15.0
|
| 383 |
+
],
|
| 384 |
+
[
|
| 385 |
+
5.0,
|
| 386 |
+
-5.0
|
| 387 |
+
],
|
| 388 |
+
[
|
| 389 |
+
5.0,
|
| 390 |
+
5.0
|
| 391 |
+
],
|
| 392 |
+
[
|
| 393 |
+
5.0,
|
| 394 |
+
15.0
|
| 395 |
+
],
|
| 396 |
+
[
|
| 397 |
+
5.0,
|
| 398 |
+
25.0
|
| 399 |
+
],
|
| 400 |
+
[
|
| 401 |
+
5.0,
|
| 402 |
+
35.0
|
| 403 |
+
],
|
| 404 |
+
[
|
| 405 |
+
5.0,
|
| 406 |
+
45.0
|
| 407 |
+
],
|
| 408 |
+
[
|
| 409 |
+
5.0,
|
| 410 |
+
55.0
|
| 411 |
+
],
|
| 412 |
+
[
|
| 413 |
+
5.0,
|
| 414 |
+
65.0
|
| 415 |
+
],
|
| 416 |
+
[
|
| 417 |
+
5.0,
|
| 418 |
+
75.0
|
| 419 |
+
],
|
| 420 |
+
[
|
| 421 |
+
5.0,
|
| 422 |
+
85.0
|
| 423 |
+
],
|
| 424 |
+
[
|
| 425 |
+
15.0,
|
| 426 |
+
-75.0
|
| 427 |
+
],
|
| 428 |
+
[
|
| 429 |
+
15.0,
|
| 430 |
+
-65.0
|
| 431 |
+
],
|
| 432 |
+
[
|
| 433 |
+
15.0,
|
| 434 |
+
-55.0
|
| 435 |
+
],
|
| 436 |
+
[
|
| 437 |
+
15.0,
|
| 438 |
+
-45.0
|
| 439 |
+
],
|
| 440 |
+
[
|
| 441 |
+
15.0,
|
| 442 |
+
-35.0
|
| 443 |
+
],
|
| 444 |
+
[
|
| 445 |
+
15.0,
|
| 446 |
+
-25.0
|
| 447 |
+
],
|
| 448 |
+
[
|
| 449 |
+
15.0,
|
| 450 |
+
-15.0
|
| 451 |
+
],
|
| 452 |
+
[
|
| 453 |
+
15.0,
|
| 454 |
+
-5.0
|
| 455 |
+
],
|
| 456 |
+
[
|
| 457 |
+
15.0,
|
| 458 |
+
5.0
|
| 459 |
+
],
|
| 460 |
+
[
|
| 461 |
+
15.0,
|
| 462 |
+
15.0
|
| 463 |
+
],
|
| 464 |
+
[
|
| 465 |
+
15.0,
|
| 466 |
+
25.0
|
| 467 |
+
],
|
| 468 |
+
[
|
| 469 |
+
15.0,
|
| 470 |
+
35.0
|
| 471 |
+
],
|
| 472 |
+
[
|
| 473 |
+
15.0,
|
| 474 |
+
45.0
|
| 475 |
+
],
|
| 476 |
+
[
|
| 477 |
+
15.0,
|
| 478 |
+
55.0
|
| 479 |
+
],
|
| 480 |
+
[
|
| 481 |
+
15.0,
|
| 482 |
+
65.0
|
| 483 |
+
],
|
| 484 |
+
[
|
| 485 |
+
15.0,
|
| 486 |
+
75.0
|
| 487 |
+
],
|
| 488 |
+
[
|
| 489 |
+
15.0,
|
| 490 |
+
85.0
|
| 491 |
+
],
|
| 492 |
+
[
|
| 493 |
+
25.0,
|
| 494 |
+
-75.0
|
| 495 |
+
],
|
| 496 |
+
[
|
| 497 |
+
25.0,
|
| 498 |
+
-65.0
|
| 499 |
+
],
|
| 500 |
+
[
|
| 501 |
+
25.0,
|
| 502 |
+
-55.0
|
| 503 |
+
],
|
| 504 |
+
[
|
| 505 |
+
25.0,
|
| 506 |
+
-45.0
|
| 507 |
+
],
|
| 508 |
+
[
|
| 509 |
+
25.0,
|
| 510 |
+
-35.0
|
| 511 |
+
],
|
| 512 |
+
[
|
| 513 |
+
25.0,
|
| 514 |
+
-25.0
|
| 515 |
+
],
|
| 516 |
+
[
|
| 517 |
+
25.0,
|
| 518 |
+
-15.0
|
| 519 |
+
],
|
| 520 |
+
[
|
| 521 |
+
25.0,
|
| 522 |
+
-5.0
|
| 523 |
+
],
|
| 524 |
+
[
|
| 525 |
+
25.0,
|
| 526 |
+
5.0
|
| 527 |
+
],
|
| 528 |
+
[
|
| 529 |
+
25.0,
|
| 530 |
+
15.0
|
| 531 |
+
],
|
| 532 |
+
[
|
| 533 |
+
25.0,
|
| 534 |
+
25.0
|
| 535 |
+
],
|
| 536 |
+
[
|
| 537 |
+
25.0,
|
| 538 |
+
35.0
|
| 539 |
+
],
|
| 540 |
+
[
|
| 541 |
+
25.0,
|
| 542 |
+
45.0
|
| 543 |
+
],
|
| 544 |
+
[
|
| 545 |
+
25.0,
|
| 546 |
+
55.0
|
| 547 |
+
],
|
| 548 |
+
[
|
| 549 |
+
25.0,
|
| 550 |
+
65.0
|
| 551 |
+
],
|
| 552 |
+
[
|
| 553 |
+
25.0,
|
| 554 |
+
75.0
|
| 555 |
+
],
|
| 556 |
+
[
|
| 557 |
+
25.0,
|
| 558 |
+
85.0
|
| 559 |
+
],
|
| 560 |
+
[
|
| 561 |
+
35.0,
|
| 562 |
+
-85.0
|
| 563 |
+
],
|
| 564 |
+
[
|
| 565 |
+
35.0,
|
| 566 |
+
-75.0
|
| 567 |
+
],
|
| 568 |
+
[
|
| 569 |
+
35.0,
|
| 570 |
+
-65.0
|
| 571 |
+
],
|
| 572 |
+
[
|
| 573 |
+
35.0,
|
| 574 |
+
-55.0
|
| 575 |
+
],
|
| 576 |
+
[
|
| 577 |
+
35.0,
|
| 578 |
+
-45.0
|
| 579 |
+
],
|
| 580 |
+
[
|
| 581 |
+
35.0,
|
| 582 |
+
-35.0
|
| 583 |
+
],
|
| 584 |
+
[
|
| 585 |
+
35.0,
|
| 586 |
+
-25.0
|
| 587 |
+
],
|
| 588 |
+
[
|
| 589 |
+
35.0,
|
| 590 |
+
-15.0
|
| 591 |
+
],
|
| 592 |
+
[
|
| 593 |
+
35.0,
|
| 594 |
+
-5.0
|
| 595 |
+
],
|
| 596 |
+
[
|
| 597 |
+
35.0,
|
| 598 |
+
5.0
|
| 599 |
+
],
|
| 600 |
+
[
|
| 601 |
+
35.0,
|
| 602 |
+
15.0
|
| 603 |
+
],
|
| 604 |
+
[
|
| 605 |
+
35.0,
|
| 606 |
+
25.0
|
| 607 |
+
],
|
| 608 |
+
[
|
| 609 |
+
35.0,
|
| 610 |
+
35.0
|
| 611 |
+
],
|
| 612 |
+
[
|
| 613 |
+
35.0,
|
| 614 |
+
45.0
|
| 615 |
+
],
|
| 616 |
+
[
|
| 617 |
+
35.0,
|
| 618 |
+
55.0
|
| 619 |
+
],
|
| 620 |
+
[
|
| 621 |
+
35.0,
|
| 622 |
+
65.0
|
| 623 |
+
],
|
| 624 |
+
[
|
| 625 |
+
35.0,
|
| 626 |
+
75.0
|
| 627 |
+
],
|
| 628 |
+
[
|
| 629 |
+
45.0,
|
| 630 |
+
-95.0
|
| 631 |
+
],
|
| 632 |
+
[
|
| 633 |
+
45.0,
|
| 634 |
+
-85.0
|
| 635 |
+
],
|
| 636 |
+
[
|
| 637 |
+
45.0,
|
| 638 |
+
-75.0
|
| 639 |
+
],
|
| 640 |
+
[
|
| 641 |
+
45.0,
|
| 642 |
+
-65.0
|
| 643 |
+
],
|
| 644 |
+
[
|
| 645 |
+
45.0,
|
| 646 |
+
-55.0
|
| 647 |
+
],
|
| 648 |
+
[
|
| 649 |
+
45.0,
|
| 650 |
+
-45.0
|
| 651 |
+
],
|
| 652 |
+
[
|
| 653 |
+
45.0,
|
| 654 |
+
-35.0
|
| 655 |
+
],
|
| 656 |
+
[
|
| 657 |
+
45.0,
|
| 658 |
+
-25.0
|
| 659 |
+
],
|
| 660 |
+
[
|
| 661 |
+
45.0,
|
| 662 |
+
-15.0
|
| 663 |
+
],
|
| 664 |
+
[
|
| 665 |
+
45.0,
|
| 666 |
+
-5.0
|
| 667 |
+
],
|
| 668 |
+
[
|
| 669 |
+
45.0,
|
| 670 |
+
5.0
|
| 671 |
+
],
|
| 672 |
+
[
|
| 673 |
+
45.0,
|
| 674 |
+
15.0
|
| 675 |
+
],
|
| 676 |
+
[
|
| 677 |
+
45.0,
|
| 678 |
+
25.0
|
| 679 |
+
],
|
| 680 |
+
[
|
| 681 |
+
45.0,
|
| 682 |
+
35.0
|
| 683 |
+
],
|
| 684 |
+
[
|
| 685 |
+
45.0,
|
| 686 |
+
45.0
|
| 687 |
+
],
|
| 688 |
+
[
|
| 689 |
+
45.0,
|
| 690 |
+
55.0
|
| 691 |
+
],
|
| 692 |
+
[
|
| 693 |
+
45.0,
|
| 694 |
+
65.0
|
| 695 |
+
],
|
| 696 |
+
[
|
| 697 |
+
45.0,
|
| 698 |
+
75.0
|
| 699 |
+
],
|
| 700 |
+
[
|
| 701 |
+
55.0,
|
| 702 |
+
-95.0
|
| 703 |
+
],
|
| 704 |
+
[
|
| 705 |
+
55.0,
|
| 706 |
+
-85.0
|
| 707 |
+
],
|
| 708 |
+
[
|
| 709 |
+
55.0,
|
| 710 |
+
-75.0
|
| 711 |
+
],
|
| 712 |
+
[
|
| 713 |
+
55.0,
|
| 714 |
+
-65.0
|
| 715 |
+
],
|
| 716 |
+
[
|
| 717 |
+
55.0,
|
| 718 |
+
-55.0
|
| 719 |
+
],
|
| 720 |
+
[
|
| 721 |
+
55.0,
|
| 722 |
+
-45.0
|
| 723 |
+
],
|
| 724 |
+
[
|
| 725 |
+
55.0,
|
| 726 |
+
-35.0
|
| 727 |
+
],
|
| 728 |
+
[
|
| 729 |
+
55.0,
|
| 730 |
+
-25.0
|
| 731 |
+
],
|
| 732 |
+
[
|
| 733 |
+
55.0,
|
| 734 |
+
-15.0
|
| 735 |
+
],
|
| 736 |
+
[
|
| 737 |
+
55.0,
|
| 738 |
+
-5.0
|
| 739 |
+
],
|
| 740 |
+
[
|
| 741 |
+
55.0,
|
| 742 |
+
5.0
|
| 743 |
+
],
|
| 744 |
+
[
|
| 745 |
+
55.0,
|
| 746 |
+
15.0
|
| 747 |
+
],
|
| 748 |
+
[
|
| 749 |
+
55.0,
|
| 750 |
+
25.0
|
| 751 |
+
],
|
| 752 |
+
[
|
| 753 |
+
55.0,
|
| 754 |
+
35.0
|
| 755 |
+
],
|
| 756 |
+
[
|
| 757 |
+
55.0,
|
| 758 |
+
45.0
|
| 759 |
+
],
|
| 760 |
+
[
|
| 761 |
+
55.0,
|
| 762 |
+
55.0
|
| 763 |
+
],
|
| 764 |
+
[
|
| 765 |
+
55.0,
|
| 766 |
+
65.0
|
| 767 |
+
],
|
| 768 |
+
[
|
| 769 |
+
55.0,
|
| 770 |
+
75.0
|
| 771 |
+
],
|
| 772 |
+
[
|
| 773 |
+
65.0,
|
| 774 |
+
-105.0
|
| 775 |
+
],
|
| 776 |
+
[
|
| 777 |
+
65.0,
|
| 778 |
+
-95.0
|
| 779 |
+
],
|
| 780 |
+
[
|
| 781 |
+
65.0,
|
| 782 |
+
-85.0
|
| 783 |
+
],
|
| 784 |
+
[
|
| 785 |
+
65.0,
|
| 786 |
+
-75.0
|
| 787 |
+
],
|
| 788 |
+
[
|
| 789 |
+
65.0,
|
| 790 |
+
-65.0
|
| 791 |
+
],
|
| 792 |
+
[
|
| 793 |
+
65.0,
|
| 794 |
+
-55.0
|
| 795 |
+
],
|
| 796 |
+
[
|
| 797 |
+
65.0,
|
| 798 |
+
-45.0
|
| 799 |
+
],
|
| 800 |
+
[
|
| 801 |
+
65.0,
|
| 802 |
+
-35.0
|
| 803 |
+
],
|
| 804 |
+
[
|
| 805 |
+
65.0,
|
| 806 |
+
-25.0
|
| 807 |
+
],
|
| 808 |
+
[
|
| 809 |
+
65.0,
|
| 810 |
+
-15.0
|
| 811 |
+
],
|
| 812 |
+
[
|
| 813 |
+
65.0,
|
| 814 |
+
-5.0
|
| 815 |
+
],
|
| 816 |
+
[
|
| 817 |
+
65.0,
|
| 818 |
+
5.0
|
| 819 |
+
],
|
| 820 |
+
[
|
| 821 |
+
65.0,
|
| 822 |
+
15.0
|
| 823 |
+
],
|
| 824 |
+
[
|
| 825 |
+
65.0,
|
| 826 |
+
25.0
|
| 827 |
+
],
|
| 828 |
+
[
|
| 829 |
+
65.0,
|
| 830 |
+
35.0
|
| 831 |
+
],
|
| 832 |
+
[
|
| 833 |
+
65.0,
|
| 834 |
+
45.0
|
| 835 |
+
],
|
| 836 |
+
[
|
| 837 |
+
65.0,
|
| 838 |
+
55.0
|
| 839 |
+
],
|
| 840 |
+
[
|
| 841 |
+
65.0,
|
| 842 |
+
65.0
|
| 843 |
+
],
|
| 844 |
+
[
|
| 845 |
+
75.0,
|
| 846 |
+
-105.0
|
| 847 |
+
],
|
| 848 |
+
[
|
| 849 |
+
75.0,
|
| 850 |
+
-95.0
|
| 851 |
+
],
|
| 852 |
+
[
|
| 853 |
+
75.0,
|
| 854 |
+
-45.0
|
| 855 |
+
],
|
| 856 |
+
[
|
| 857 |
+
75.0,
|
| 858 |
+
-35.0
|
| 859 |
+
],
|
| 860 |
+
[
|
| 861 |
+
75.0,
|
| 862 |
+
-25.0
|
| 863 |
+
],
|
| 864 |
+
[
|
| 865 |
+
75.0,
|
| 866 |
+
-15.0
|
| 867 |
+
],
|
| 868 |
+
[
|
| 869 |
+
75.0,
|
| 870 |
+
-5.0
|
| 871 |
+
],
|
| 872 |
+
[
|
| 873 |
+
75.0,
|
| 874 |
+
5.0
|
| 875 |
+
],
|
| 876 |
+
[
|
| 877 |
+
75.0,
|
| 878 |
+
15.0
|
| 879 |
+
],
|
| 880 |
+
[
|
| 881 |
+
75.0,
|
| 882 |
+
25.0
|
| 883 |
+
],
|
| 884 |
+
[
|
| 885 |
+
75.0,
|
| 886 |
+
35.0
|
| 887 |
+
],
|
| 888 |
+
[
|
| 889 |
+
75.0,
|
| 890 |
+
45.0
|
| 891 |
+
],
|
| 892 |
+
[
|
| 893 |
+
75.0,
|
| 894 |
+
55.0
|
| 895 |
+
],
|
| 896 |
+
[
|
| 897 |
+
75.0,
|
| 898 |
+
65.0
|
| 899 |
+
],
|
| 900 |
+
[
|
| 901 |
+
85.0,
|
| 902 |
+
-45.0
|
| 903 |
+
],
|
| 904 |
+
[
|
| 905 |
+
85.0,
|
| 906 |
+
-35.0
|
| 907 |
+
],
|
| 908 |
+
[
|
| 909 |
+
85.0,
|
| 910 |
+
-25.0
|
| 911 |
+
],
|
| 912 |
+
[
|
| 913 |
+
85.0,
|
| 914 |
+
-15.0
|
| 915 |
+
],
|
| 916 |
+
[
|
| 917 |
+
85.0,
|
| 918 |
+
-5.0
|
| 919 |
+
],
|
| 920 |
+
[
|
| 921 |
+
85.0,
|
| 922 |
+
5.0
|
| 923 |
+
],
|
| 924 |
+
[
|
| 925 |
+
85.0,
|
| 926 |
+
35.0
|
| 927 |
+
],
|
| 928 |
+
[
|
| 929 |
+
85.0,
|
| 930 |
+
45.0
|
| 931 |
+
],
|
| 932 |
+
[
|
| 933 |
+
85.0,
|
| 934 |
+
65.0
|
| 935 |
+
],
|
| 936 |
+
[
|
| 937 |
+
95.0,
|
| 938 |
+
-55.0
|
| 939 |
+
],
|
| 940 |
+
[
|
| 941 |
+
95.0,
|
| 942 |
+
-45.0
|
| 943 |
+
],
|
| 944 |
+
[
|
| 945 |
+
95.0,
|
| 946 |
+
-35.0
|
| 947 |
+
]
|
| 948 |
+
],
|
| 949 |
+
"context_dilations": [
|
| 950 |
+
2,
|
| 951 |
+
4,
|
| 952 |
+
8
|
| 953 |
+
],
|
| 954 |
+
"context_mid_ch": 96,
|
| 955 |
+
"in_ch": 1
|
| 956 |
+
}
|
experiments/round5-20260927/best/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:0e4c417375684a044860f8af3ac3a2fb44e1a5729254ca33abda758ea71aea6e
|
| 3 |
+
size 15909320
|
experiments/round5-20260927/selection.json
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"arm": "v2_baseline",
|
| 3 |
+
"step": 0,
|
| 4 |
+
"accepted": false
|
| 5 |
+
}
|
experiments/round5-20260927/source/PROTOCOL.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Round 4 — semantic guidance and broad color patches
|
| 2 |
+
|
| 3 |
+
Status: experiment protocol, not a result or release claim.
|
| 4 |
+
|
| 5 |
+
## Research and rationale
|
| 6 |
+
|
| 7 |
+
DDColor uses multi-scale semantic features and color queries to reduce color bleeding: https://arxiv.org/abs/2212.11613 . The authors' model notes warn that their colorfulness loss can generate unwanted color blocks and describe an artistic checkpoint trained without that loss: https://github.com/piddnad/DDColor/blob/master/MODEL_ZOO.md . This does not prove that our different round-3 chroma-retention loss caused all artifacts, but it motivates removing extra saturation pressure and measuring the effect.
|
| 8 |
+
|
| 9 |
+
The previous fine-edge proxy missed broad colored clouds. Round 3 reduced missed color but increased neutral-region spill, and its actual images still contained patches. This run therefore measures color differences across regions separated by 4–64 pixels, in areas where ground-truth color is locally consistent, before and after the existing guided decoder. It still requires visual review: no scalar proxy proves that blotches are gone.
|
| 10 |
+
|
| 11 |
+
## Design
|
| 12 |
+
|
| 13 |
+
- Start from v2 revision 704fa80d792c3d759db91daa00b2dcfe6f0f6412, checked against SHA256 0e4c417375684a044860f8af3ac3a2fb44e1a5729254ca33abda758ea71aea6e. Do not continue the round-3 color-strength loss.
|
| 14 |
+
- Compare matched 3,000-step structure-only and structure-plus-distillation arms: seed 409, batch 16, learning rate 3e-5 decaying to 3e-6, frozen BatchNorm, BF16 student convolution, gradient clipping 1.
|
| 15 |
+
- Use the existing 16,230-image training pool; fixed 150 validation and 300 test memberships. These are previous development holdouts. Teacher pretraining overlap is unknown, especially for Imagenette derived from ImageNet.
|
| 16 |
+
- Ordinary-gray training inputs with flips and mild luminance augmentation. Gray and deterministic film-style inputs for validation/test.
|
| 17 |
+
- Remove inverse-frequency class reweighting and explicit chroma-strength pressure. Keep soft color-bin classification; add supervised multiscale color-gradient matching and a neutral-target penalty.
|
| 18 |
+
- Distillation arm also learns local ab predictions from piddnad/ddcolor_artistic, with confidence reduced when teacher colors strongly disagree with original ground truth. Teacher predictions are fallible and are not labels for original historical colors.
|
| 19 |
+
- Teacher: upstream DDColor code pinned at 2adb63f2656ac41cbdf7b894cddd94121a3faf13, checkpoint revision resolved once and recorded. Neutral sRGB input at 512 square, raw teacher Lab ab output. Full precision inference. Training teacher targets cached at 64 square after area pooling; student remains 256 square. The grayscale input contract follows the upstream pipeline; fixed 256-to-512 resizing is an experimental comparison choice.
|
| 20 |
+
- Evaluate the full pretrained teacher independently as a possible larger deployment option. No claim that distilling it guarantees the small model inherits its semantic capabilities.
|
| 21 |
+
|
| 22 |
+
## Selection gates
|
| 23 |
+
|
| 24 |
+
Every 500 steps, both gray/film validation conditions must satisfy all gates relative to v2: at least 15% less broad-patch excess after guided decoding; at least 10% less raw broad-patch excess; ab error no more than 5% higher; neutral spill and missed color no more than 1 percentage point higher; color coverage at least 95% of baseline. Rank eligible checkpoints by broad-patch proxy with smaller neutral/error terms. Keep baseline if none qualifies. Final test and real grayscale diagnostic montage are required before release.
|
| 25 |
+
|
| 26 |
+
## Persistence and compute
|
| 27 |
+
|
| 28 |
+
One L4 job, four-hour hard timeout and internal 3.5-hour graceful deadline checked through cache/training/evaluation loops. Expected roughly 1–3 hours, dependent on teacher inference speed. At the verified $0.80/hour L4 rate, four hours is about $3.20. No automatic write to main or stable: current direct-write permission was rejected in the previous attempt.
|
| 29 |
+
|
| 30 |
+
The job's finally block exports selected weights (or baseline if rejected), metrics, provenance, source and visual probes as a ZIP in ARTIFACT_BEGIN/ARTIFACT_CHUNK/ARTIFACT_END log records with SHA256. An ordinary failure is exported with status.json; hardware failure or a hard kill can still interrupt export. The teacher checkpoint is referenced by immutable Hub revision rather than duplicated. Optimizer state and the disposable teacher cache are not exported. Recover logs with an adequate tail and verify checksum before extraction.
|
| 31 |
+
|
| 32 |
+
Known limitation: the matched experiment isolates the extra teacher term, not every other change from v2. No claim of a clean causal attribution to any one of the shared changes.
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
## Round5 rerun — 2026-09-27
|
| 36 |
+
The previous round4 result was lost because stdout exceeded the archive cap. This rerun keeps the matched experiment design. Every 250 steps it commits model/config, optimizer and scheduler states to main under experiments/round5-20260927 and downloads changed files at the returned commit to verify SHA256. Evaluation and best selection are saved every 500 steps. Initial baseline and sources are saved before teacher/GPU work. Failed persistence aborts training after three attempts. The app's root checkpoint is only replaced after visual review; all candidates are available on main. Stdout contains compact status and metrics only. Optimizer state is saved for manual continuation, not an exact data-order resume implementation.
|
experiments/round5-20260927/source/data.py
ADDED
|
@@ -0,0 +1,116 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Fixed color-bin semantics and reproducible dataset membership."""
|
| 2 |
+
import hashlib
|
| 3 |
+
import io
|
| 4 |
+
import json
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
import numpy as np
|
| 7 |
+
from PIL import Image
|
| 8 |
+
import pyarrow.parquet as pq
|
| 9 |
+
from scipy.spatial import cKDTree
|
| 10 |
+
from skimage.color import rgb2lab
|
| 11 |
+
import torch
|
| 12 |
+
from torch.utils.data import Dataset
|
| 13 |
+
|
| 14 |
+
class ColorBins:
|
| 15 |
+
def __init__(self, centers, weights=None):
|
| 16 |
+
self.centers = np.asarray(centers, np.float32)
|
| 17 |
+
if self.centers.ndim != 2 or self.centers.shape[1] != 2:
|
| 18 |
+
raise ValueError("Expected Q by 2 color centers")
|
| 19 |
+
self.tree = cKDTree(self.centers)
|
| 20 |
+
self.weights = np.ones(len(self.centers), np.float32) if weights is None else np.asarray(weights,np.float32)
|
| 21 |
+
|
| 22 |
+
def encode(self, ab, k=5, sigma=5):
|
| 23 |
+
k = min(k, len(self.centers))
|
| 24 |
+
if sigma <= 0 or k < 1:
|
| 25 |
+
raise ValueError("Invalid encoding parameters")
|
| 26 |
+
dist, idx = self.tree.query(ab.reshape(-1,2), k=k)
|
| 27 |
+
dist, idx = dist.reshape(-1,k), idx.reshape(-1,k)
|
| 28 |
+
logw = -dist**2/(2*sigma**2)
|
| 29 |
+
logw -= logw.max(1, keepdims=True)
|
| 30 |
+
weight = np.exp(logw); weight /= weight.sum(1,keepdims=True)
|
| 31 |
+
shape = (*ab.shape[:-1], k)
|
| 32 |
+
return idx.reshape(shape).astype(np.int64), weight.reshape(shape).astype(np.float32)
|
| 33 |
+
|
| 34 |
+
def load_table(path):
|
| 35 |
+
table = pq.read_table(path)
|
| 36 |
+
return table
|
| 37 |
+
|
| 38 |
+
def validate_manifest(path, table, manifest):
|
| 39 |
+
with open(path,'rb') as f:
|
| 40 |
+
digest=hashlib.file_digest(f,'sha256').hexdigest()
|
| 41 |
+
if digest!=manifest['dataset_sha256'] or len(table)!=manifest['rows']:
|
| 42 |
+
raise ValueError('Dataset differs from pinned split manifest')
|
| 43 |
+
groups=[list(map(int,manifest[k])) for k in ['train','validation','test']]
|
| 44 |
+
if any(not g or len(g)!=len(set(g)) or min(g)<0 or max(g)>=len(table) for g in groups):
|
| 45 |
+
raise ValueError('Invalid or repeated indices in manifest')
|
| 46 |
+
if any(set(groups[i])&set(groups[j]) for i,j in [(0,1),(0,2),(1,2)]):
|
| 47 |
+
raise ValueError('Train/validation/test overlap')
|
| 48 |
+
|
| 49 |
+
def decode_image(table, index):
|
| 50 |
+
obj = table['image'][int(index)].as_py()
|
| 51 |
+
return Image.open(io.BytesIO(obj['bytes'])).convert('RGB')
|
| 52 |
+
|
| 53 |
+
def legacy_split(n, seed=0):
|
| 54 |
+
# Exactly reproduce datasets.Dataset.train_test_split(test_size=.05, seed=0).
|
| 55 |
+
order = np.random.default_rng(seed).permutation(n)
|
| 56 |
+
nval = int(np.ceil(.05*n))
|
| 57 |
+
return order[nval:], order[:nval]
|
| 58 |
+
|
| 59 |
+
def stratified_subset(indices, labels, per_class, seed):
|
| 60 |
+
rng = np.random.default_rng(seed)
|
| 61 |
+
selected = []
|
| 62 |
+
for label in sorted(set(labels)):
|
| 63 |
+
group = np.asarray([i for i in indices if labels[i] == label])
|
| 64 |
+
selected.extend(rng.choice(group, min(per_class,len(group)), replace=False).tolist())
|
| 65 |
+
return np.asarray(selected, dtype=np.int64)
|
| 66 |
+
|
| 67 |
+
class PhotoDataset(Dataset):
|
| 68 |
+
def __init__(self, table, indices, bins, size=256, augment=False):
|
| 69 |
+
self.table, self.indices, self.bins = table, list(map(int,indices)), bins
|
| 70 |
+
self.size, self.augment = size, augment
|
| 71 |
+
|
| 72 |
+
def __len__(self):
|
| 73 |
+
return len(self.indices)
|
| 74 |
+
|
| 75 |
+
def __getitem__(self, item):
|
| 76 |
+
# Keep square preprocessing for controlled comparison with historical runs.
|
| 77 |
+
# Production inference preserves aspect ratio separately.
|
| 78 |
+
image = decode_image(self.table,self.indices[item]).resize((self.size,self.size),Image.Resampling.BILINEAR)
|
| 79 |
+
arr = np.asarray(image,dtype=np.float32)/255
|
| 80 |
+
if self.augment and torch.rand(()) < .5:
|
| 81 |
+
arr = arr[:,::-1].copy()
|
| 82 |
+
lab = rgb2lab(arr).astype(np.float32)
|
| 83 |
+
idx, weight = self.bins.encode(lab[...,1:])
|
| 84 |
+
return (torch.from_numpy((lab[...,:1]/50-1).transpose(2,0,1).copy()),
|
| 85 |
+
torch.from_numpy(lab[...,1:].transpose(2,0,1).copy()),
|
| 86 |
+
torch.from_numpy(idx),torch.from_numpy(weight))
|
| 87 |
+
|
| 88 |
+
def estimate_weights(table, indices, bins, mix=.7, limit=1000, seed=0):
|
| 89 |
+
"""Re-estimate the prior ON checkpoint bins; never replace the bin vocabulary."""
|
| 90 |
+
if not 0 <= mix <= 1:
|
| 91 |
+
raise ValueError("Rebalance mixture must be in [0,1]")
|
| 92 |
+
rng=np.random.default_rng(seed)
|
| 93 |
+
chosen=rng.choice(indices,min(limit,len(indices)),replace=False)
|
| 94 |
+
counts=np.ones(len(bins.centers),np.float64)*1e-3
|
| 95 |
+
for index in chosen:
|
| 96 |
+
image=decode_image(table,index).resize((32,32))
|
| 97 |
+
ab=rgb2lab(np.asarray(image,dtype=np.float32)/255)[...,1:]
|
| 98 |
+
# Low chroma is a heuristic, not proof that an image was originally B&W.
|
| 99 |
+
if np.linalg.norm(ab,axis=-1).mean()<3:
|
| 100 |
+
continue
|
| 101 |
+
idx,weight=bins.encode(ab)
|
| 102 |
+
np.add.at(counts,idx.ravel(),weight.ravel())
|
| 103 |
+
prior=counts/counts.sum()
|
| 104 |
+
weights=1/((1-mix)*prior+mix/len(prior))
|
| 105 |
+
weights/=np.sum(prior*weights)
|
| 106 |
+
return weights.astype(np.float32),prior.astype(np.float32)
|
| 107 |
+
|
| 108 |
+
def write_manifest(path, table, train, validation, test, dataset_sha256):
|
| 109 |
+
groups=[set(map(int,g)) for g in [train,validation,test]]
|
| 110 |
+
assert not groups[0]&groups[1] and not groups[0]&groups[2] and not groups[1]&groups[2]
|
| 111 |
+
content={'dataset':'johnowhitaker/imagenette2-320','dataset_sha256':dataset_sha256,
|
| 112 |
+
'rows':len(table),'split_seed':0,'legacy_holdout_fraction':.05,
|
| 113 |
+
'train':list(map(int,train)),'validation':list(map(int,validation)), 'test':list(map(int,test)),
|
| 114 |
+
'caveat':'Matches the recent seed-0 holdout; older upstream training exposure is not established.'}
|
| 115 |
+
Path(path).write_text(json.dumps(content,indent=2))
|
| 116 |
+
return content
|
experiments/round5-20260927/source/inference.py
ADDED
|
@@ -0,0 +1,64 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Aspect-preserving inference; infer chroma globally and retain original luminance."""
|
| 2 |
+
import argparse
|
| 3 |
+
import math
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
import numpy as np
|
| 6 |
+
from PIL import Image, ImageOps
|
| 7 |
+
import torch
|
| 8 |
+
import torch.nn.functional as F
|
| 9 |
+
from skimage.color import rgb2lab, lab2rgb
|
| 10 |
+
from model import load_model
|
| 11 |
+
from spatial import guided_chroma
|
| 12 |
+
|
| 13 |
+
@torch.inference_mode()
|
| 14 |
+
def colorize(model, image, size=256, temperature=0.38, saturation=1.0, flip_tta=False,
|
| 15 |
+
guided_radius=8, guided_epsilon=.001):
|
| 16 |
+
if size < 8 or not math.isfinite(saturation) or saturation < 0:
|
| 17 |
+
raise ValueError("size must be >=8 and saturation finite and nonnegative")
|
| 18 |
+
image = ImageOps.exif_transpose(image).convert("RGB")
|
| 19 |
+
rgb = np.asarray(image, dtype=np.float32) / 255.0
|
| 20 |
+
luminance = rgb2lab(rgb)[..., 0].astype(np.float32)
|
| 21 |
+
h, w = luminance.shape
|
| 22 |
+
scale = min(size / max(h, w), 1.0)
|
| 23 |
+
target = (max(8, round(h*scale)), max(8, round(w*scale)))
|
| 24 |
+
device = next(model.parameters()).device
|
| 25 |
+
L = torch.from_numpy(luminance)[None, None].to(device) / 50 - 1
|
| 26 |
+
small = F.interpolate(L, size=target, mode="bilinear", align_corners=False, antialias=True)
|
| 27 |
+
logits = model(small)
|
| 28 |
+
if flip_tta:
|
| 29 |
+
logits = (logits + model(small.flip(-1)).flip(-1)) * 0.5
|
| 30 |
+
ab = model.decode(logits, temperature)
|
| 31 |
+
ab = guided_chroma(small, ab, guided_radius, guided_epsilon)
|
| 32 |
+
ab = F.interpolate(ab, size=(h,w), mode="bilinear", align_corners=False)
|
| 33 |
+
ab = ab[0].permute(1,2,0).cpu().numpy() * saturation
|
| 34 |
+
lab = np.concatenate([luminance[...,None], ab], axis=-1)
|
| 35 |
+
result = np.clip(lab2rgb(lab), 0, 1)
|
| 36 |
+
return Image.fromarray(np.rint(result * 255).astype(np.uint8))
|
| 37 |
+
|
| 38 |
+
def main():
|
| 39 |
+
p = argparse.ArgumentParser(__doc__)
|
| 40 |
+
p.add_argument("images", nargs="+")
|
| 41 |
+
p.add_argument("--model", required=True)
|
| 42 |
+
p.add_argument("--revision", default=None)
|
| 43 |
+
p.add_argument("--size", type=int, default=256)
|
| 44 |
+
p.add_argument("--temperature", type=float, default=.38)
|
| 45 |
+
p.add_argument("--saturation", type=float, default=1)
|
| 46 |
+
p.add_argument("--flip-tta", action="store_true")
|
| 47 |
+
p.add_argument("--guided-radius",type=int,default=8,help="Chroma smoothing radius at model resolution;0 disables")
|
| 48 |
+
p.add_argument("--guided-epsilon",type=float,default=.001)
|
| 49 |
+
p.add_argument("--output-dir", default="colorized")
|
| 50 |
+
p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
|
| 51 |
+
a = p.parse_args()
|
| 52 |
+
torch.set_num_threads(min(torch.get_num_threads(),4))
|
| 53 |
+
model = load_model(a.model, a.revision, a.device)
|
| 54 |
+
out = Path(a.output_dir); out.mkdir(parents=True, exist_ok=True)
|
| 55 |
+
for file in a.images:
|
| 56 |
+
with Image.open(file) as im:
|
| 57 |
+
result = colorize(model, im, a.size, a.temperature, a.saturation, a.flip_tta,
|
| 58 |
+
a.guided_radius,a.guided_epsilon)
|
| 59 |
+
dest = out / (Path(file).stem + "_colorized.png")
|
| 60 |
+
result.save(dest); print(dest)
|
| 61 |
+
|
| 62 |
+
if __name__ == "__main__":
|
| 63 |
+
main()
|
| 64 |
+
|
experiments/round5-20260927/source/metrics.py
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Metrics are fidelity/artifact proxies; none certifies plausible color by itself."""
|
| 2 |
+
import numpy as np
|
| 3 |
+
import torch
|
| 4 |
+
|
| 5 |
+
def soft_ce(logits, idx, weights, class_weights=None):
|
| 6 |
+
# FP32 reduction under autocast; no full target distribution materialized.
|
| 7 |
+
logits=logits.float()
|
| 8 |
+
idx=idx.permute(0,3,1,2)
|
| 9 |
+
weights=weights.permute(0,3,1,2)
|
| 10 |
+
nll=-((logits.gather(1,idx)-logits.logsumexp(1,keepdim=True))*weights).sum(1)
|
| 11 |
+
if class_weights is not None:
|
| 12 |
+
nll=nll*class_weights[idx[:,0]]
|
| 13 |
+
return nll.mean()
|
| 14 |
+
|
| 15 |
+
def per_image_metrics(L, pred, target):
|
| 16 |
+
error=(pred-target).square().sum(1).sqrt().mean((1,2))
|
| 17 |
+
chroma=pred.square().sum(1).sqrt().mean((1,2))
|
| 18 |
+
true_chroma=target.square().sum(1).sqrt().mean((1,2))
|
| 19 |
+
seams=[]
|
| 20 |
+
for dim in [-1,-2]:
|
| 21 |
+
dp=pred.diff(dim=dim).square().sum(1).sqrt()
|
| 22 |
+
dt=target.diff(dim=dim).square().sum(1).sqrt()
|
| 23 |
+
dl=L.diff(dim=dim).abs()[:,0]*50
|
| 24 |
+
# Penalize excess chroma discontinuity where luminance is flat.
|
| 25 |
+
mask=(dl<2).float()
|
| 26 |
+
seams.append(((dp-dt).relu()*mask).sum((1,2))/mask.sum((1,2)).clamp_min(1))
|
| 27 |
+
return {'ab_error':error,'chroma':chroma,'target_chroma':true_chroma,
|
| 28 |
+
'excess_chroma_edge':sum(seams)/2}
|
| 29 |
+
|
| 30 |
+
def summarize(rows):
|
| 31 |
+
keys=['ce','ab_error','chroma','target_chroma','excess_chroma_edge']
|
| 32 |
+
result={k:float(np.mean([r[k] for r in rows])) for k in keys}
|
| 33 |
+
result['n']=len(rows)
|
| 34 |
+
result['chroma_ratio']=result['chroma']/max(result['target_chroma'],1e-9)
|
| 35 |
+
return result
|
experiments/round5-20260927/source/model.py
ADDED
|
@@ -0,0 +1,122 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Shared, checkpoint-compatible Mini U-Net. Parameters remain under 4M."""
|
| 2 |
+
import json
|
| 3 |
+
import math
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
import numpy as np
|
| 6 |
+
import torch
|
| 7 |
+
from torch import nn
|
| 8 |
+
import torch.nn.functional as F
|
| 9 |
+
from huggingface_hub import PyTorchModelHubMixin, snapshot_download
|
| 10 |
+
from safetensors.torch import load_file, save_file
|
| 11 |
+
def double_conv(in_ch, out_ch):
|
| 12 |
+
return nn.Sequential(
|
| 13 |
+
nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False),
|
| 14 |
+
nn.BatchNorm2d(out_ch),
|
| 15 |
+
nn.ReLU(inplace=True),
|
| 16 |
+
nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False),
|
| 17 |
+
nn.BatchNorm2d(out_ch),
|
| 18 |
+
nn.ReLU(inplace=True),
|
| 19 |
+
)
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
class DilatedContextBlock(nn.Module):
|
| 23 |
+
def __init__(self, channels, mid_ch=96, dilations=(2, 4, 8)):
|
| 24 |
+
super().__init__()
|
| 25 |
+
self.proj_in = nn.Sequential(
|
| 26 |
+
nn.Conv2d(channels, mid_ch, 1, bias=False),
|
| 27 |
+
nn.BatchNorm2d(mid_ch),
|
| 28 |
+
nn.ReLU(inplace=True),
|
| 29 |
+
)
|
| 30 |
+
layers = []
|
| 31 |
+
for d in dilations:
|
| 32 |
+
layers += [
|
| 33 |
+
nn.Conv2d(mid_ch, mid_ch, 3, padding=d, dilation=d, bias=False),
|
| 34 |
+
nn.BatchNorm2d(mid_ch),
|
| 35 |
+
nn.ReLU(inplace=True),
|
| 36 |
+
]
|
| 37 |
+
self.dilated = nn.Sequential(*layers)
|
| 38 |
+
self.proj_out = nn.Sequential(
|
| 39 |
+
nn.Conv2d(mid_ch, channels, 1, bias=False),
|
| 40 |
+
nn.BatchNorm2d(channels),
|
| 41 |
+
)
|
| 42 |
+
nn.init.zeros_(self.proj_out[-1].weight)
|
| 43 |
+
self.relu = nn.ReLU(inplace=True)
|
| 44 |
+
|
| 45 |
+
def forward(self, x):
|
| 46 |
+
y = self.proj_in(x)
|
| 47 |
+
y = self.dilated(y)
|
| 48 |
+
y = self.proj_out(y)
|
| 49 |
+
return self.relu(x + y)
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
class SmallUNetColorizer(
|
| 53 |
+
nn.Module,
|
| 54 |
+
PyTorchModelHubMixin,
|
| 55 |
+
pipeline_tag="image-to-image",
|
| 56 |
+
license="apache-2.0",
|
| 57 |
+
tags=["colorization", "unet", "image-to-image", "classification"],
|
| 58 |
+
):
|
| 59 |
+
def __init__(self, bin_centers, in_ch: int = 1, base: int = 44,
|
| 60 |
+
context_mid_ch: int = 96, context_dilations=(2, 4, 8)):
|
| 61 |
+
super().__init__()
|
| 62 |
+
self.in_ch, self.base = in_ch, base
|
| 63 |
+
num_bins = len(bin_centers)
|
| 64 |
+
self.num_bins = num_bins
|
| 65 |
+
self.register_buffer("bin_centers", torch.tensor(bin_centers, dtype=torch.float32))
|
| 66 |
+
|
| 67 |
+
self.enc1 = double_conv(in_ch, base)
|
| 68 |
+
self.enc2 = double_conv(base, base * 2)
|
| 69 |
+
self.enc3 = double_conv(base * 2, base * 4)
|
| 70 |
+
self.enc4 = double_conv(base * 4, base * 8)
|
| 71 |
+
self.pool = nn.MaxPool2d(2)
|
| 72 |
+
self.context = DilatedContextBlock(base * 8, mid_ch=context_mid_ch,
|
| 73 |
+
dilations=tuple(context_dilations))
|
| 74 |
+
self.up3 = nn.ConvTranspose2d(base * 8, base * 4, 2, stride=2)
|
| 75 |
+
self.dec3 = double_conv(base * 8, base * 4)
|
| 76 |
+
self.up2 = nn.ConvTranspose2d(base * 4, base * 2, 2, stride=2)
|
| 77 |
+
self.dec2 = double_conv(base * 4, base * 2)
|
| 78 |
+
self.up1 = nn.ConvTranspose2d(base * 2, base, 2, stride=2)
|
| 79 |
+
self.dec1 = double_conv(base * 2, base)
|
| 80 |
+
self.out_conv = nn.Conv2d(base, num_bins, 1)
|
| 81 |
+
|
| 82 |
+
def forward(self, x):
|
| 83 |
+
h, w = x.shape[-2:]
|
| 84 |
+
x = F.pad(x, (0, (-w) % 8, 0, (-h) % 8), mode="replicate")
|
| 85 |
+
e1 = self.enc1(x)
|
| 86 |
+
e2 = self.enc2(self.pool(e1))
|
| 87 |
+
e3 = self.enc3(self.pool(e2))
|
| 88 |
+
e4 = self.context(self.enc4(self.pool(e3)))
|
| 89 |
+
d3 = self.dec3(torch.cat([self.up3(e4), e3], dim=1))
|
| 90 |
+
d2 = self.dec2(torch.cat([self.up2(d3), e2], dim=1))
|
| 91 |
+
d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1))
|
| 92 |
+
return self.out_conv(d1)[..., :h, :w]
|
| 93 |
+
|
| 94 |
+
def decode(self, logits, temperature: float = 0.38):
|
| 95 |
+
if not math.isfinite(temperature) or temperature <= 0:
|
| 96 |
+
raise ValueError("temperature must be finite and positive")
|
| 97 |
+
probs_t = F.softmax(logits.float() / temperature, dim=1)
|
| 98 |
+
return torch.einsum("bqhw,qc->bchw", probs_t, self.bin_centers)
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
def load_model(source, revision=None, device="cpu"):
|
| 102 |
+
path = Path(source)
|
| 103 |
+
if not path.is_dir():
|
| 104 |
+
path = Path(snapshot_download(source, revision=revision,
|
| 105 |
+
allow_patterns=["config.json", "model.safetensors"]))
|
| 106 |
+
config = json.loads((path / "config.json").read_text())
|
| 107 |
+
state = load_file(str(path / "model.safetensors"))
|
| 108 |
+
centers = torch.tensor(config["bin_centers"], dtype=torch.float32)
|
| 109 |
+
if not torch.equal(centers, state["bin_centers"]):
|
| 110 |
+
raise ValueError("Checkpoint config and state color bins differ; refusing ambiguous decode")
|
| 111 |
+
model = SmallUNetColorizer(**config)
|
| 112 |
+
model.load_state_dict(state, strict=True)
|
| 113 |
+
model.to(device).eval()
|
| 114 |
+
return model
|
| 115 |
+
|
| 116 |
+
def save_model(model, path):
|
| 117 |
+
path = Path(path); path.mkdir(parents=True, exist_ok=True)
|
| 118 |
+
model.save_pretrained(path)
|
| 119 |
+
# Mixin config can retain constructor bins; use the actual authoritative buffer.
|
| 120 |
+
cfg = json.loads((path / "config.json").read_text())
|
| 121 |
+
cfg["bin_centers"] = model.bin_centers.detach().cpu().tolist()
|
| 122 |
+
(path / "config.json").write_text(json.dumps(cfg, indent=2))
|
experiments/round5-20260927/source/persistence.py
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Commit run artifacts atomically and verify every changed file by SHA256."""
|
| 2 |
+
import hashlib, json, os, time
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
from huggingface_hub import HfApi, CommitOperationAdd, hf_hub_download
|
| 5 |
+
|
| 6 |
+
REPO = 'User-2468/mini-unet-colorizer'
|
| 7 |
+
PREFIX = 'experiments/round5-20260927'
|
| 8 |
+
|
| 9 |
+
class DurableRun:
|
| 10 |
+
def __init__(self, root):
|
| 11 |
+
if not os.environ.get('HF_TOKEN'):
|
| 12 |
+
raise RuntimeError('A write token is required before starting')
|
| 13 |
+
self.api = HfApi(token=os.environ['HF_TOKEN'])
|
| 14 |
+
if self.api.whoami()['name'] != 'User-2468':
|
| 15 |
+
raise RuntimeError('Unexpected account')
|
| 16 |
+
self.root = Path(root)
|
| 17 |
+
self.saved = {}
|
| 18 |
+
|
| 19 |
+
def sync(self, reason):
|
| 20 |
+
files = [p for p in sorted(self.root.rglob('*')) if p.is_file()]
|
| 21 |
+
hashes = {p.relative_to(self.root).as_posix(): hashlib.sha256(p.read_bytes()).hexdigest() for p in files}
|
| 22 |
+
changed = [p for p in files if self.saved.get(p.relative_to(self.root).as_posix()) != hashes[p.relative_to(self.root).as_posix()]]
|
| 23 |
+
if not changed:
|
| 24 |
+
return
|
| 25 |
+
for attempt in range(3):
|
| 26 |
+
try:
|
| 27 |
+
head = self.api.model_info(REPO, revision='main').sha
|
| 28 |
+
operations = [CommitOperationAdd(path_in_repo=PREFIX+'/'+p.relative_to(self.root).as_posix(), path_or_fileobj=str(p)) for p in changed]
|
| 29 |
+
operations.append(CommitOperationAdd(path_in_repo=PREFIX+'/SHA256SUMS.json', path_or_fileobj=json.dumps(hashes,sort_keys=True,indent=2).encode()))
|
| 30 |
+
commit = self.api.create_commit(repo_id=REPO, revision='main', parent_commit=head, operations=operations, commit_message='Colorizer round5: '+reason)
|
| 31 |
+
for p in changed:
|
| 32 |
+
name = p.relative_to(self.root).as_posix()
|
| 33 |
+
downloaded = hf_hub_download(REPO, PREFIX+'/'+name, revision=commit.oid, force_download=True, token=os.environ['HF_TOKEN'])
|
| 34 |
+
if hashlib.sha256(Path(downloaded).read_bytes()).hexdigest() != hashes[name]:
|
| 35 |
+
raise RuntimeError('Remote checksum mismatch: '+name)
|
| 36 |
+
self.saved = hashes
|
| 37 |
+
print('PERSISTED',reason,commit.oid,'verified_files',len(changed),flush=True)
|
| 38 |
+
return
|
| 39 |
+
except Exception:
|
| 40 |
+
if attempt == 2:
|
| 41 |
+
raise
|
| 42 |
+
time.sleep(2**attempt)
|
experiments/round5-20260927/source/previous_manifest.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
experiments/round5-20260927/source/spatial.py
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Controlled spatial decoding alternatives; no learned parameters."""
|
| 2 |
+
import math
|
| 3 |
+
import torch
|
| 4 |
+
import torch.nn.functional as F
|
| 5 |
+
|
| 6 |
+
def box_mean(x,radius):
|
| 7 |
+
# Border-normalized local windows, also work on images smaller than kernel.
|
| 8 |
+
return F.avg_pool2d(x,2*radius+1,stride=1,padding=radius,count_include_pad=False)
|
| 9 |
+
|
| 10 |
+
def guided_chroma(L,ab,radius=4,epsilon=.001):
|
| 11 |
+
"""Scalar luminance-guided local linear filter (He et al., ECCV 2010).
|
| 12 |
+
|
| 13 |
+
L is the network's [-1,1] luminance. Epsilon is in [0,1] luminance units.
|
| 14 |
+
Uses luminance only; reference colors never enter inference.
|
| 15 |
+
"""
|
| 16 |
+
if not isinstance(radius,int) or radius<0 or not math.isfinite(epsilon) or epsilon<=0:
|
| 17 |
+
raise ValueError('Invalid guided-filter radius or epsilon')
|
| 18 |
+
if radius==0:return ab
|
| 19 |
+
I=(L.float()+1)/2;p=ab.float()
|
| 20 |
+
mi=box_mean(I,radius);mp=box_mean(p,radius)
|
| 21 |
+
var=(box_mean(I*I,radius)-mi*mi).clamp_min(0)
|
| 22 |
+
cov=box_mean(I*p,radius)-mi*mp
|
| 23 |
+
a=cov/(var+epsilon);b=mp-a*mi
|
| 24 |
+
return box_mean(a,radius)*I+box_mean(b,radius)
|
| 25 |
+
|
| 26 |
+
def spatial_decode(model,logits,L,temperature=.38,pool=1,radius=0,epsilon=.001):
|
| 27 |
+
if pool<1 or not isinstance(pool,int):raise ValueError('pool must be a positive integer')
|
| 28 |
+
if pool>1:
|
| 29 |
+
# Average evidence before annealing, then upsample chroma, not RGB.
|
| 30 |
+
small=F.avg_pool2d(logits,pool,ceil_mode=True,count_include_pad=False)
|
| 31 |
+
ab=model.decode(small,temperature)
|
| 32 |
+
ab=F.interpolate(ab,size=L.shape[-2:],mode='bilinear',align_corners=False)
|
| 33 |
+
else:ab=model.decode(logits,temperature)
|
| 34 |
+
return guided_chroma(L,ab,radius,epsilon)
|
experiments/round5-20260927/source/train.py
ADDED
|
@@ -0,0 +1,259 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Round 4: matched multiscale structure and DDColor distillation experiments."""
|
| 2 |
+
import os,sys,json,time,random,hashlib,traceback,shutil
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
import numpy as np
|
| 5 |
+
import cv2
|
| 6 |
+
from PIL import Image,ImageDraw
|
| 7 |
+
from skimage import data as sample_data
|
| 8 |
+
from skimage.color import rgb2lab,lab2rgb
|
| 9 |
+
import pyarrow as pa
|
| 10 |
+
import pyarrow.parquet as pq
|
| 11 |
+
import torch
|
| 12 |
+
import torch.nn.functional as F
|
| 13 |
+
from torch.utils.data import Dataset,DataLoader
|
| 14 |
+
from huggingface_hub import HfApi,hf_hub_download,PyTorchModelHubMixin
|
| 15 |
+
from model import load_model,save_model
|
| 16 |
+
from data import decode_image,ColorBins
|
| 17 |
+
from metrics import soft_ce,per_image_metrics
|
| 18 |
+
from spatial import guided_chroma
|
| 19 |
+
|
| 20 |
+
from persistence import DurableRun
|
| 21 |
+
OUT=Path('round5');OUT.mkdir(exist_ok=True)
|
| 22 |
+
DURABLE=None
|
| 23 |
+
START=time.monotonic();DEADLINE=START+3.5*3600
|
| 24 |
+
DEVICE='cuda';SEED=409;STEPS=3000
|
| 25 |
+
BASE_REV='704fa80d792c3d759db91daa00b2dcfe6f0f6412'
|
| 26 |
+
BASE_HASH='0e4c417375684a044860f8af3ac3a2fb44e1a5729254ca33abda758ea71aea6e'
|
| 27 |
+
TEACHER_ID='piddnad/ddcolor_artistic'
|
| 28 |
+
torch.set_num_threads(6)
|
| 29 |
+
|
| 30 |
+
def deadline():
|
| 31 |
+
if time.monotonic()>DEADLINE:raise TimeoutError('Graceful export before remote hard timeout')
|
| 32 |
+
|
| 33 |
+
def gray_arrays(im,style='gray'):
|
| 34 |
+
im=im.resize((256,256),Image.Resampling.BILINEAR)
|
| 35 |
+
rgb=np.asarray(im,np.float32)/255
|
| 36 |
+
lab=rgb2lab(rgb).astype(np.float32)
|
| 37 |
+
if style=='film':
|
| 38 |
+
g=(rgb*np.array([.42,.45,.13],np.float32)).sum(-1)
|
| 39 |
+
g=np.clip((g**1.15-.5)*.8+.5,0,1)
|
| 40 |
+
else:g=np.asarray(im.convert('L'),np.float32)/255
|
| 41 |
+
# Feed the same neutral sRGB input to teacher and student.
|
| 42 |
+
gray=np.repeat(g[...,None],3,-1).astype(np.float32)
|
| 43 |
+
L=rgb2lab(gray)[...,0].astype(np.float32)
|
| 44 |
+
return (L[None]/50-1).copy(),lab[...,1:].transpose(2,0,1).copy(),gray
|
| 45 |
+
|
| 46 |
+
class Photos(Dataset):
|
| 47 |
+
def __init__(self,table,ids,bins=None,cache=None,augment=False,style='gray'):
|
| 48 |
+
self.table,self.ids,self.bins,self.cache,self.augment,self.style=table,ids,bins,cache,augment,style
|
| 49 |
+
def __len__(self):return len(self.ids)
|
| 50 |
+
def __getitem__(self,k):
|
| 51 |
+
L,ab,gray=gray_arrays(decode_image(self.table,self.ids[k]),self.style)
|
| 52 |
+
teach=np.array(self.cache[k],np.float32) if self.cache is not None else np.zeros((2,64,64),np.float32)
|
| 53 |
+
if self.augment:
|
| 54 |
+
if random.random()<.5:L=L[...,::-1].copy();ab=ab[...,::-1].copy();teach=teach[...,::-1].copy()
|
| 55 |
+
if random.random()<.3:L=np.clip(L*random.uniform(.85,1.15)+random.uniform(-.08,.08),-1,1)
|
| 56 |
+
if self.bins is not None:
|
| 57 |
+
idx,w=self.bins.encode(ab.transpose(1,2,0));return torch.from_numpy(L),torch.from_numpy(ab),torch.from_numpy(idx),torch.from_numpy(w),torch.from_numpy(teach)
|
| 58 |
+
return torch.from_numpy(L),torch.from_numpy(ab),torch.from_numpy(gray.transpose(2,0,1)),self.ids[k]
|
| 59 |
+
|
| 60 |
+
def patch_excess(p,t):
|
| 61 |
+
"""Excess color differences across 4–64px regions, in target-flat areas."""
|
| 62 |
+
values=[]
|
| 63 |
+
for scale in [4,16]:
|
| 64 |
+
a=F.avg_pool2d(p,scale);b=F.avg_pool2d(t,scale)
|
| 65 |
+
for offset in [1,4]:
|
| 66 |
+
for dim in [-1,-2]:
|
| 67 |
+
x=a.narrow(dim,offset,a.shape[dim]-offset)-a.narrow(dim,0,a.shape[dim]-offset)
|
| 68 |
+
y=b.narrow(dim,offset,b.shape[dim]-offset)-b.narrow(dim,0,b.shape[dim]-offset)
|
| 69 |
+
dx=x.norm(dim=1);dy=y.norm(dim=1);mask=(dy<3).float()
|
| 70 |
+
values.append(((dx-dy-1).relu()*mask).sum((1,2))/mask.sum((1,2)).clamp_min(1))
|
| 71 |
+
return sum(values)/len(values)
|
| 72 |
+
|
| 73 |
+
def structural_loss(p,t):
|
| 74 |
+
terms=[]
|
| 75 |
+
for scale in [4,16]:
|
| 76 |
+
a=F.avg_pool2d(p,scale);b=F.avg_pool2d(t,scale)
|
| 77 |
+
for offset in [1,4]:
|
| 78 |
+
for dim in [-1,-2]:
|
| 79 |
+
x=a.narrow(dim,offset,a.shape[dim]-offset)-a.narrow(dim,0,a.shape[dim]-offset)
|
| 80 |
+
y=b.narrow(dim,offset,b.shape[dim]-offset)-b.narrow(dim,0,b.shape[dim]-offset)
|
| 81 |
+
terms.append(F.smooth_l1_loss(x/10,y/10,beta=.3))
|
| 82 |
+
return sum(terms)/len(terms)
|
| 83 |
+
|
| 84 |
+
def additional_losses(p,t,teacher):
|
| 85 |
+
small=F.avg_pool2d(p,4);target=F.avg_pool2d(t,4)
|
| 86 |
+
# Teacher is fallible; reduce guidance when it conflicts strongly with ground truth.
|
| 87 |
+
confidence=torch.exp(-(teacher-target).norm(dim=1,keepdim=True)/30).detach()
|
| 88 |
+
kd=(F.smooth_l1_loss(small/20,teacher/20,beta=.5,reduction='none')*confidence).mean()
|
| 89 |
+
neutral=(t.norm(dim=1)<3).float()
|
| 90 |
+
neutral_loss=(p.norm(dim=1)*neutral).sum()/neutral.sum().clamp_min(1)/20
|
| 91 |
+
return structural_loss(p,t),kd,neutral_loss
|
| 92 |
+
|
| 93 |
+
def measure(L,p,t,raw=None):
|
| 94 |
+
r=per_image_metrics(L,p,t);C=p.norm(dim=1);T=t.norm(dim=1)
|
| 95 |
+
for name,mask,event in [('missed_color',T>12,C<5),('neutral_spill',T<3,C>10)]:
|
| 96 |
+
r[name]=(mask&event).sum((1,2))/mask.sum((1,2)).clamp_min(1)
|
| 97 |
+
r['color_coverage']=(C>10).float().mean((1,2));r['patch_excess']=patch_excess(p,t)
|
| 98 |
+
r['raw_patch_excess']=patch_excess(p if raw is None else raw,t)
|
| 99 |
+
return r
|
| 100 |
+
|
| 101 |
+
@torch.inference_mode()
|
| 102 |
+
def teacher_predict(teacher,gray):
|
| 103 |
+
gray=F.interpolate(gray.to(DEVICE),size=(512,512),mode='bilinear',align_corners=False)
|
| 104 |
+
# Full precision is deliberate: spectral normalization / attention are not assumed BF16-safe.
|
| 105 |
+
out=teacher(gray).float()
|
| 106 |
+
assert out.shape[1]==2 and torch.isfinite(out).all()
|
| 107 |
+
return F.interpolate(out,size=(256,256),mode='bilinear',align_corners=False)
|
| 108 |
+
|
| 109 |
+
@torch.inference_mode()
|
| 110 |
+
def evaluate(model,table,ids,is_teacher=False):
|
| 111 |
+
result={};model.eval()
|
| 112 |
+
for style in ['gray','film']:
|
| 113 |
+
rows=[]
|
| 114 |
+
for L,t,gray,indices in DataLoader(Photos(table,ids,style=style),batch_size=2 if is_teacher else 8,num_workers=4):
|
| 115 |
+
deadline();L,t=L.to(DEVICE),t.to(DEVICE)
|
| 116 |
+
raw=teacher_predict(model,gray) if is_teacher else model.decode(model(L),.38)
|
| 117 |
+
pred=raw if is_teacher else guided_chroma(L,raw,8)
|
| 118 |
+
metrics=measure(L,pred,t,raw)
|
| 119 |
+
for j,index in enumerate(indices):rows.append({'index':int(index)}|{k:float(v[j]) for k,v in metrics.items()})
|
| 120 |
+
keep=[x for x in rows if x['target_chroma']>=5]
|
| 121 |
+
result[style]={'n_total':len(rows),'n_color':len(keep),'summary':{k:float(np.mean([x[k] for x in keep])) for k in keep[0] if k!='index'},'per_image':rows}
|
| 122 |
+
return result
|
| 123 |
+
|
| 124 |
+
def selection(result,baseline):
|
| 125 |
+
scores=[];eligible=True
|
| 126 |
+
for style in ['gray','film']:
|
| 127 |
+
s=result[style]['summary'];b=baseline[style]['summary']
|
| 128 |
+
eligible &= (s['ab_error']<=b['ab_error']*1.05 and s['neutral_spill']<=b['neutral_spill']+.01
|
| 129 |
+
and s['missed_color']<=b['missed_color']+.01 and s['color_coverage']>=.95*b['color_coverage']
|
| 130 |
+
and s['patch_excess']<.85*b['patch_excess'] and s['raw_patch_excess']<.9*b['raw_patch_excess'])
|
| 131 |
+
scores.append(s['patch_excess']/max(b['patch_excess'],1e-6)+.3*s['neutral_spill']+.2*s['ab_error']/b['ab_error'])
|
| 132 |
+
return float(np.mean(scores)),bool(eligible)
|
| 133 |
+
|
| 134 |
+
@torch.inference_mode()
|
| 135 |
+
def probes(model,tag,teacher=False):
|
| 136 |
+
names=['astronaut','coffee','chelsea','rocket','camera','coins','moon'];images=[]
|
| 137 |
+
for name in names:
|
| 138 |
+
original=Image.fromarray(getattr(sample_data,name)()).convert('RGB')
|
| 139 |
+
L,_,gray=gray_arrays(original);x=torch.from_numpy(L)[None].to(DEVICE)
|
| 140 |
+
ab=teacher_predict(model,torch.from_numpy(gray.transpose(2,0,1))[None]) if teacher else guided_chroma(x,model.decode(model(x),.38),8)
|
| 141 |
+
lab=np.concatenate([(L[0]*50+50)[...,None],ab[0].cpu().numpy().transpose(1,2,0)],-1)
|
| 142 |
+
rgb=np.uint8(np.clip(lab2rgb(lab),0,1)*255)
|
| 143 |
+
images.append((name,Image.fromarray(np.uint8(gray*255)),Image.fromarray(rgb)))
|
| 144 |
+
canvas=Image.new('RGB',(512,len(names)*280),'white');d=ImageDraw.Draw(canvas)
|
| 145 |
+
for i,(name,gray,col) in enumerate(images):
|
| 146 |
+
canvas.paste(gray,(0,i*280+24));canvas.paste(col,(256,i*280+24));d.text((4,i*280+4),name+' input | '+tag,fill='black')
|
| 147 |
+
canvas.save(OUT/(tag+'.jpg'),quality=90)
|
| 148 |
+
|
| 149 |
+
def write(name,obj):
|
| 150 |
+
(OUT/name).write_text(json.dumps(obj,indent=2));return obj
|
| 151 |
+
|
| 152 |
+
def run():
|
| 153 |
+
global DURABLE
|
| 154 |
+
DURABLE=DurableRun(OUT)
|
| 155 |
+
source=OUT/'source';source.mkdir(exist_ok=True)
|
| 156 |
+
for p in Path('.').glob('*.py'):shutil.copy2(p,source/p.name)
|
| 157 |
+
shutil.copy2('previous_manifest.json',source/'previous_manifest.json')
|
| 158 |
+
shutil.copy2('PROTOCOL.md',source/'PROTOCOL.md')
|
| 159 |
+
shutil.copytree('vendor_ddcolor',source/'vendor_ddcolor',dirs_exist_ok=True,ignore=shutil.ignore_patterns('__pycache__'))
|
| 160 |
+
write('status.json',{'status':'initializing'})
|
| 161 |
+
torch.manual_seed(SEED);np.random.seed(SEED);random.seed(SEED);torch.backends.cudnn.benchmark=True
|
| 162 |
+
root=Path('init');root.mkdir(exist_ok=True)
|
| 163 |
+
for name in ['model.safetensors','config.json']:
|
| 164 |
+
root.joinpath(name).write_bytes(Path(hf_hub_download('User-2468/mini-unet-colorizer',name,revision=BASE_REV)).read_bytes())
|
| 165 |
+
assert hashlib.sha256((root/'model.safetensors').read_bytes()).hexdigest()==BASE_HASH
|
| 166 |
+
model=load_model(root,device=DEVICE);save_model(model,OUT/'best')
|
| 167 |
+
write('selection.json',{'arm':'v2_baseline','step':0,'accepted':False})
|
| 168 |
+
DURABLE.sync('baseline checkpoint and source before GPU work')
|
| 169 |
+
# Pin source by bundled git revision and resolve the teacher model revision exactly once.
|
| 170 |
+
from ddcolor import DDColor
|
| 171 |
+
class DDColorHF(DDColor,PyTorchModelHubMixin):
|
| 172 |
+
def __init__(self,config=None,**kw):super().__init__(**({**config,**kw} if isinstance(config,dict) else kw))
|
| 173 |
+
teacher_rev=HfApi().model_info(TEACHER_ID).sha
|
| 174 |
+
teacher=DDColorHF.from_pretrained(TEACHER_ID,revision=teacher_rev).to(DEVICE).eval()
|
| 175 |
+
for p in teacher.parameters():p.requires_grad_(False)
|
| 176 |
+
# Verify the actual GPU teacher batch and student backward path before data preparation.
|
| 177 |
+
assert torch.cuda.is_available()
|
| 178 |
+
print('PREFLIGHT teacher forward',flush=True)
|
| 179 |
+
smoke_teacher=teacher_predict(teacher,torch.full((4,3,256,256),.5))
|
| 180 |
+
assert smoke_teacher.shape==(4,2,256,256)
|
| 181 |
+
del smoke_teacher
|
| 182 |
+
model.eval();model.zero_grad(set_to_none=True)
|
| 183 |
+
with torch.autocast('cuda',dtype=torch.bfloat16):
|
| 184 |
+
smoke_logits=model(torch.zeros(2,1,256,256,device=DEVICE))
|
| 185 |
+
smoke_pred=model.decode(smoke_logits,.38)
|
| 186 |
+
smoke_losses=additional_losses(smoke_pred,torch.zeros_like(smoke_pred),torch.zeros(2,2,64,64,device=DEVICE))
|
| 187 |
+
sum(smoke_losses).backward()
|
| 188 |
+
assert all(torch.isfinite(p.grad).all() for p in model.parameters() if p.grad is not None)
|
| 189 |
+
model.zero_grad(set_to_none=True);del smoke_logits,smoke_pred,smoke_losses
|
| 190 |
+
print('PREFLIGHT PASSED: teacher batch 4, student BF16 backward, finite gradients, '+torch.cuda.get_device_name(),flush=True)
|
| 191 |
+
write('provenance.json',{'teacher_id':TEACHER_ID,'teacher_revision':teacher_rev,'teacher_parameters':sum(p.numel() for p in teacher.parameters()),'ddcolor_git':'2adb63f2656ac41cbdf7b894cddd94121a3faf13','base_revision':BASE_REV,'base_sha256':BASE_HASH,'steps_per_arm':STEPS,'seed':SEED,'teacher_input_size':512,'warning':'Existing development holdouts; teacher upstream training overlap unknown; no production certification.'})
|
| 192 |
+
probes(model,'v2');probes(teacher,'ddcolor_artistic',True)
|
| 193 |
+
DURABLE.sync('teacher preflight and visual baselines')
|
| 194 |
+
tables=[]
|
| 195 |
+
for repo,rev,files in [('johnowhitaker/imagenette2-320','771c1310a2487e8076ede6b7d6307244aa8400af',['default/train/0000.parquet']),('detection-datasets/coco','26ddc382fe75dfc2a0655b5977e296ea10efebce',['default/train/0000.parquet','default/train/0001.parquet'])]:
|
| 196 |
+
for file in files:tables.append(pq.read_table(hf_hub_download(repo,file,repo_type='dataset',revision=rev),columns=['image']))
|
| 197 |
+
table=pa.concat_tables(tables);manifest=json.loads(Path('previous_manifest.json').read_text())
|
| 198 |
+
ids=manifest['train'];val=manifest['validation'];test=manifest['test']
|
| 199 |
+
assert all(not set(a)&set(b) for a,b in [(ids,val),(ids,test),(val,test)])
|
| 200 |
+
write('manifest.json',manifest)
|
| 201 |
+
baseline=write('baseline_validation.json',evaluate(model,table,val))
|
| 202 |
+
write('teacher_validation.json',evaluate(teacher,table,val,True))
|
| 203 |
+
write('teacher_test.json',evaluate(teacher,table,test,True))
|
| 204 |
+
print('BASELINE',json.dumps({k:v['summary'] for k,v in baseline.items()}),flush=True)
|
| 205 |
+
DURABLE.sync('baseline and teacher evaluation')
|
| 206 |
+
cache=np.lib.format.open_memmap('teacher_cache.npy',mode='w+',dtype=np.float16,shape=(len(ids),2,64,64))
|
| 207 |
+
offset=0
|
| 208 |
+
for L,t,gray,indices in DataLoader(Photos(table,ids),batch_size=4,num_workers=4):
|
| 209 |
+
deadline();ab=F.avg_pool2d(teacher_predict(teacher,gray),4).cpu().numpy();cache[offset:offset+len(ab)]=ab;offset+=len(ab)
|
| 210 |
+
if offset%400==0:print('TEACHER_CACHE',offset,len(ids),'seconds',int(time.monotonic()-START),flush=True)
|
| 211 |
+
cache.flush();del teacher;torch.cuda.empty_cache()
|
| 212 |
+
bins=ColorBins(model.bin_centers.cpu().numpy());best_score=float('inf');history=[]
|
| 213 |
+
for arm,kd_weight in [('structure_control',0.),('structure_distilled',1.)]:
|
| 214 |
+
torch.manual_seed(SEED);np.random.seed(SEED);random.seed(SEED)
|
| 215 |
+
model=load_model(root,device=DEVICE)
|
| 216 |
+
optimizer=torch.optim.AdamW(model.parameters(),lr=3e-5,weight_decay=1e-4)
|
| 217 |
+
sched=torch.optim.lr_scheduler.CosineAnnealingLR(optimizer,STEPS,eta_min=3e-6)
|
| 218 |
+
loader=DataLoader(Photos(table,ids,bins,cache,True),batch_size=16,shuffle=True,num_workers=6,pin_memory=True,persistent_workers=True,generator=torch.Generator().manual_seed(SEED));it=iter(loader)
|
| 219 |
+
for step in range(1,STEPS+1):
|
| 220 |
+
deadline();model.train()
|
| 221 |
+
for m in model.modules():
|
| 222 |
+
if isinstance(m,torch.nn.BatchNorm2d):m.eval()
|
| 223 |
+
try:batch=next(it)
|
| 224 |
+
except StopIteration:it=iter(loader);batch=next(it)
|
| 225 |
+
L,t,idx,w,teach=[x.to(DEVICE,non_blocking=True) for x in batch]
|
| 226 |
+
optimizer.zero_grad(set_to_none=True)
|
| 227 |
+
with torch.autocast('cuda',dtype=torch.bfloat16):z=model(L)
|
| 228 |
+
p=model.decode(z,.38);structure,kd,neutral=additional_losses(p,t,teach)
|
| 229 |
+
loss=soft_ce(z,idx,w)+structure+kd_weight*kd+.3*neutral
|
| 230 |
+
if not torch.isfinite(loss):raise RuntimeError('Non-finite training loss')
|
| 231 |
+
loss.backward();torch.nn.utils.clip_grad_norm_(model.parameters(),1,error_if_nonfinite=True);optimizer.step();sched.step()
|
| 232 |
+
if step%100==0:print('STEP',arm,step,'loss',float(loss),'seconds',int(time.monotonic()-START),flush=True)
|
| 233 |
+
if step%250==0:
|
| 234 |
+
save_model(model,OUT/arm/'latest')
|
| 235 |
+
torch.save({'optimizer':optimizer.state_dict(),'scheduler':sched.state_dict(),'arm':arm,'step':step,'seed':SEED},OUT/arm/'latest'/'training_state.pt')
|
| 236 |
+
write('progress.json',{'arm':arm,'step':step,'elapsed_seconds':time.monotonic()-START})
|
| 237 |
+
DURABLE.sync(arm+' step '+str(step))
|
| 238 |
+
if step%500==0:
|
| 239 |
+
result=evaluate(model,table,val);score,ok=selection(result,baseline)
|
| 240 |
+
record={'arm':arm,'step':step,'eligible':ok,'score':score,'summary':{k:v['summary'] for k,v in result.items()}}
|
| 241 |
+
history.append(record);write('history.json',history);write(f'{arm}_{step}_validation.json',result);probes(model,f'{arm}_{step}')
|
| 242 |
+
if ok and score<best_score:
|
| 243 |
+
best_score=score;save_model(model,OUT/'best');write('selection.json',{'arm':arm,'step':step,'accepted':True,'score':score})
|
| 244 |
+
print('VALIDATION',json.dumps(record),flush=True)
|
| 245 |
+
DURABLE.sync(arm+' evaluation '+str(step))
|
| 246 |
+
del it,loader,model,optimizer;torch.cuda.empty_cache()
|
| 247 |
+
for tag,path in [('v2',root),('selected',OUT/'best')]:
|
| 248 |
+
model=load_model(path,device=DEVICE);write(tag+'_test.json',evaluate(model,table,test));probes(model,tag+'_final')
|
| 249 |
+
write('status.json',{'status':'completed','elapsed_seconds':time.monotonic()-START})
|
| 250 |
+
chosen=json.loads((OUT/'selection.json').read_text())
|
| 251 |
+
(OUT/'best'/'README.md').write_text('# Mini U-Net round5 candidate\n\nExperimental checkpoint; visual review required before deployment.\n\nSelection: '+json.dumps(chosen)+'\n\nSee ../source/PROTOCOL.md, ../provenance.json, ../history.json and ../selected_test.json. Trained on Imagenette and COCO; DDColor teacher uses data with possible evaluation overlap. Predictions are plausible colors, not recovered historical colors. Architecture and decoder remain compatible with the existing Space.\n')
|
| 252 |
+
DURABLE.sync('completed checkpoint, model card and evaluation')
|
| 253 |
+
|
| 254 |
+
if __name__=='__main__':
|
| 255 |
+
try:run()
|
| 256 |
+
except BaseException as e:
|
| 257 |
+
write('status.json',{'status':'interrupted_or_failed','error_type':type(e).__name__,'traceback':traceback.format_exc().replace(os.environ.get('HF_TOKEN','__NO_TOKEN__'),'[REDACTED]')})
|
| 258 |
+
if DURABLE is not None:DURABLE.sync('interrupted status and available checkpoints')
|
| 259 |
+
raise
|
experiments/round5-20260927/source/vendor_ddcolor/LICENSE
ADDED
|
@@ -0,0 +1,201 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Apache License
|
| 2 |
+
Version 2.0, January 2004
|
| 3 |
+
http://www.apache.org/licenses/
|
| 4 |
+
|
| 5 |
+
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
| 6 |
+
|
| 7 |
+
1. Definitions.
|
| 8 |
+
|
| 9 |
+
"License" shall mean the terms and conditions for use, reproduction,
|
| 10 |
+
and distribution as defined by Sections 1 through 9 of this document.
|
| 11 |
+
|
| 12 |
+
"Licensor" shall mean the copyright owner or entity authorized by
|
| 13 |
+
the copyright owner that is granting the License.
|
| 14 |
+
|
| 15 |
+
"Legal Entity" shall mean the union of the acting entity and all
|
| 16 |
+
other entities that control, are controlled by, or are under common
|
| 17 |
+
control with that entity. For the purposes of this definition,
|
| 18 |
+
"control" means (i) the power, direct or indirect, to cause the
|
| 19 |
+
direction or management of such entity, whether by contract or
|
| 20 |
+
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
| 21 |
+
outstanding shares, or (iii) beneficial ownership of such entity.
|
| 22 |
+
|
| 23 |
+
"You" (or "Your") shall mean an individual or Legal Entity
|
| 24 |
+
exercising permissions granted by this License.
|
| 25 |
+
|
| 26 |
+
"Source" form shall mean the preferred form for making modifications,
|
| 27 |
+
including but not limited to software source code, documentation
|
| 28 |
+
source, and configuration files.
|
| 29 |
+
|
| 30 |
+
"Object" form shall mean any form resulting from mechanical
|
| 31 |
+
transformation or translation of a Source form, including but
|
| 32 |
+
not limited to compiled object code, generated documentation,
|
| 33 |
+
and conversions to other media types.
|
| 34 |
+
|
| 35 |
+
"Work" shall mean the work of authorship, whether in Source or
|
| 36 |
+
Object form, made available under the License, as indicated by a
|
| 37 |
+
copyright notice that is included in or attached to the work
|
| 38 |
+
(an example is provided in the Appendix below).
|
| 39 |
+
|
| 40 |
+
"Derivative Works" shall mean any work, whether in Source or Object
|
| 41 |
+
form, that is based on (or derived from) the Work and for which the
|
| 42 |
+
editorial revisions, annotations, elaborations, or other modifications
|
| 43 |
+
represent, as a whole, an original work of authorship. For the purposes
|
| 44 |
+
of this License, Derivative Works shall not include works that remain
|
| 45 |
+
separable from, or merely link (or bind by name) to the interfaces of,
|
| 46 |
+
the Work and Derivative Works thereof.
|
| 47 |
+
|
| 48 |
+
"Contribution" shall mean any work of authorship, including
|
| 49 |
+
the original version of the Work and any modifications or additions
|
| 50 |
+
to that Work or Derivative Works thereof, that is intentionally
|
| 51 |
+
submitted to Licensor for inclusion in the Work by the copyright owner
|
| 52 |
+
or by an individual or Legal Entity authorized to submit on behalf of
|
| 53 |
+
the copyright owner. For the purposes of this definition, "submitted"
|
| 54 |
+
means any form of electronic, verbal, or written communication sent
|
| 55 |
+
to the Licensor or its representatives, including but not limited to
|
| 56 |
+
communication on electronic mailing lists, source code control systems,
|
| 57 |
+
and issue tracking systems that are managed by, or on behalf of, the
|
| 58 |
+
Licensor for the purpose of discussing and improving the Work, but
|
| 59 |
+
excluding communication that is conspicuously marked or otherwise
|
| 60 |
+
designated in writing by the copyright owner as "Not a Contribution."
|
| 61 |
+
|
| 62 |
+
"Contributor" shall mean Licensor and any individual or Legal Entity
|
| 63 |
+
on behalf of whom a Contribution has been received by Licensor and
|
| 64 |
+
subsequently incorporated within the Work.
|
| 65 |
+
|
| 66 |
+
2. Grant of Copyright License. Subject to the terms and conditions of
|
| 67 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 68 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 69 |
+
copyright license to reproduce, prepare Derivative Works of,
|
| 70 |
+
publicly display, publicly perform, sublicense, and distribute the
|
| 71 |
+
Work and such Derivative Works in Source or Object form.
|
| 72 |
+
|
| 73 |
+
3. Grant of Patent License. Subject to the terms and conditions of
|
| 74 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 75 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 76 |
+
(except as stated in this section) patent license to make, have made,
|
| 77 |
+
use, offer to sell, sell, import, and otherwise transfer the Work,
|
| 78 |
+
where such license applies only to those patent claims licensable
|
| 79 |
+
by such Contributor that are necessarily infringed by their
|
| 80 |
+
Contribution(s) alone or by combination of their Contribution(s)
|
| 81 |
+
with the Work to which such Contribution(s) was submitted. If You
|
| 82 |
+
institute patent litigation against any entity (including a
|
| 83 |
+
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
| 84 |
+
or a Contribution incorporated within the Work constitutes direct
|
| 85 |
+
or contributory patent infringement, then any patent licenses
|
| 86 |
+
granted to You under this License for that Work shall terminate
|
| 87 |
+
as of the date such litigation is filed.
|
| 88 |
+
|
| 89 |
+
4. Redistribution. You may reproduce and distribute copies of the
|
| 90 |
+
Work or Derivative Works thereof in any medium, with or without
|
| 91 |
+
modifications, and in Source or Object form, provided that You
|
| 92 |
+
meet the following conditions:
|
| 93 |
+
|
| 94 |
+
(a) You must give any other recipients of the Work or
|
| 95 |
+
Derivative Works a copy of this License; and
|
| 96 |
+
|
| 97 |
+
(b) You must cause any modified files to carry prominent notices
|
| 98 |
+
stating that You changed the files; and
|
| 99 |
+
|
| 100 |
+
(c) You must retain, in the Source form of any Derivative Works
|
| 101 |
+
that You distribute, all copyright, patent, trademark, and
|
| 102 |
+
attribution notices from the Source form of the Work,
|
| 103 |
+
excluding those notices that do not pertain to any part of
|
| 104 |
+
the Derivative Works; and
|
| 105 |
+
|
| 106 |
+
(d) If the Work includes a "NOTICE" text file as part of its
|
| 107 |
+
distribution, then any Derivative Works that You distribute must
|
| 108 |
+
include a readable copy of the attribution notices contained
|
| 109 |
+
within such NOTICE file, excluding those notices that do not
|
| 110 |
+
pertain to any part of the Derivative Works, in at least one
|
| 111 |
+
of the following places: within a NOTICE text file distributed
|
| 112 |
+
as part of the Derivative Works; within the Source form or
|
| 113 |
+
documentation, if provided along with the Derivative Works; or,
|
| 114 |
+
within a display generated by the Derivative Works, if and
|
| 115 |
+
wherever such third-party notices normally appear. The contents
|
| 116 |
+
of the NOTICE file are for informational purposes only and
|
| 117 |
+
do not modify the License. You may add Your own attribution
|
| 118 |
+
notices within Derivative Works that You distribute, alongside
|
| 119 |
+
or as an addendum to the NOTICE text from the Work, provided
|
| 120 |
+
that such additional attribution notices cannot be construed
|
| 121 |
+
as modifying the License.
|
| 122 |
+
|
| 123 |
+
You may add Your own copyright statement to Your modifications and
|
| 124 |
+
may provide additional or different license terms and conditions
|
| 125 |
+
for use, reproduction, or distribution of Your modifications, or
|
| 126 |
+
for any such Derivative Works as a whole, provided Your use,
|
| 127 |
+
reproduction, and distribution of the Work otherwise complies with
|
| 128 |
+
the conditions stated in this License.
|
| 129 |
+
|
| 130 |
+
5. Submission of Contributions. Unless You explicitly state otherwise,
|
| 131 |
+
any Contribution intentionally submitted for inclusion in the Work
|
| 132 |
+
by You to the Licensor shall be under the terms and conditions of
|
| 133 |
+
this License, without any additional terms or conditions.
|
| 134 |
+
Notwithstanding the above, nothing herein shall supersede or modify
|
| 135 |
+
the terms of any separate license agreement you may have executed
|
| 136 |
+
with Licensor regarding such Contributions.
|
| 137 |
+
|
| 138 |
+
6. Trademarks. This License does not grant permission to use the trade
|
| 139 |
+
names, trademarks, service marks, or product names of the Licensor,
|
| 140 |
+
except as required for reasonable and customary use in describing the
|
| 141 |
+
origin of the Work and reproducing the content of the NOTICE file.
|
| 142 |
+
|
| 143 |
+
7. Disclaimer of Warranty. Unless required by applicable law or
|
| 144 |
+
agreed to in writing, Licensor provides the Work (and each
|
| 145 |
+
Contributor provides its Contributions) on an "AS IS" BASIS,
|
| 146 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
| 147 |
+
implied, including, without limitation, any warranties or conditions
|
| 148 |
+
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
| 149 |
+
PARTICULAR PURPOSE. You are solely responsible for determining the
|
| 150 |
+
appropriateness of using or redistributing the Work and assume any
|
| 151 |
+
risks associated with Your exercise of permissions under this License.
|
| 152 |
+
|
| 153 |
+
8. Limitation of Liability. In no event and under no legal theory,
|
| 154 |
+
whether in tort (including negligence), contract, or otherwise,
|
| 155 |
+
unless required by applicable law (such as deliberate and grossly
|
| 156 |
+
negligent acts) or agreed to in writing, shall any Contributor be
|
| 157 |
+
liable to You for damages, including any direct, indirect, special,
|
| 158 |
+
incidental, or consequential damages of any character arising as a
|
| 159 |
+
result of this License or out of the use or inability to use the
|
| 160 |
+
Work (including but not limited to damages for loss of goodwill,
|
| 161 |
+
work stoppage, computer failure or malfunction, or any and all
|
| 162 |
+
other commercial damages or losses), even if such Contributor
|
| 163 |
+
has been advised of the possibility of such damages.
|
| 164 |
+
|
| 165 |
+
9. Accepting Warranty or Additional Liability. While redistributing
|
| 166 |
+
the Work or Derivative Works thereof, You may choose to offer,
|
| 167 |
+
and charge a fee for, acceptance of support, warranty, indemnity,
|
| 168 |
+
or other liability obligations and/or rights consistent with this
|
| 169 |
+
License. However, in accepting such obligations, You may act only
|
| 170 |
+
on Your own behalf and on Your sole responsibility, not on behalf
|
| 171 |
+
of any other Contributor, and only if You agree to indemnify,
|
| 172 |
+
defend, and hold each Contributor harmless for any liability
|
| 173 |
+
incurred by, or claims asserted against, such Contributor by reason
|
| 174 |
+
of your accepting any such warranty or additional liability.
|
| 175 |
+
|
| 176 |
+
END OF TERMS AND CONDITIONS
|
| 177 |
+
|
| 178 |
+
APPENDIX: How to apply the Apache License to your work.
|
| 179 |
+
|
| 180 |
+
To apply the Apache License to your work, attach the following
|
| 181 |
+
boilerplate notice, with the fields enclosed by brackets "[]"
|
| 182 |
+
replaced with your own identifying information. (Don't include
|
| 183 |
+
the brackets!) The text should be enclosed in the appropriate
|
| 184 |
+
comment syntax for the file format. We also recommend that a
|
| 185 |
+
file or class name and description of purpose be included on the
|
| 186 |
+
same "printed page" as the copyright notice for easier
|
| 187 |
+
identification within third-party archives.
|
| 188 |
+
|
| 189 |
+
Copyright [yyyy] [name of copyright owner]
|
| 190 |
+
|
| 191 |
+
Licensed under the Apache License, Version 2.0 (the "License");
|
| 192 |
+
you may not use this file except in compliance with the License.
|
| 193 |
+
You may obtain a copy of the License at
|
| 194 |
+
|
| 195 |
+
http://www.apache.org/licenses/LICENSE-2.0
|
| 196 |
+
|
| 197 |
+
Unless required by applicable law or agreed to in writing, software
|
| 198 |
+
distributed under the License is distributed on an "AS IS" BASIS,
|
| 199 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 200 |
+
See the License for the specific language governing permissions and
|
| 201 |
+
limitations under the License.
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/__init__.py
ADDED
|
@@ -0,0 +1,16 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""BasicSR (vendored)
|
| 2 |
+
|
| 3 |
+
This repo's inference scripts only need a small subset under `basicsr.archs...`.
|
| 4 |
+
Upstream BasicSR's `basicsr/__init__.py` often does `import *` from archs/data/losses/metrics/models/train/utils,
|
| 5 |
+
which pulls in many training-only dependencies during inference import.
|
| 6 |
+
|
| 7 |
+
We keep this `__init__` lightweight to avoid import-time side effects.
|
| 8 |
+
Training code should explicitly import the needed submodules.
|
| 9 |
+
"""
|
| 10 |
+
|
| 11 |
+
# flake8: noqa
|
| 12 |
+
try:
|
| 13 |
+
from .version import __gitsha__, __version__ # type: ignore
|
| 14 |
+
except Exception:
|
| 15 |
+
__gitsha__ = None
|
| 16 |
+
__version__ = None
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/__init__.py
ADDED
|
@@ -0,0 +1,41 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import importlib
|
| 2 |
+
import logging
|
| 3 |
+
import os
|
| 4 |
+
from copy import deepcopy
|
| 5 |
+
from os import path as osp
|
| 6 |
+
|
| 7 |
+
from basicsr.utils.registry import ARCH_REGISTRY
|
| 8 |
+
|
| 9 |
+
__all__ = ['build_network']
|
| 10 |
+
|
| 11 |
+
_ARCHS_IMPORTED = False
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def _ensure_arch_modules_imported():
|
| 15 |
+
"""Lazy import arch modules for registry.
|
| 16 |
+
|
| 17 |
+
In upstream BasicSR, importing `basicsr.archs` scans and imports all `*_arch.py`
|
| 18 |
+
modules eagerly to populate the registry. That adds import overhead and may
|
| 19 |
+
pull in extra dependencies in inference-only scenarios.
|
| 20 |
+
Here we make it lazy: only scan/import when `build_network` is actually called.
|
| 21 |
+
"""
|
| 22 |
+
global _ARCHS_IMPORTED
|
| 23 |
+
if _ARCHS_IMPORTED:
|
| 24 |
+
return
|
| 25 |
+
arch_folder = osp.dirname(osp.abspath(__file__))
|
| 26 |
+
arch_filenames = []
|
| 27 |
+
for name in os.listdir(arch_folder):
|
| 28 |
+
if name.endswith("_arch.py"):
|
| 29 |
+
arch_filenames.append(osp.splitext(name)[0])
|
| 30 |
+
for file_name in arch_filenames:
|
| 31 |
+
importlib.import_module(f'basicsr.archs.{file_name}')
|
| 32 |
+
_ARCHS_IMPORTED = True
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def build_network(opt):
|
| 36 |
+
_ensure_arch_modules_imported()
|
| 37 |
+
opt = deepcopy(opt)
|
| 38 |
+
network_type = opt.pop('type')
|
| 39 |
+
net = ARCH_REGISTRY.get(network_type)(**opt)
|
| 40 |
+
logging.getLogger('basicsr').info(f'Network [{net.__class__.__name__}] is created.')
|
| 41 |
+
return net
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/__init__.py
ADDED
|
File without changes
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/convnext.py
ADDED
|
@@ -0,0 +1,206 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright (c) Meta Platforms, Inc. and affiliates.
|
| 2 |
+
|
| 3 |
+
# All rights reserved.
|
| 4 |
+
|
| 5 |
+
# This source code is licensed under the license found in the
|
| 6 |
+
# LICENSE file in the root directory of this source tree.
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
import torch
|
| 10 |
+
import torch.nn as nn
|
| 11 |
+
import torch.nn.functional as F
|
| 12 |
+
|
| 13 |
+
# ---- Optional dependency: timm ----
|
| 14 |
+
# This file only needs two small helpers from timm: `trunc_normal_` and `DropPath`.
|
| 15 |
+
# To reduce inference dependencies, we provide a pure-PyTorch fallback implementation.
|
| 16 |
+
try:
|
| 17 |
+
from timm.layers import trunc_normal_, DropPath # type: ignore
|
| 18 |
+
except Exception:
|
| 19 |
+
import math
|
| 20 |
+
|
| 21 |
+
def trunc_normal_(tensor, mean=0.0, std=1.0, a=-2.0, b=2.0):
|
| 22 |
+
"""Fills the input Tensor with values drawn from a truncated normal distribution.
|
| 23 |
+
|
| 24 |
+
Fallback implementation when timm is not available.
|
| 25 |
+
"""
|
| 26 |
+
# Prefer PyTorch built-in if present
|
| 27 |
+
if hasattr(torch.nn.init, "trunc_normal_"):
|
| 28 |
+
return torch.nn.init.trunc_normal_(tensor, mean=mean, std=std, a=a, b=b)
|
| 29 |
+
|
| 30 |
+
# Based on PyTorch's internal implementation pattern
|
| 31 |
+
def norm_cdf(x):
|
| 32 |
+
return (1.0 + math.erf(x / math.sqrt(2.0))) / 2.0
|
| 33 |
+
|
| 34 |
+
with torch.no_grad():
|
| 35 |
+
l = norm_cdf((a - mean) / std)
|
| 36 |
+
u = norm_cdf((b - mean) / std)
|
| 37 |
+
|
| 38 |
+
tensor.uniform_(2 * l - 1, 2 * u - 1)
|
| 39 |
+
tensor.erfinv_()
|
| 40 |
+
|
| 41 |
+
tensor.mul_(std * math.sqrt(2.0))
|
| 42 |
+
tensor.add_(mean)
|
| 43 |
+
tensor.clamp_(min=a, max=b)
|
| 44 |
+
return tensor
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
class DropPath(nn.Module):
|
| 48 |
+
"""Stochastic Depth per sample (when applied in main path of residual blocks)."""
|
| 49 |
+
|
| 50 |
+
def __init__(self, drop_prob: float = 0.0):
|
| 51 |
+
super().__init__()
|
| 52 |
+
self.drop_prob = float(drop_prob)
|
| 53 |
+
|
| 54 |
+
def forward(self, x):
|
| 55 |
+
if self.drop_prob == 0.0 or not self.training:
|
| 56 |
+
return x
|
| 57 |
+
keep_prob = 1.0 - self.drop_prob
|
| 58 |
+
shape = (x.shape[0],) + (1,) * (x.ndim - 1)
|
| 59 |
+
random_tensor = keep_prob + torch.rand(
|
| 60 |
+
shape, dtype=x.dtype, device=x.device
|
| 61 |
+
)
|
| 62 |
+
random_tensor.floor_()
|
| 63 |
+
return x.div(keep_prob) * random_tensor
|
| 64 |
+
|
| 65 |
+
class Block(nn.Module):
|
| 66 |
+
r""" ConvNeXt Block. There are two equivalent implementations:
|
| 67 |
+
(1) DwConv -> LayerNorm (channels_first) -> 1x1 Conv -> GELU -> 1x1 Conv; all in (N, C, H, W)
|
| 68 |
+
(2) DwConv -> Permute to (N, H, W, C); LayerNorm (channels_last) -> Linear -> GELU -> Linear; Permute back
|
| 69 |
+
We use (2) as we find it slightly faster in PyTorch
|
| 70 |
+
|
| 71 |
+
Args:
|
| 72 |
+
dim (int): Number of input channels.
|
| 73 |
+
drop_path (float): Stochastic depth rate. Default: 0.0
|
| 74 |
+
layer_scale_init_value (float): Init value for Layer Scale. Default: 1e-6.
|
| 75 |
+
"""
|
| 76 |
+
def __init__(self, dim, drop_path=0., layer_scale_init_value=1e-6):
|
| 77 |
+
super().__init__()
|
| 78 |
+
self.dwconv = nn.Conv2d(dim, dim, kernel_size=7, padding=3, groups=dim) # depthwise conv
|
| 79 |
+
self.norm = LayerNorm(dim, eps=1e-6)
|
| 80 |
+
self.pwconv1 = nn.Linear(dim, 4 * dim) # pointwise/1x1 convs, implemented with linear layers
|
| 81 |
+
self.act = nn.GELU()
|
| 82 |
+
self.pwconv2 = nn.Linear(4 * dim, dim)
|
| 83 |
+
self.gamma = nn.Parameter(layer_scale_init_value * torch.ones((dim)),
|
| 84 |
+
requires_grad=True) if layer_scale_init_value > 0 else None
|
| 85 |
+
self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
|
| 86 |
+
|
| 87 |
+
def forward(self, x):
|
| 88 |
+
input = x
|
| 89 |
+
x = self.dwconv(x)
|
| 90 |
+
x = x.permute(0, 2, 3, 1) # (N, C, H, W) -> (N, H, W, C)
|
| 91 |
+
x = self.norm(x)
|
| 92 |
+
x = self.pwconv1(x)
|
| 93 |
+
x = self.act(x)
|
| 94 |
+
x = self.pwconv2(x)
|
| 95 |
+
if self.gamma is not None:
|
| 96 |
+
x = self.gamma * x
|
| 97 |
+
x = x.permute(0, 3, 1, 2) # (N, H, W, C) -> (N, C, H, W)
|
| 98 |
+
|
| 99 |
+
x = input + self.drop_path(x)
|
| 100 |
+
return x
|
| 101 |
+
|
| 102 |
+
class ConvNeXt(nn.Module):
|
| 103 |
+
r""" ConvNeXt
|
| 104 |
+
A PyTorch impl of : `A ConvNet for the 2020s` -
|
| 105 |
+
https://arxiv.org/pdf/2201.03545.pdf
|
| 106 |
+
Args:
|
| 107 |
+
in_chans (int): Number of input image channels. Default: 3
|
| 108 |
+
num_classes (int): Number of classes for classification head. Default: 1000
|
| 109 |
+
depths (tuple(int)): Number of blocks at each stage. Default: [3, 3, 9, 3]
|
| 110 |
+
dims (int): Feature dimension at each stage. Default: [96, 192, 384, 768]
|
| 111 |
+
drop_path_rate (float): Stochastic depth rate. Default: 0.
|
| 112 |
+
layer_scale_init_value (float): Init value for Layer Scale. Default: 1e-6.
|
| 113 |
+
head_init_scale (float): Init scaling value for classifier weights and biases. Default: 1.
|
| 114 |
+
"""
|
| 115 |
+
def __init__(self, in_chans=3, num_classes=1000,
|
| 116 |
+
depths=[3, 3, 9, 3], dims=[96, 192, 384, 768], drop_path_rate=0.,
|
| 117 |
+
layer_scale_init_value=1e-6, head_init_scale=1.,
|
| 118 |
+
):
|
| 119 |
+
super().__init__()
|
| 120 |
+
|
| 121 |
+
self.downsample_layers = nn.ModuleList() # stem and 3 intermediate downsampling conv layers
|
| 122 |
+
stem = nn.Sequential(
|
| 123 |
+
nn.Conv2d(in_chans, dims[0], kernel_size=4, stride=4),
|
| 124 |
+
LayerNorm(dims[0], eps=1e-6, data_format="channels_first")
|
| 125 |
+
)
|
| 126 |
+
self.downsample_layers.append(stem)
|
| 127 |
+
for i in range(3):
|
| 128 |
+
downsample_layer = nn.Sequential(
|
| 129 |
+
LayerNorm(dims[i], eps=1e-6, data_format="channels_first"),
|
| 130 |
+
nn.Conv2d(dims[i], dims[i+1], kernel_size=2, stride=2),
|
| 131 |
+
)
|
| 132 |
+
self.downsample_layers.append(downsample_layer)
|
| 133 |
+
|
| 134 |
+
self.stages = nn.ModuleList() # 4 feature resolution stages, each consisting of multiple residual blocks
|
| 135 |
+
dp_rates=[x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))]
|
| 136 |
+
cur = 0
|
| 137 |
+
for i in range(4):
|
| 138 |
+
stage = nn.Sequential(
|
| 139 |
+
*[Block(dim=dims[i], drop_path=dp_rates[cur + j],
|
| 140 |
+
layer_scale_init_value=layer_scale_init_value) for j in range(depths[i])]
|
| 141 |
+
)
|
| 142 |
+
self.stages.append(stage)
|
| 143 |
+
cur += depths[i]
|
| 144 |
+
|
| 145 |
+
# add norm layers for each output
|
| 146 |
+
out_indices = (0, 1, 2, 3)
|
| 147 |
+
for i in out_indices:
|
| 148 |
+
layer = LayerNorm(dims[i], eps=1e-6, data_format="channels_first")
|
| 149 |
+
# layer = nn.Identity()
|
| 150 |
+
layer_name = f'norm{i}'
|
| 151 |
+
self.add_module(layer_name, layer)
|
| 152 |
+
|
| 153 |
+
self.norm = nn.LayerNorm(dims[-1], eps=1e-6) # final norm layer
|
| 154 |
+
# self.head_cls = nn.Linear(dims[-1], 4)
|
| 155 |
+
|
| 156 |
+
self.apply(self._init_weights)
|
| 157 |
+
# self.head_cls.weight.data.mul_(head_init_scale)
|
| 158 |
+
# self.head_cls.bias.data.mul_(head_init_scale)
|
| 159 |
+
|
| 160 |
+
def _init_weights(self, m):
|
| 161 |
+
if isinstance(m, (nn.Conv2d, nn.Linear)):
|
| 162 |
+
trunc_normal_(m.weight, std=.02)
|
| 163 |
+
nn.init.constant_(m.bias, 0)
|
| 164 |
+
|
| 165 |
+
def forward_features(self, x):
|
| 166 |
+
for i in range(4):
|
| 167 |
+
x = self.downsample_layers[i](x)
|
| 168 |
+
x = self.stages[i](x)
|
| 169 |
+
|
| 170 |
+
# add extra norm
|
| 171 |
+
norm_layer = getattr(self, f'norm{i}')
|
| 172 |
+
# x = norm_layer(x)
|
| 173 |
+
norm_layer(x)
|
| 174 |
+
|
| 175 |
+
return self.norm(x.mean([-2, -1])) # global average pooling, (N, C, H, W) -> (N, C)
|
| 176 |
+
|
| 177 |
+
def forward(self, x):
|
| 178 |
+
x = self.forward_features(x)
|
| 179 |
+
# x = self.head_cls(x)
|
| 180 |
+
return x
|
| 181 |
+
|
| 182 |
+
class LayerNorm(nn.Module):
|
| 183 |
+
r""" LayerNorm that supports two data formats: channels_last (default) or channels_first.
|
| 184 |
+
The ordering of the dimensions in the inputs. channels_last corresponds to inputs with
|
| 185 |
+
shape (batch_size, height, width, channels) while channels_first corresponds to inputs
|
| 186 |
+
with shape (batch_size, channels, height, width).
|
| 187 |
+
"""
|
| 188 |
+
def __init__(self, normalized_shape, eps=1e-6, data_format="channels_last"):
|
| 189 |
+
super().__init__()
|
| 190 |
+
self.weight = nn.Parameter(torch.ones(normalized_shape))
|
| 191 |
+
self.bias = nn.Parameter(torch.zeros(normalized_shape))
|
| 192 |
+
self.eps = eps
|
| 193 |
+
self.data_format = data_format
|
| 194 |
+
if self.data_format not in ["channels_last", "channels_first"]:
|
| 195 |
+
raise NotImplementedError
|
| 196 |
+
self.normalized_shape = (normalized_shape, )
|
| 197 |
+
|
| 198 |
+
def forward(self, x):
|
| 199 |
+
if self.data_format == "channels_last": # B H W C
|
| 200 |
+
return F.layer_norm(x, self.normalized_shape, self.weight, self.bias, self.eps)
|
| 201 |
+
elif self.data_format == "channels_first": # B C H W
|
| 202 |
+
u = x.mean(1, keepdim=True)
|
| 203 |
+
s = (x - u).pow(2).mean(1, keepdim=True)
|
| 204 |
+
x = (x - u) / torch.sqrt(s + self.eps)
|
| 205 |
+
x = self.weight[:, None, None] * x + self.bias[:, None, None]
|
| 206 |
+
return x
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/position_encoding.py
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright (c) Facebook, Inc. and its affiliates.
|
| 2 |
+
# Modified from: https://github.com/facebookresearch/detr/blob/master/models/position_encoding.py
|
| 3 |
+
"""
|
| 4 |
+
Various positional encodings for the transformer.
|
| 5 |
+
"""
|
| 6 |
+
import math
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
from torch import nn
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class PositionEmbeddingSine(nn.Module):
|
| 13 |
+
"""
|
| 14 |
+
This is a more standard version of the position embedding, very similar to the one
|
| 15 |
+
used by the Attention is all you need paper, generalized to work on images.
|
| 16 |
+
"""
|
| 17 |
+
|
| 18 |
+
def __init__(self, num_pos_feats=64, temperature=10000, normalize=False, scale=None):
|
| 19 |
+
super().__init__()
|
| 20 |
+
self.num_pos_feats = num_pos_feats
|
| 21 |
+
self.temperature = temperature
|
| 22 |
+
self.normalize = normalize
|
| 23 |
+
if scale is not None and normalize is False:
|
| 24 |
+
raise ValueError("normalize should be True if scale is passed")
|
| 25 |
+
if scale is None:
|
| 26 |
+
scale = 2 * math.pi
|
| 27 |
+
self.scale = scale
|
| 28 |
+
|
| 29 |
+
def forward(self, x, mask=None):
|
| 30 |
+
if mask is None:
|
| 31 |
+
mask = torch.zeros((x.size(0), x.size(2), x.size(3)), device=x.device, dtype=torch.bool)
|
| 32 |
+
not_mask = ~mask
|
| 33 |
+
y_embed = not_mask.cumsum(1, dtype=torch.float32)
|
| 34 |
+
x_embed = not_mask.cumsum(2, dtype=torch.float32)
|
| 35 |
+
if self.normalize:
|
| 36 |
+
eps = 1e-6
|
| 37 |
+
y_embed = y_embed / (y_embed[:, -1:, :] + eps) * self.scale
|
| 38 |
+
x_embed = x_embed / (x_embed[:, :, -1:] + eps) * self.scale
|
| 39 |
+
|
| 40 |
+
dim_t = torch.arange(self.num_pos_feats, dtype=torch.float32, device=x.device)
|
| 41 |
+
dim_t = self.temperature ** (2 * (dim_t // 2) / self.num_pos_feats)
|
| 42 |
+
|
| 43 |
+
pos_x = x_embed[:, :, :, None] / dim_t
|
| 44 |
+
pos_y = y_embed[:, :, :, None] / dim_t
|
| 45 |
+
pos_x = torch.stack(
|
| 46 |
+
(pos_x[:, :, :, 0::2].sin(), pos_x[:, :, :, 1::2].cos()), dim=4
|
| 47 |
+
).flatten(3)
|
| 48 |
+
pos_y = torch.stack(
|
| 49 |
+
(pos_y[:, :, :, 0::2].sin(), pos_y[:, :, :, 1::2].cos()), dim=4
|
| 50 |
+
).flatten(3)
|
| 51 |
+
pos = torch.cat((pos_y, pos_x), dim=3).permute(0, 3, 1, 2)
|
| 52 |
+
return pos
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/transformer.py
ADDED
|
@@ -0,0 +1,368 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright (c) Facebook, Inc. and its affiliates.
|
| 2 |
+
# Modified from: https://github.com/facebookresearch/detr/blob/master/models/transformer.py
|
| 3 |
+
"""
|
| 4 |
+
Transformer class.
|
| 5 |
+
Copy-paste from torch.nn.Transformer with modifications:
|
| 6 |
+
* positional encodings are passed in MHattention
|
| 7 |
+
* extra LN at the end of encoder is removed
|
| 8 |
+
* decoder returns a stack of activations from all decoding layers
|
| 9 |
+
"""
|
| 10 |
+
import copy
|
| 11 |
+
from typing import List, Optional
|
| 12 |
+
|
| 13 |
+
import torch
|
| 14 |
+
import torch.nn.functional as F
|
| 15 |
+
from torch import Tensor, nn
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
class Transformer(nn.Module):
|
| 19 |
+
def __init__(
|
| 20 |
+
self,
|
| 21 |
+
d_model=512,
|
| 22 |
+
nhead=8,
|
| 23 |
+
num_encoder_layers=6,
|
| 24 |
+
num_decoder_layers=6,
|
| 25 |
+
dim_feedforward=2048,
|
| 26 |
+
dropout=0.1,
|
| 27 |
+
activation="relu",
|
| 28 |
+
normalize_before=False,
|
| 29 |
+
return_intermediate_dec=False,
|
| 30 |
+
):
|
| 31 |
+
super().__init__()
|
| 32 |
+
|
| 33 |
+
encoder_layer = TransformerEncoderLayer(
|
| 34 |
+
d_model, nhead, dim_feedforward, dropout, activation, normalize_before
|
| 35 |
+
)
|
| 36 |
+
encoder_norm = nn.LayerNorm(d_model) if normalize_before else None
|
| 37 |
+
self.encoder = TransformerEncoder(encoder_layer, num_encoder_layers, encoder_norm)
|
| 38 |
+
|
| 39 |
+
decoder_layer = TransformerDecoderLayer(
|
| 40 |
+
d_model, nhead, dim_feedforward, dropout, activation, normalize_before
|
| 41 |
+
)
|
| 42 |
+
decoder_norm = nn.LayerNorm(d_model)
|
| 43 |
+
self.decoder = TransformerDecoder(
|
| 44 |
+
decoder_layer,
|
| 45 |
+
num_decoder_layers,
|
| 46 |
+
decoder_norm,
|
| 47 |
+
return_intermediate=return_intermediate_dec,
|
| 48 |
+
)
|
| 49 |
+
|
| 50 |
+
self._reset_parameters()
|
| 51 |
+
|
| 52 |
+
self.d_model = d_model
|
| 53 |
+
self.nhead = nhead
|
| 54 |
+
|
| 55 |
+
def _reset_parameters(self):
|
| 56 |
+
for p in self.parameters():
|
| 57 |
+
if p.dim() > 1:
|
| 58 |
+
nn.init.xavier_uniform_(p)
|
| 59 |
+
|
| 60 |
+
def forward(self, src, mask, query_embed, pos_embed):
|
| 61 |
+
# flatten NxCxHxW to HWxNxC
|
| 62 |
+
bs, c, h, w = src.shape
|
| 63 |
+
src = src.flatten(2).permute(2, 0, 1)
|
| 64 |
+
pos_embed = pos_embed.flatten(2).permute(2, 0, 1)
|
| 65 |
+
query_embed = query_embed.unsqueeze(1).repeat(1, bs, 1)
|
| 66 |
+
if mask is not None:
|
| 67 |
+
mask = mask.flatten(1)
|
| 68 |
+
|
| 69 |
+
tgt = torch.zeros_like(query_embed)
|
| 70 |
+
memory = self.encoder(src, src_key_padding_mask=mask, pos=pos_embed)
|
| 71 |
+
hs = self.decoder(
|
| 72 |
+
tgt, memory, memory_key_padding_mask=mask, pos=pos_embed, query_pos=query_embed
|
| 73 |
+
)
|
| 74 |
+
return hs.transpose(1, 2), memory.permute(1, 2, 0).view(bs, c, h, w)
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
class TransformerEncoder(nn.Module):
|
| 78 |
+
def __init__(self, encoder_layer, num_layers, norm=None):
|
| 79 |
+
super().__init__()
|
| 80 |
+
self.layers = _get_clones(encoder_layer, num_layers)
|
| 81 |
+
self.num_layers = num_layers
|
| 82 |
+
self.norm = norm
|
| 83 |
+
|
| 84 |
+
def forward(
|
| 85 |
+
self,
|
| 86 |
+
src,
|
| 87 |
+
mask: Optional[Tensor] = None,
|
| 88 |
+
src_key_padding_mask: Optional[Tensor] = None,
|
| 89 |
+
pos: Optional[Tensor] = None,
|
| 90 |
+
):
|
| 91 |
+
output = src
|
| 92 |
+
|
| 93 |
+
for layer in self.layers:
|
| 94 |
+
output = layer(
|
| 95 |
+
output, src_mask=mask, src_key_padding_mask=src_key_padding_mask, pos=pos
|
| 96 |
+
)
|
| 97 |
+
|
| 98 |
+
if self.norm is not None:
|
| 99 |
+
output = self.norm(output)
|
| 100 |
+
|
| 101 |
+
return output
|
| 102 |
+
|
| 103 |
+
|
| 104 |
+
class TransformerDecoder(nn.Module):
|
| 105 |
+
def __init__(self, decoder_layer, num_layers, norm=None, return_intermediate=False):
|
| 106 |
+
super().__init__()
|
| 107 |
+
self.layers = _get_clones(decoder_layer, num_layers)
|
| 108 |
+
self.num_layers = num_layers
|
| 109 |
+
self.norm = norm
|
| 110 |
+
self.return_intermediate = return_intermediate
|
| 111 |
+
|
| 112 |
+
def forward(
|
| 113 |
+
self,
|
| 114 |
+
tgt,
|
| 115 |
+
memory,
|
| 116 |
+
tgt_mask: Optional[Tensor] = None,
|
| 117 |
+
memory_mask: Optional[Tensor] = None,
|
| 118 |
+
tgt_key_padding_mask: Optional[Tensor] = None,
|
| 119 |
+
memory_key_padding_mask: Optional[Tensor] = None,
|
| 120 |
+
pos: Optional[Tensor] = None,
|
| 121 |
+
query_pos: Optional[Tensor] = None,
|
| 122 |
+
):
|
| 123 |
+
output = tgt
|
| 124 |
+
|
| 125 |
+
intermediate = []
|
| 126 |
+
|
| 127 |
+
for layer in self.layers:
|
| 128 |
+
output = layer(
|
| 129 |
+
output,
|
| 130 |
+
memory,
|
| 131 |
+
tgt_mask=tgt_mask,
|
| 132 |
+
memory_mask=memory_mask,
|
| 133 |
+
tgt_key_padding_mask=tgt_key_padding_mask,
|
| 134 |
+
memory_key_padding_mask=memory_key_padding_mask,
|
| 135 |
+
pos=pos,
|
| 136 |
+
query_pos=query_pos,
|
| 137 |
+
)
|
| 138 |
+
if self.return_intermediate:
|
| 139 |
+
intermediate.append(self.norm(output))
|
| 140 |
+
|
| 141 |
+
if self.norm is not None:
|
| 142 |
+
output = self.norm(output)
|
| 143 |
+
if self.return_intermediate:
|
| 144 |
+
intermediate.pop()
|
| 145 |
+
intermediate.append(output)
|
| 146 |
+
|
| 147 |
+
if self.return_intermediate:
|
| 148 |
+
return torch.stack(intermediate)
|
| 149 |
+
|
| 150 |
+
return output.unsqueeze(0)
|
| 151 |
+
|
| 152 |
+
|
| 153 |
+
class TransformerEncoderLayer(nn.Module):
|
| 154 |
+
def __init__(
|
| 155 |
+
self,
|
| 156 |
+
d_model,
|
| 157 |
+
nhead,
|
| 158 |
+
dim_feedforward=2048,
|
| 159 |
+
dropout=0.1,
|
| 160 |
+
activation="relu",
|
| 161 |
+
normalize_before=False,
|
| 162 |
+
):
|
| 163 |
+
super().__init__()
|
| 164 |
+
self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)
|
| 165 |
+
# Implementation of Feedforward model
|
| 166 |
+
self.linear1 = nn.Linear(d_model, dim_feedforward)
|
| 167 |
+
self.dropout = nn.Dropout(dropout)
|
| 168 |
+
self.linear2 = nn.Linear(dim_feedforward, d_model)
|
| 169 |
+
|
| 170 |
+
self.norm1 = nn.LayerNorm(d_model)
|
| 171 |
+
self.norm2 = nn.LayerNorm(d_model)
|
| 172 |
+
self.dropout1 = nn.Dropout(dropout)
|
| 173 |
+
self.dropout2 = nn.Dropout(dropout)
|
| 174 |
+
|
| 175 |
+
self.activation = _get_activation_fn(activation)
|
| 176 |
+
self.normalize_before = normalize_before
|
| 177 |
+
|
| 178 |
+
def with_pos_embed(self, tensor, pos: Optional[Tensor]):
|
| 179 |
+
return tensor if pos is None else tensor + pos
|
| 180 |
+
|
| 181 |
+
def forward_post(
|
| 182 |
+
self,
|
| 183 |
+
src,
|
| 184 |
+
src_mask: Optional[Tensor] = None,
|
| 185 |
+
src_key_padding_mask: Optional[Tensor] = None,
|
| 186 |
+
pos: Optional[Tensor] = None,
|
| 187 |
+
):
|
| 188 |
+
q = k = self.with_pos_embed(src, pos)
|
| 189 |
+
src2 = self.self_attn(
|
| 190 |
+
q, k, value=src, attn_mask=src_mask, key_padding_mask=src_key_padding_mask
|
| 191 |
+
)[0]
|
| 192 |
+
src = src + self.dropout1(src2)
|
| 193 |
+
src = self.norm1(src)
|
| 194 |
+
src2 = self.linear2(self.dropout(self.activation(self.linear1(src))))
|
| 195 |
+
src = src + self.dropout2(src2)
|
| 196 |
+
src = self.norm2(src)
|
| 197 |
+
return src
|
| 198 |
+
|
| 199 |
+
def forward_pre(
|
| 200 |
+
self,
|
| 201 |
+
src,
|
| 202 |
+
src_mask: Optional[Tensor] = None,
|
| 203 |
+
src_key_padding_mask: Optional[Tensor] = None,
|
| 204 |
+
pos: Optional[Tensor] = None,
|
| 205 |
+
):
|
| 206 |
+
src2 = self.norm1(src)
|
| 207 |
+
q = k = self.with_pos_embed(src2, pos)
|
| 208 |
+
src2 = self.self_attn(
|
| 209 |
+
q, k, value=src2, attn_mask=src_mask, key_padding_mask=src_key_padding_mask
|
| 210 |
+
)[0]
|
| 211 |
+
src = src + self.dropout1(src2)
|
| 212 |
+
src2 = self.norm2(src)
|
| 213 |
+
src2 = self.linear2(self.dropout(self.activation(self.linear1(src2))))
|
| 214 |
+
src = src + self.dropout2(src2)
|
| 215 |
+
return src
|
| 216 |
+
|
| 217 |
+
def forward(
|
| 218 |
+
self,
|
| 219 |
+
src,
|
| 220 |
+
src_mask: Optional[Tensor] = None,
|
| 221 |
+
src_key_padding_mask: Optional[Tensor] = None,
|
| 222 |
+
pos: Optional[Tensor] = None,
|
| 223 |
+
):
|
| 224 |
+
if self.normalize_before:
|
| 225 |
+
return self.forward_pre(src, src_mask, src_key_padding_mask, pos)
|
| 226 |
+
return self.forward_post(src, src_mask, src_key_padding_mask, pos)
|
| 227 |
+
|
| 228 |
+
|
| 229 |
+
class TransformerDecoderLayer(nn.Module):
|
| 230 |
+
def __init__(
|
| 231 |
+
self,
|
| 232 |
+
d_model,
|
| 233 |
+
nhead,
|
| 234 |
+
dim_feedforward=2048,
|
| 235 |
+
dropout=0.1,
|
| 236 |
+
activation="relu",
|
| 237 |
+
normalize_before=False,
|
| 238 |
+
):
|
| 239 |
+
super().__init__()
|
| 240 |
+
self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)
|
| 241 |
+
self.multihead_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)
|
| 242 |
+
# Implementation of Feedforward model
|
| 243 |
+
self.linear1 = nn.Linear(d_model, dim_feedforward)
|
| 244 |
+
self.dropout = nn.Dropout(dropout)
|
| 245 |
+
self.linear2 = nn.Linear(dim_feedforward, d_model)
|
| 246 |
+
|
| 247 |
+
self.norm1 = nn.LayerNorm(d_model)
|
| 248 |
+
self.norm2 = nn.LayerNorm(d_model)
|
| 249 |
+
self.norm3 = nn.LayerNorm(d_model)
|
| 250 |
+
self.dropout1 = nn.Dropout(dropout)
|
| 251 |
+
self.dropout2 = nn.Dropout(dropout)
|
| 252 |
+
self.dropout3 = nn.Dropout(dropout)
|
| 253 |
+
|
| 254 |
+
self.activation = _get_activation_fn(activation)
|
| 255 |
+
self.normalize_before = normalize_before
|
| 256 |
+
|
| 257 |
+
def with_pos_embed(self, tensor, pos: Optional[Tensor]):
|
| 258 |
+
return tensor if pos is None else tensor + pos
|
| 259 |
+
|
| 260 |
+
def forward_post(
|
| 261 |
+
self,
|
| 262 |
+
tgt,
|
| 263 |
+
memory,
|
| 264 |
+
tgt_mask: Optional[Tensor] = None,
|
| 265 |
+
memory_mask: Optional[Tensor] = None,
|
| 266 |
+
tgt_key_padding_mask: Optional[Tensor] = None,
|
| 267 |
+
memory_key_padding_mask: Optional[Tensor] = None,
|
| 268 |
+
pos: Optional[Tensor] = None,
|
| 269 |
+
query_pos: Optional[Tensor] = None,
|
| 270 |
+
):
|
| 271 |
+
q = k = self.with_pos_embed(tgt, query_pos)
|
| 272 |
+
tgt2 = self.self_attn(
|
| 273 |
+
q, k, value=tgt, attn_mask=tgt_mask, key_padding_mask=tgt_key_padding_mask
|
| 274 |
+
)[0]
|
| 275 |
+
tgt = tgt + self.dropout1(tgt2)
|
| 276 |
+
tgt = self.norm1(tgt)
|
| 277 |
+
tgt2 = self.multihead_attn(
|
| 278 |
+
query=self.with_pos_embed(tgt, query_pos),
|
| 279 |
+
key=self.with_pos_embed(memory, pos),
|
| 280 |
+
value=memory,
|
| 281 |
+
attn_mask=memory_mask,
|
| 282 |
+
key_padding_mask=memory_key_padding_mask,
|
| 283 |
+
)[0]
|
| 284 |
+
tgt = tgt + self.dropout2(tgt2)
|
| 285 |
+
tgt = self.norm2(tgt)
|
| 286 |
+
tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt))))
|
| 287 |
+
tgt = tgt + self.dropout3(tgt2)
|
| 288 |
+
tgt = self.norm3(tgt)
|
| 289 |
+
return tgt
|
| 290 |
+
|
| 291 |
+
def forward_pre(
|
| 292 |
+
self,
|
| 293 |
+
tgt,
|
| 294 |
+
memory,
|
| 295 |
+
tgt_mask: Optional[Tensor] = None,
|
| 296 |
+
memory_mask: Optional[Tensor] = None,
|
| 297 |
+
tgt_key_padding_mask: Optional[Tensor] = None,
|
| 298 |
+
memory_key_padding_mask: Optional[Tensor] = None,
|
| 299 |
+
pos: Optional[Tensor] = None,
|
| 300 |
+
query_pos: Optional[Tensor] = None,
|
| 301 |
+
):
|
| 302 |
+
tgt2 = self.norm1(tgt)
|
| 303 |
+
q = k = self.with_pos_embed(tgt2, query_pos)
|
| 304 |
+
tgt2 = self.self_attn(
|
| 305 |
+
q, k, value=tgt2, attn_mask=tgt_mask, key_padding_mask=tgt_key_padding_mask
|
| 306 |
+
)[0]
|
| 307 |
+
tgt = tgt + self.dropout1(tgt2)
|
| 308 |
+
tgt2 = self.norm2(tgt)
|
| 309 |
+
tgt2 = self.multihead_attn(
|
| 310 |
+
query=self.with_pos_embed(tgt2, query_pos),
|
| 311 |
+
key=self.with_pos_embed(memory, pos),
|
| 312 |
+
value=memory,
|
| 313 |
+
attn_mask=memory_mask,
|
| 314 |
+
key_padding_mask=memory_key_padding_mask,
|
| 315 |
+
)[0]
|
| 316 |
+
tgt = tgt + self.dropout2(tgt2)
|
| 317 |
+
tgt2 = self.norm3(tgt)
|
| 318 |
+
tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt2))))
|
| 319 |
+
tgt = tgt + self.dropout3(tgt2)
|
| 320 |
+
return tgt
|
| 321 |
+
|
| 322 |
+
def forward(
|
| 323 |
+
self,
|
| 324 |
+
tgt,
|
| 325 |
+
memory,
|
| 326 |
+
tgt_mask: Optional[Tensor] = None,
|
| 327 |
+
memory_mask: Optional[Tensor] = None,
|
| 328 |
+
tgt_key_padding_mask: Optional[Tensor] = None,
|
| 329 |
+
memory_key_padding_mask: Optional[Tensor] = None,
|
| 330 |
+
pos: Optional[Tensor] = None,
|
| 331 |
+
query_pos: Optional[Tensor] = None,
|
| 332 |
+
):
|
| 333 |
+
if self.normalize_before:
|
| 334 |
+
return self.forward_pre(
|
| 335 |
+
tgt,
|
| 336 |
+
memory,
|
| 337 |
+
tgt_mask,
|
| 338 |
+
memory_mask,
|
| 339 |
+
tgt_key_padding_mask,
|
| 340 |
+
memory_key_padding_mask,
|
| 341 |
+
pos,
|
| 342 |
+
query_pos,
|
| 343 |
+
)
|
| 344 |
+
return self.forward_post(
|
| 345 |
+
tgt,
|
| 346 |
+
memory,
|
| 347 |
+
tgt_mask,
|
| 348 |
+
memory_mask,
|
| 349 |
+
tgt_key_padding_mask,
|
| 350 |
+
memory_key_padding_mask,
|
| 351 |
+
pos,
|
| 352 |
+
query_pos,
|
| 353 |
+
)
|
| 354 |
+
|
| 355 |
+
|
| 356 |
+
def _get_clones(module, N):
|
| 357 |
+
return nn.ModuleList([copy.deepcopy(module) for i in range(N)])
|
| 358 |
+
|
| 359 |
+
|
| 360 |
+
def _get_activation_fn(activation):
|
| 361 |
+
"""Return an activation function given a string"""
|
| 362 |
+
if activation == "relu":
|
| 363 |
+
return F.relu
|
| 364 |
+
if activation == "gelu":
|
| 365 |
+
return F.gelu
|
| 366 |
+
if activation == "glu":
|
| 367 |
+
return F.glu
|
| 368 |
+
raise RuntimeError(f"activation should be relu/gelu, not {activation}.")
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/transformer_utils.py
ADDED
|
@@ -0,0 +1,192 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import Optional
|
| 2 |
+
from torch import nn, Tensor
|
| 3 |
+
from torch.nn import functional as F
|
| 4 |
+
|
| 5 |
+
class SelfAttentionLayer(nn.Module):
|
| 6 |
+
|
| 7 |
+
def __init__(self, d_model, nhead, dropout=0.0,
|
| 8 |
+
activation="relu", normalize_before=False):
|
| 9 |
+
super().__init__()
|
| 10 |
+
self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)
|
| 11 |
+
|
| 12 |
+
self.norm = nn.LayerNorm(d_model)
|
| 13 |
+
self.dropout = nn.Dropout(dropout)
|
| 14 |
+
|
| 15 |
+
self.activation = _get_activation_fn(activation)
|
| 16 |
+
self.normalize_before = normalize_before
|
| 17 |
+
|
| 18 |
+
self._reset_parameters()
|
| 19 |
+
|
| 20 |
+
def _reset_parameters(self):
|
| 21 |
+
for p in self.parameters():
|
| 22 |
+
if p.dim() > 1:
|
| 23 |
+
nn.init.xavier_uniform_(p)
|
| 24 |
+
|
| 25 |
+
def with_pos_embed(self, tensor, pos: Optional[Tensor]):
|
| 26 |
+
return tensor if pos is None else tensor + pos
|
| 27 |
+
|
| 28 |
+
def forward_post(self, tgt,
|
| 29 |
+
tgt_mask: Optional[Tensor] = None,
|
| 30 |
+
tgt_key_padding_mask: Optional[Tensor] = None,
|
| 31 |
+
query_pos: Optional[Tensor] = None):
|
| 32 |
+
q = k = self.with_pos_embed(tgt, query_pos)
|
| 33 |
+
tgt2 = self.self_attn(q, k, value=tgt, attn_mask=tgt_mask,
|
| 34 |
+
key_padding_mask=tgt_key_padding_mask)[0]
|
| 35 |
+
tgt = tgt + self.dropout(tgt2)
|
| 36 |
+
tgt = self.norm(tgt)
|
| 37 |
+
|
| 38 |
+
return tgt
|
| 39 |
+
|
| 40 |
+
def forward_pre(self, tgt,
|
| 41 |
+
tgt_mask: Optional[Tensor] = None,
|
| 42 |
+
tgt_key_padding_mask: Optional[Tensor] = None,
|
| 43 |
+
query_pos: Optional[Tensor] = None):
|
| 44 |
+
tgt2 = self.norm(tgt)
|
| 45 |
+
q = k = self.with_pos_embed(tgt2, query_pos)
|
| 46 |
+
tgt2 = self.self_attn(q, k, value=tgt2, attn_mask=tgt_mask,
|
| 47 |
+
key_padding_mask=tgt_key_padding_mask)[0]
|
| 48 |
+
tgt = tgt + self.dropout(tgt2)
|
| 49 |
+
|
| 50 |
+
return tgt
|
| 51 |
+
|
| 52 |
+
def forward(self, tgt,
|
| 53 |
+
tgt_mask: Optional[Tensor] = None,
|
| 54 |
+
tgt_key_padding_mask: Optional[Tensor] = None,
|
| 55 |
+
query_pos: Optional[Tensor] = None):
|
| 56 |
+
if self.normalize_before:
|
| 57 |
+
return self.forward_pre(tgt, tgt_mask,
|
| 58 |
+
tgt_key_padding_mask, query_pos)
|
| 59 |
+
return self.forward_post(tgt, tgt_mask,
|
| 60 |
+
tgt_key_padding_mask, query_pos)
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
class CrossAttentionLayer(nn.Module):
|
| 64 |
+
|
| 65 |
+
def __init__(self, d_model, nhead, dropout=0.0,
|
| 66 |
+
activation="relu", normalize_before=False):
|
| 67 |
+
super().__init__()
|
| 68 |
+
self.multihead_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)
|
| 69 |
+
|
| 70 |
+
self.norm = nn.LayerNorm(d_model)
|
| 71 |
+
self.dropout = nn.Dropout(dropout)
|
| 72 |
+
|
| 73 |
+
self.activation = _get_activation_fn(activation)
|
| 74 |
+
self.normalize_before = normalize_before
|
| 75 |
+
|
| 76 |
+
self._reset_parameters()
|
| 77 |
+
|
| 78 |
+
def _reset_parameters(self):
|
| 79 |
+
for p in self.parameters():
|
| 80 |
+
if p.dim() > 1:
|
| 81 |
+
nn.init.xavier_uniform_(p)
|
| 82 |
+
|
| 83 |
+
def with_pos_embed(self, tensor, pos: Optional[Tensor]):
|
| 84 |
+
return tensor if pos is None else tensor + pos
|
| 85 |
+
|
| 86 |
+
def forward_post(self, tgt, memory,
|
| 87 |
+
memory_mask: Optional[Tensor] = None,
|
| 88 |
+
memory_key_padding_mask: Optional[Tensor] = None,
|
| 89 |
+
pos: Optional[Tensor] = None,
|
| 90 |
+
query_pos: Optional[Tensor] = None):
|
| 91 |
+
tgt2 = self.multihead_attn(query=self.with_pos_embed(tgt, query_pos),
|
| 92 |
+
key=self.with_pos_embed(memory, pos),
|
| 93 |
+
value=memory, attn_mask=memory_mask,
|
| 94 |
+
key_padding_mask=memory_key_padding_mask)[0]
|
| 95 |
+
tgt = tgt + self.dropout(tgt2)
|
| 96 |
+
tgt = self.norm(tgt)
|
| 97 |
+
|
| 98 |
+
return tgt
|
| 99 |
+
|
| 100 |
+
def forward_pre(self, tgt, memory,
|
| 101 |
+
memory_mask: Optional[Tensor] = None,
|
| 102 |
+
memory_key_padding_mask: Optional[Tensor] = None,
|
| 103 |
+
pos: Optional[Tensor] = None,
|
| 104 |
+
query_pos: Optional[Tensor] = None):
|
| 105 |
+
tgt2 = self.norm(tgt)
|
| 106 |
+
tgt2 = self.multihead_attn(query=self.with_pos_embed(tgt2, query_pos),
|
| 107 |
+
key=self.with_pos_embed(memory, pos),
|
| 108 |
+
value=memory, attn_mask=memory_mask,
|
| 109 |
+
key_padding_mask=memory_key_padding_mask)[0]
|
| 110 |
+
tgt = tgt + self.dropout(tgt2)
|
| 111 |
+
|
| 112 |
+
return tgt
|
| 113 |
+
|
| 114 |
+
def forward(self, tgt, memory,
|
| 115 |
+
memory_mask: Optional[Tensor] = None,
|
| 116 |
+
memory_key_padding_mask: Optional[Tensor] = None,
|
| 117 |
+
pos: Optional[Tensor] = None,
|
| 118 |
+
query_pos: Optional[Tensor] = None):
|
| 119 |
+
if self.normalize_before:
|
| 120 |
+
return self.forward_pre(tgt, memory, memory_mask,
|
| 121 |
+
memory_key_padding_mask, pos, query_pos)
|
| 122 |
+
return self.forward_post(tgt, memory, memory_mask,
|
| 123 |
+
memory_key_padding_mask, pos, query_pos)
|
| 124 |
+
|
| 125 |
+
|
| 126 |
+
class FFNLayer(nn.Module):
|
| 127 |
+
|
| 128 |
+
def __init__(self, d_model, dim_feedforward=2048, dropout=0.0,
|
| 129 |
+
activation="relu", normalize_before=False):
|
| 130 |
+
super().__init__()
|
| 131 |
+
# Implementation of Feedforward model
|
| 132 |
+
self.linear1 = nn.Linear(d_model, dim_feedforward)
|
| 133 |
+
self.dropout = nn.Dropout(dropout)
|
| 134 |
+
self.linear2 = nn.Linear(dim_feedforward, d_model)
|
| 135 |
+
|
| 136 |
+
self.norm = nn.LayerNorm(d_model)
|
| 137 |
+
|
| 138 |
+
self.activation = _get_activation_fn(activation)
|
| 139 |
+
self.normalize_before = normalize_before
|
| 140 |
+
|
| 141 |
+
self._reset_parameters()
|
| 142 |
+
|
| 143 |
+
def _reset_parameters(self):
|
| 144 |
+
for p in self.parameters():
|
| 145 |
+
if p.dim() > 1:
|
| 146 |
+
nn.init.xavier_uniform_(p)
|
| 147 |
+
|
| 148 |
+
def with_pos_embed(self, tensor, pos: Optional[Tensor]):
|
| 149 |
+
return tensor if pos is None else tensor + pos
|
| 150 |
+
|
| 151 |
+
def forward_post(self, tgt):
|
| 152 |
+
tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt))))
|
| 153 |
+
tgt = tgt + self.dropout(tgt2)
|
| 154 |
+
tgt = self.norm(tgt)
|
| 155 |
+
return tgt
|
| 156 |
+
|
| 157 |
+
def forward_pre(self, tgt):
|
| 158 |
+
tgt2 = self.norm(tgt)
|
| 159 |
+
tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt2))))
|
| 160 |
+
tgt = tgt + self.dropout(tgt2)
|
| 161 |
+
return tgt
|
| 162 |
+
|
| 163 |
+
def forward(self, tgt):
|
| 164 |
+
if self.normalize_before:
|
| 165 |
+
return self.forward_pre(tgt)
|
| 166 |
+
return self.forward_post(tgt)
|
| 167 |
+
|
| 168 |
+
|
| 169 |
+
def _get_activation_fn(activation):
|
| 170 |
+
"""Return an activation function given a string"""
|
| 171 |
+
if activation == "relu":
|
| 172 |
+
return F.relu
|
| 173 |
+
if activation == "gelu":
|
| 174 |
+
return F.gelu
|
| 175 |
+
if activation == "glu":
|
| 176 |
+
return F.glu
|
| 177 |
+
raise RuntimeError(F"activation should be relu/gelu, not {activation}.")
|
| 178 |
+
|
| 179 |
+
|
| 180 |
+
class MLP(nn.Module):
|
| 181 |
+
""" Very simple multi-layer perceptron (also called FFN)"""
|
| 182 |
+
|
| 183 |
+
def __init__(self, input_dim, hidden_dim, output_dim, num_layers):
|
| 184 |
+
super().__init__()
|
| 185 |
+
self.num_layers = num_layers
|
| 186 |
+
h = [hidden_dim] * (num_layers - 1)
|
| 187 |
+
self.layers = nn.ModuleList(nn.Linear(n, k) for n, k in zip([input_dim] + h, h + [output_dim]))
|
| 188 |
+
|
| 189 |
+
def forward(self, x):
|
| 190 |
+
for i, layer in enumerate(self.layers):
|
| 191 |
+
x = F.relu(layer(x)) if i < self.num_layers - 1 else layer(x)
|
| 192 |
+
return x
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/unet.py
ADDED
|
@@ -0,0 +1,208 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from enum import Enum
|
| 2 |
+
import torch
|
| 3 |
+
import torch.nn as nn
|
| 4 |
+
from torch.nn import functional as F
|
| 5 |
+
import collections
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
NormType = Enum('NormType', 'Batch BatchZero Weight Spectral')
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
class Hook:
|
| 12 |
+
feature = None
|
| 13 |
+
|
| 14 |
+
def __init__(self, module):
|
| 15 |
+
self.hook = module.register_forward_hook(self.hook_fn)
|
| 16 |
+
|
| 17 |
+
def hook_fn(self, module, input, output):
|
| 18 |
+
if isinstance(output, torch.Tensor):
|
| 19 |
+
self.feature = output
|
| 20 |
+
elif isinstance(output, collections.OrderedDict):
|
| 21 |
+
self.feature = output['out']
|
| 22 |
+
|
| 23 |
+
def remove(self):
|
| 24 |
+
self.hook.remove()
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
class SelfAttention(nn.Module):
|
| 28 |
+
"Self attention layer for nd."
|
| 29 |
+
|
| 30 |
+
def __init__(self, n_channels: int):
|
| 31 |
+
super().__init__()
|
| 32 |
+
self.query = conv1d(n_channels, n_channels // 8)
|
| 33 |
+
self.key = conv1d(n_channels, n_channels // 8)
|
| 34 |
+
self.value = conv1d(n_channels, n_channels)
|
| 35 |
+
self.gamma = nn.Parameter(torch.tensor([0.]))
|
| 36 |
+
|
| 37 |
+
def forward(self, x):
|
| 38 |
+
#Notation from https://arxiv.org/pdf/1805.08318.pdf
|
| 39 |
+
size = x.size()
|
| 40 |
+
x = x.view(*size[:2], -1)
|
| 41 |
+
f, g, h = self.query(x), self.key(x), self.value(x)
|
| 42 |
+
beta = F.softmax(torch.bmm(f.permute(0, 2, 1).contiguous(), g), dim=1)
|
| 43 |
+
o = self.gamma * torch.bmm(h, beta) + x
|
| 44 |
+
return o.view(*size).contiguous()
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def batchnorm_2d(nf: int, norm_type: NormType = NormType.Batch):
|
| 48 |
+
"A batchnorm2d layer with `nf` features initialized depending on `norm_type`."
|
| 49 |
+
bn = nn.BatchNorm2d(nf)
|
| 50 |
+
with torch.no_grad():
|
| 51 |
+
bn.bias.fill_(1e-3)
|
| 52 |
+
bn.weight.fill_(0. if norm_type == NormType.BatchZero else 1.)
|
| 53 |
+
return bn
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def init_default(m: nn.Module, func=nn.init.kaiming_normal_) -> None:
|
| 57 |
+
"Initialize `m` weights with `func` and set `bias` to 0."
|
| 58 |
+
if func:
|
| 59 |
+
if hasattr(m, 'weight'): func(m.weight)
|
| 60 |
+
if hasattr(m, 'bias') and hasattr(m.bias, 'data'): m.bias.data.fill_(0.)
|
| 61 |
+
return m
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def icnr(x, scale=2, init=nn.init.kaiming_normal_):
|
| 65 |
+
"ICNR init of `x`, with `scale` and `init` function."
|
| 66 |
+
ni, nf, h, w = x.shape
|
| 67 |
+
ni2 = int(ni / (scale**2))
|
| 68 |
+
k = init(torch.zeros([ni2, nf, h, w])).transpose(0, 1)
|
| 69 |
+
k = k.contiguous().view(ni2, nf, -1)
|
| 70 |
+
k = k.repeat(1, 1, scale**2)
|
| 71 |
+
k = k.contiguous().view([nf, ni, h, w]).transpose(0, 1)
|
| 72 |
+
x.data.copy_(k)
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
def conv1d(ni: int, no: int, ks: int = 1, stride: int = 1, padding: int = 0, bias: bool = False):
|
| 76 |
+
"Create and initialize a `nn.Conv1d` layer with spectral normalization."
|
| 77 |
+
conv = nn.Conv1d(ni, no, ks, stride=stride, padding=padding, bias=bias)
|
| 78 |
+
nn.init.kaiming_normal_(conv.weight)
|
| 79 |
+
if bias: conv.bias.data.zero_()
|
| 80 |
+
return nn.utils.spectral_norm(conv)
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
def custom_conv_layer(
|
| 84 |
+
ni: int,
|
| 85 |
+
nf: int,
|
| 86 |
+
ks: int = 3,
|
| 87 |
+
stride: int = 1,
|
| 88 |
+
padding: int = None,
|
| 89 |
+
bias: bool = None,
|
| 90 |
+
is_1d: bool = False,
|
| 91 |
+
norm_type=NormType.Batch,
|
| 92 |
+
use_activ: bool = True,
|
| 93 |
+
transpose: bool = False,
|
| 94 |
+
init=nn.init.kaiming_normal_,
|
| 95 |
+
self_attention: bool = False,
|
| 96 |
+
extra_bn: bool = False,
|
| 97 |
+
):
|
| 98 |
+
"Create a sequence of convolutional (`ni` to `nf`), ReLU (if `use_activ`) and batchnorm (if `bn`) layers."
|
| 99 |
+
if padding is None:
|
| 100 |
+
padding = (ks - 1) // 2 if not transpose else 0
|
| 101 |
+
bn = norm_type in (NormType.Batch, NormType.BatchZero) or extra_bn == True
|
| 102 |
+
if bias is None:
|
| 103 |
+
bias = not bn
|
| 104 |
+
conv_func = nn.ConvTranspose2d if transpose else nn.Conv1d if is_1d else nn.Conv2d
|
| 105 |
+
conv = init_default(
|
| 106 |
+
conv_func(ni, nf, kernel_size=ks, bias=bias, stride=stride, padding=padding),
|
| 107 |
+
init,
|
| 108 |
+
)
|
| 109 |
+
|
| 110 |
+
if norm_type == NormType.Weight:
|
| 111 |
+
conv = nn.utils.weight_norm(conv)
|
| 112 |
+
elif norm_type == NormType.Spectral:
|
| 113 |
+
conv = nn.utils.spectral_norm(conv)
|
| 114 |
+
layers = [conv]
|
| 115 |
+
if use_activ:
|
| 116 |
+
layers.append(nn.ReLU(True))
|
| 117 |
+
if bn:
|
| 118 |
+
layers.append((nn.BatchNorm1d if is_1d else nn.BatchNorm2d)(nf))
|
| 119 |
+
if self_attention:
|
| 120 |
+
layers.append(SelfAttention(nf))
|
| 121 |
+
return nn.Sequential(*layers)
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
def conv_layer(ni: int,
|
| 125 |
+
nf: int,
|
| 126 |
+
ks: int = 3,
|
| 127 |
+
stride: int = 1,
|
| 128 |
+
padding: int = None,
|
| 129 |
+
bias: bool = None,
|
| 130 |
+
is_1d: bool = False,
|
| 131 |
+
norm_type=NormType.Batch,
|
| 132 |
+
use_activ: bool = True,
|
| 133 |
+
transpose: bool = False,
|
| 134 |
+
init=nn.init.kaiming_normal_,
|
| 135 |
+
self_attention: bool = False):
|
| 136 |
+
"Create a sequence of convolutional (`ni` to `nf`), ReLU (if `use_activ`) and batchnorm (if `bn`) layers."
|
| 137 |
+
if padding is None: padding = (ks - 1) // 2 if not transpose else 0
|
| 138 |
+
bn = norm_type in (NormType.Batch, NormType.BatchZero)
|
| 139 |
+
if bias is None: bias = not bn
|
| 140 |
+
conv_func = nn.ConvTranspose2d if transpose else nn.Conv1d if is_1d else nn.Conv2d
|
| 141 |
+
conv = init_default(conv_func(ni, nf, kernel_size=ks, bias=bias, stride=stride, padding=padding), init)
|
| 142 |
+
if norm_type == NormType.Weight: conv = nn.utils.weight_norm(conv)
|
| 143 |
+
elif norm_type == NormType.Spectral: conv = nn.utils.spectral_norm(conv)
|
| 144 |
+
layers = [conv]
|
| 145 |
+
if use_activ: layers.append(nn.ReLU(True))
|
| 146 |
+
if bn: layers.append((nn.BatchNorm1d if is_1d else nn.BatchNorm2d)(nf))
|
| 147 |
+
if self_attention: layers.append(SelfAttention(nf))
|
| 148 |
+
return nn.Sequential(*layers)
|
| 149 |
+
|
| 150 |
+
|
| 151 |
+
def _conv(ni: int, nf: int, ks: int = 3, stride: int = 1, **kwargs):
|
| 152 |
+
return conv_layer(ni, nf, ks=ks, stride=stride, norm_type=NormType.Spectral, **kwargs)
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
class CustomPixelShuffle_ICNR(nn.Module):
|
| 156 |
+
"Upsample by `scale` from `ni` filters to `nf` (default `ni`), using `nn.PixelShuffle`, `icnr` init, and `weight_norm`."
|
| 157 |
+
|
| 158 |
+
def __init__(self,
|
| 159 |
+
ni: int,
|
| 160 |
+
nf: int = None,
|
| 161 |
+
scale: int = 2,
|
| 162 |
+
blur: bool = True,
|
| 163 |
+
norm_type=NormType.Spectral,
|
| 164 |
+
extra_bn=False):
|
| 165 |
+
super().__init__()
|
| 166 |
+
self.conv = custom_conv_layer(
|
| 167 |
+
ni, nf * (scale**2), ks=1, use_activ=False, norm_type=norm_type, extra_bn=extra_bn)
|
| 168 |
+
icnr(self.conv[0].weight)
|
| 169 |
+
self.shuf = nn.PixelShuffle(scale)
|
| 170 |
+
self.do_blur = blur
|
| 171 |
+
# Blurring over (h*w) kernel
|
| 172 |
+
# "Super-Resolution using Convolutional Neural Networks without Any Checkerboard Artifacts"
|
| 173 |
+
# - https://arxiv.org/abs/1806.02658
|
| 174 |
+
self.pad = nn.ReplicationPad2d((1, 0, 1, 0))
|
| 175 |
+
self.blur = nn.AvgPool2d(2, stride=1)
|
| 176 |
+
self.relu = nn.ReLU(True)
|
| 177 |
+
|
| 178 |
+
def forward(self, x):
|
| 179 |
+
x = self.shuf(self.relu(self.conv(x)))
|
| 180 |
+
return self.blur(self.pad(x)) if self.do_blur else x
|
| 181 |
+
|
| 182 |
+
|
| 183 |
+
class UnetBlockWide(nn.Module):
|
| 184 |
+
"A quasi-UNet block, using `PixelShuffle_ICNR upsampling`."
|
| 185 |
+
|
| 186 |
+
def __init__(self,
|
| 187 |
+
up_in_c: int,
|
| 188 |
+
x_in_c: int,
|
| 189 |
+
n_out: int,
|
| 190 |
+
hook,
|
| 191 |
+
blur: bool = False,
|
| 192 |
+
self_attention: bool = False,
|
| 193 |
+
norm_type=NormType.Spectral):
|
| 194 |
+
super().__init__()
|
| 195 |
+
|
| 196 |
+
self.hook = hook
|
| 197 |
+
up_out = n_out
|
| 198 |
+
self.shuf = CustomPixelShuffle_ICNR(up_in_c, up_out, blur=blur, norm_type=norm_type, extra_bn=True)
|
| 199 |
+
self.bn = batchnorm_2d(x_in_c)
|
| 200 |
+
ni = up_out + x_in_c
|
| 201 |
+
self.conv = custom_conv_layer(ni, n_out, norm_type=norm_type, self_attention=self_attention, extra_bn=True)
|
| 202 |
+
self.relu = nn.ReLU()
|
| 203 |
+
|
| 204 |
+
def forward(self, up_in):
|
| 205 |
+
s = self.hook.feature
|
| 206 |
+
up_out = self.shuf(up_in)
|
| 207 |
+
cat_x = self.relu(torch.cat([up_out, self.bn(s)], dim=1))
|
| 208 |
+
return self.conv(cat_x)
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/util.py
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
import torch
|
| 3 |
+
from skimage import color
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
def rgb2lab(img_rgb):
|
| 7 |
+
img_lab = color.rgb2lab(img_rgb)
|
| 8 |
+
return img_lab[:, :, :1], img_lab[:, :, 1:]
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def tensor_lab2rgb(labs, illuminant="D65", observer="2"):
|
| 12 |
+
"""
|
| 13 |
+
Args:
|
| 14 |
+
lab : (B, C, H, W)
|
| 15 |
+
Returns:
|
| 16 |
+
tuple : (B, C, H, W)
|
| 17 |
+
"""
|
| 18 |
+
illuminants = \
|
| 19 |
+
{"A": {'2': (1.098466069456375, 1, 0.3558228003436005),
|
| 20 |
+
'10': (1.111420406956693, 1, 0.3519978321919493)},
|
| 21 |
+
"D50": {'2': (0.9642119944211994, 1, 0.8251882845188288),
|
| 22 |
+
'10': (0.9672062750333777, 1, 0.8142801513128616)},
|
| 23 |
+
"D55": {'2': (0.956797052643698, 1, 0.9214805860173273),
|
| 24 |
+
'10': (0.9579665682254781, 1, 0.9092525159847462)},
|
| 25 |
+
"D65": {'2': (0.95047, 1., 1.08883), # This was: `lab_ref_white`
|
| 26 |
+
'10': (0.94809667673716, 1, 1.0730513595166162)},
|
| 27 |
+
"D75": {'2': (0.9497220898840717, 1, 1.226393520724154),
|
| 28 |
+
'10': (0.9441713925645873, 1, 1.2064272211720228)},
|
| 29 |
+
"E": {'2': (1.0, 1.0, 1.0),
|
| 30 |
+
'10': (1.0, 1.0, 1.0)}}
|
| 31 |
+
xyz_from_rgb = np.array([[0.412453, 0.357580, 0.180423], [0.212671, 0.715160, 0.072169],
|
| 32 |
+
[0.019334, 0.119193, 0.950227]])
|
| 33 |
+
|
| 34 |
+
rgb_from_xyz = np.array([[3.240481340, -0.96925495, 0.055646640], [-1.53715152, 1.875990000, -0.20404134],
|
| 35 |
+
[-0.49853633, 0.041555930, 1.057311070]])
|
| 36 |
+
B, C, H, W = labs.shape
|
| 37 |
+
arrs = labs.permute((0, 2, 3, 1)).contiguous() # (B, 3, H, W) -> (B, H, W, 3)
|
| 38 |
+
L, a, b = arrs[:, :, :, 0:1], arrs[:, :, :, 1:2], arrs[:, :, :, 2:]
|
| 39 |
+
y = (L + 16.) / 116.
|
| 40 |
+
x = (a / 500.) + y
|
| 41 |
+
z = y - (b / 200.)
|
| 42 |
+
invalid = z.data < 0
|
| 43 |
+
z[invalid] = 0
|
| 44 |
+
xyz = torch.cat([x, y, z], dim=3)
|
| 45 |
+
mask = xyz.data > 0.2068966
|
| 46 |
+
mask_xyz = xyz.clone()
|
| 47 |
+
mask_xyz[mask] = torch.pow(xyz[mask], 3.0)
|
| 48 |
+
mask_xyz[~mask] = (xyz[~mask] - 16.0 / 116.) / 7.787
|
| 49 |
+
xyz_ref_white = illuminants[illuminant][observer]
|
| 50 |
+
for i in range(C):
|
| 51 |
+
mask_xyz[:, :, :, i] = mask_xyz[:, :, :, i] * xyz_ref_white[i]
|
| 52 |
+
|
| 53 |
+
rgb_trans = torch.mm(mask_xyz.view(-1, 3), torch.from_numpy(rgb_from_xyz).type_as(xyz)).view(B, H, W, C)
|
| 54 |
+
rgb = rgb_trans.permute((0, 3, 1, 2)).contiguous()
|
| 55 |
+
mask = rgb.data > 0.0031308
|
| 56 |
+
mask_rgb = rgb.clone()
|
| 57 |
+
mask_rgb[mask] = 1.055 * torch.pow(rgb[mask], 1 / 2.4) - 0.055
|
| 58 |
+
mask_rgb[~mask] = rgb[~mask] * 12.92
|
| 59 |
+
neg_mask = mask_rgb.data < 0
|
| 60 |
+
large_mask = mask_rgb.data > 1
|
| 61 |
+
mask_rgb[neg_mask] = 0
|
| 62 |
+
mask_rgb[large_mask] = 1
|
| 63 |
+
return mask_rgb
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/__init__.py
ADDED
|
@@ -0,0 +1,37 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .diffjpeg import DiffJPEG
|
| 2 |
+
from .file_client import FileClient
|
| 3 |
+
from .img_process_util import USMSharp, usm_sharp
|
| 4 |
+
from .img_util import crop_border, imfrombytes, img2tensor, imwrite, tensor2img
|
| 5 |
+
from .logger import AvgTimer, MessageLogger, get_env_info, get_root_logger, init_tb_logger, init_wandb_logger
|
| 6 |
+
from .misc import check_resume, get_time_str, make_exp_dirs, mkdir_and_rename, scandir, set_random_seed, sizeof_fmt
|
| 7 |
+
|
| 8 |
+
__all__ = [
|
| 9 |
+
# file_client.py
|
| 10 |
+
'FileClient',
|
| 11 |
+
# img_util.py
|
| 12 |
+
'img2tensor',
|
| 13 |
+
'tensor2img',
|
| 14 |
+
'imfrombytes',
|
| 15 |
+
'imwrite',
|
| 16 |
+
'crop_border',
|
| 17 |
+
# logger.py
|
| 18 |
+
'MessageLogger',
|
| 19 |
+
'AvgTimer',
|
| 20 |
+
'init_tb_logger',
|
| 21 |
+
'init_wandb_logger',
|
| 22 |
+
'get_root_logger',
|
| 23 |
+
'get_env_info',
|
| 24 |
+
# misc.py
|
| 25 |
+
'set_random_seed',
|
| 26 |
+
'get_time_str',
|
| 27 |
+
'mkdir_and_rename',
|
| 28 |
+
'make_exp_dirs',
|
| 29 |
+
'scandir',
|
| 30 |
+
'check_resume',
|
| 31 |
+
'sizeof_fmt',
|
| 32 |
+
# diffjpeg
|
| 33 |
+
'DiffJPEG',
|
| 34 |
+
# img_process_util
|
| 35 |
+
'USMSharp',
|
| 36 |
+
'usm_sharp'
|
| 37 |
+
]
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/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 basicsr.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]))
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/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
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/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)
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/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
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/img_util.py
ADDED
|
@@ -0,0 +1,227 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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, ...]
|
| 173 |
+
|
| 174 |
+
|
| 175 |
+
def tensor_lab2rgb(labs, illuminant="D65", observer="2"):
|
| 176 |
+
"""
|
| 177 |
+
Args:
|
| 178 |
+
lab : (B, C, H, W)
|
| 179 |
+
Returns:
|
| 180 |
+
tuple : (C, H, W)
|
| 181 |
+
"""
|
| 182 |
+
illuminants = \
|
| 183 |
+
{"A": {'2': (1.098466069456375, 1, 0.3558228003436005),
|
| 184 |
+
'10': (1.111420406956693, 1, 0.3519978321919493)},
|
| 185 |
+
"D50": {'2': (0.9642119944211994, 1, 0.8251882845188288),
|
| 186 |
+
'10': (0.9672062750333777, 1, 0.8142801513128616)},
|
| 187 |
+
"D55": {'2': (0.956797052643698, 1, 0.9214805860173273),
|
| 188 |
+
'10': (0.9579665682254781, 1, 0.9092525159847462)},
|
| 189 |
+
"D65": {'2': (0.95047, 1., 1.08883), # This was: `lab_ref_white`
|
| 190 |
+
'10': (0.94809667673716, 1, 1.0730513595166162)},
|
| 191 |
+
"D75": {'2': (0.9497220898840717, 1, 1.226393520724154),
|
| 192 |
+
'10': (0.9441713925645873, 1, 1.2064272211720228)},
|
| 193 |
+
"E": {'2': (1.0, 1.0, 1.0),
|
| 194 |
+
'10': (1.0, 1.0, 1.0)}}
|
| 195 |
+
xyz_from_rgb = np.array([[0.412453, 0.357580, 0.180423], [0.212671, 0.715160, 0.072169],
|
| 196 |
+
[0.019334, 0.119193, 0.950227]])
|
| 197 |
+
|
| 198 |
+
rgb_from_xyz = np.array([[3.240481340, -0.96925495, 0.055646640], [-1.53715152, 1.875990000, -0.20404134],
|
| 199 |
+
[-0.49853633, 0.041555930, 1.057311070]])
|
| 200 |
+
B, C, H, W = labs.shape
|
| 201 |
+
arrs = labs.permute((0, 2, 3, 1)).contiguous() # (B, 3, H, W) -> (B, H, W, 3)
|
| 202 |
+
L, a, b = arrs[:, :, :, 0:1], arrs[:, :, :, 1:2], arrs[:, :, :, 2:]
|
| 203 |
+
y = (L + 16.) / 116.
|
| 204 |
+
x = (a / 500.) + y
|
| 205 |
+
z = y - (b / 200.)
|
| 206 |
+
invalid = z.data < 0
|
| 207 |
+
z[invalid] = 0
|
| 208 |
+
xyz = torch.cat([x, y, z], dim=3)
|
| 209 |
+
mask = xyz.data > 0.2068966
|
| 210 |
+
mask_xyz = xyz.clone()
|
| 211 |
+
mask_xyz[mask] = torch.pow(xyz[mask], 3.0)
|
| 212 |
+
mask_xyz[~mask] = (xyz[~mask] - 16.0 / 116.) / 7.787
|
| 213 |
+
xyz_ref_white = illuminants[illuminant][observer]
|
| 214 |
+
for i in range(C):
|
| 215 |
+
mask_xyz[:, :, :, i] = mask_xyz[:, :, :, i] * xyz_ref_white[i]
|
| 216 |
+
|
| 217 |
+
rgb_trans = torch.mm(mask_xyz.view(-1, 3), torch.from_numpy(rgb_from_xyz).type_as(xyz)).view(B, H, W, C)
|
| 218 |
+
rgb = rgb_trans.permute((0, 3, 1, 2)).contiguous()
|
| 219 |
+
mask = rgb.data > 0.0031308
|
| 220 |
+
mask_rgb = rgb.clone()
|
| 221 |
+
mask_rgb[mask] = 1.055 * torch.pow(rgb[mask], 1 / 2.4) - 0.055
|
| 222 |
+
mask_rgb[~mask] = rgb[~mask] * 12.92
|
| 223 |
+
neg_mask = mask_rgb.data < 0
|
| 224 |
+
large_mask = mask_rgb.data > 1
|
| 225 |
+
mask_rgb[neg_mask] = 0
|
| 226 |
+
mask_rgb[large_mask] = 1
|
| 227 |
+
return mask_rgb
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/logger.py
ADDED
|
@@ -0,0 +1,209 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 = time.time()
|
| 22 |
+
|
| 23 |
+
def record(self):
|
| 24 |
+
self.count += 1
|
| 25 |
+
self.current_time = time.time() - self.start_time
|
| 26 |
+
self.total_time += self.current_time
|
| 27 |
+
# calculate average time
|
| 28 |
+
self.avg_time = self.total_time / self.count
|
| 29 |
+
# reset
|
| 30 |
+
if self.count > self.window:
|
| 31 |
+
self.count = 0
|
| 32 |
+
self.total_time = 0
|
| 33 |
+
|
| 34 |
+
def get_current_time(self):
|
| 35 |
+
return self.current_time
|
| 36 |
+
|
| 37 |
+
def get_avg_time(self):
|
| 38 |
+
return self.avg_time
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
class MessageLogger():
|
| 42 |
+
"""Message logger for printing.
|
| 43 |
+
|
| 44 |
+
Args:
|
| 45 |
+
opt (dict): Config. It contains the following keys:
|
| 46 |
+
name (str): Exp name.
|
| 47 |
+
logger (dict): Contains 'print_freq' (str) for logger interval.
|
| 48 |
+
train (dict): Contains 'total_iter' (int) for total iters.
|
| 49 |
+
use_tb_logger (bool): Use tensorboard logger.
|
| 50 |
+
start_iter (int): Start iter. Default: 1.
|
| 51 |
+
tb_logger (obj:`tb_logger`): Tensorboard logger. Default: None.
|
| 52 |
+
"""
|
| 53 |
+
|
| 54 |
+
def __init__(self, opt, start_iter=1, tb_logger=None):
|
| 55 |
+
self.exp_name = opt['name']
|
| 56 |
+
self.interval = opt['logger']['print_freq']
|
| 57 |
+
self.start_iter = start_iter
|
| 58 |
+
self.max_iters = opt['train']['total_iter']
|
| 59 |
+
self.use_tb_logger = opt['logger']['use_tb_logger']
|
| 60 |
+
self.tb_logger = tb_logger
|
| 61 |
+
self.start_time = time.time()
|
| 62 |
+
self.logger = get_root_logger()
|
| 63 |
+
|
| 64 |
+
def reset_start_time(self):
|
| 65 |
+
self.start_time = time.time()
|
| 66 |
+
|
| 67 |
+
@master_only
|
| 68 |
+
def __call__(self, log_vars):
|
| 69 |
+
"""Format logging message.
|
| 70 |
+
|
| 71 |
+
Args:
|
| 72 |
+
log_vars (dict): It contains the following keys:
|
| 73 |
+
epoch (int): Epoch number.
|
| 74 |
+
iter (int): Current iter.
|
| 75 |
+
lrs (list): List for learning rates.
|
| 76 |
+
|
| 77 |
+
time (float): Iter time.
|
| 78 |
+
data_time (float): Data time for each iter.
|
| 79 |
+
"""
|
| 80 |
+
# epoch, iter, learning rates
|
| 81 |
+
epoch = log_vars.pop('epoch')
|
| 82 |
+
current_iter = log_vars.pop('iter')
|
| 83 |
+
lrs = log_vars.pop('lrs')
|
| 84 |
+
|
| 85 |
+
message = (f'[{self.exp_name[:5]}..][epoch:{epoch:3d}, iter:{current_iter:8,d}, lr:(')
|
| 86 |
+
for v in lrs:
|
| 87 |
+
message += f'{v:.3e},'
|
| 88 |
+
message += ')] '
|
| 89 |
+
|
| 90 |
+
# time and estimated time
|
| 91 |
+
if 'time' in log_vars.keys():
|
| 92 |
+
iter_time = log_vars.pop('time')
|
| 93 |
+
data_time = log_vars.pop('data_time')
|
| 94 |
+
|
| 95 |
+
total_time = time.time() - self.start_time
|
| 96 |
+
time_sec_avg = total_time / (current_iter - self.start_iter + 1)
|
| 97 |
+
eta_sec = time_sec_avg * (self.max_iters - current_iter - 1)
|
| 98 |
+
eta_str = str(datetime.timedelta(seconds=int(eta_sec)))
|
| 99 |
+
message += f'[eta: {eta_str}, '
|
| 100 |
+
message += f'time (data): {iter_time:.3f} ({data_time:.3f})] '
|
| 101 |
+
|
| 102 |
+
# other items, especially losses
|
| 103 |
+
for k, v in log_vars.items():
|
| 104 |
+
message += f'{k}: {v:.4e} '
|
| 105 |
+
# tensorboard logger
|
| 106 |
+
if self.use_tb_logger and 'debug' not in self.exp_name:
|
| 107 |
+
if k.startswith('l_'):
|
| 108 |
+
self.tb_logger.add_scalar(f'losses/{k}', v, current_iter)
|
| 109 |
+
else:
|
| 110 |
+
self.tb_logger.add_scalar(k, v, current_iter)
|
| 111 |
+
self.logger.info(message)
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
@master_only
|
| 115 |
+
def init_tb_logger(log_dir):
|
| 116 |
+
from torch.utils.tensorboard import SummaryWriter
|
| 117 |
+
tb_logger = SummaryWriter(log_dir=log_dir)
|
| 118 |
+
return tb_logger
|
| 119 |
+
|
| 120 |
+
|
| 121 |
+
@master_only
|
| 122 |
+
def init_wandb_logger(opt):
|
| 123 |
+
"""We now only use wandb to sync tensorboard log."""
|
| 124 |
+
import wandb
|
| 125 |
+
logger = get_root_logger()
|
| 126 |
+
|
| 127 |
+
project = opt['logger']['wandb']['project']
|
| 128 |
+
resume_id = opt['logger']['wandb'].get('resume_id')
|
| 129 |
+
if resume_id:
|
| 130 |
+
wandb_id = resume_id
|
| 131 |
+
resume = 'allow'
|
| 132 |
+
logger.warning(f'Resume wandb logger with id={wandb_id}.')
|
| 133 |
+
else:
|
| 134 |
+
wandb_id = wandb.util.generate_id()
|
| 135 |
+
resume = 'never'
|
| 136 |
+
|
| 137 |
+
wandb.init(id=wandb_id, resume=resume, name=opt['name'], config=opt, project=project, sync_tensorboard=True)
|
| 138 |
+
|
| 139 |
+
logger.info(f'Use wandb logger with id={wandb_id}; project={project}.')
|
| 140 |
+
|
| 141 |
+
|
| 142 |
+
def get_root_logger(logger_name='basicsr', log_level=logging.INFO, log_file=None):
|
| 143 |
+
"""Get the root logger.
|
| 144 |
+
|
| 145 |
+
The logger will be initialized if it has not been initialized. By default a
|
| 146 |
+
StreamHandler will be added. If `log_file` is specified, a FileHandler will
|
| 147 |
+
also be added.
|
| 148 |
+
|
| 149 |
+
Args:
|
| 150 |
+
logger_name (str): root logger name. Default: 'basicsr'.
|
| 151 |
+
log_file (str | None): The log filename. If specified, a FileHandler
|
| 152 |
+
will be added to the root logger.
|
| 153 |
+
log_level (int): The root logger level. Note that only the process of
|
| 154 |
+
rank 0 is affected, while other processes will set the level to
|
| 155 |
+
"Error" and be silent most of the time.
|
| 156 |
+
|
| 157 |
+
Returns:
|
| 158 |
+
logging.Logger: The root logger.
|
| 159 |
+
"""
|
| 160 |
+
logger = logging.getLogger(logger_name)
|
| 161 |
+
# if the logger has been initialized, just return it
|
| 162 |
+
if logger_name in initialized_logger:
|
| 163 |
+
return logger
|
| 164 |
+
|
| 165 |
+
format_str = '%(asctime)s %(levelname)s: %(message)s'
|
| 166 |
+
stream_handler = logging.StreamHandler()
|
| 167 |
+
stream_handler.setFormatter(logging.Formatter(format_str))
|
| 168 |
+
logger.addHandler(stream_handler)
|
| 169 |
+
logger.propagate = False
|
| 170 |
+
rank, _ = get_dist_info()
|
| 171 |
+
if rank != 0:
|
| 172 |
+
logger.setLevel('ERROR')
|
| 173 |
+
elif log_file is not None:
|
| 174 |
+
logger.setLevel(log_level)
|
| 175 |
+
# add file handler
|
| 176 |
+
file_handler = logging.FileHandler(log_file, 'w')
|
| 177 |
+
file_handler.setFormatter(logging.Formatter(format_str))
|
| 178 |
+
file_handler.setLevel(log_level)
|
| 179 |
+
logger.addHandler(file_handler)
|
| 180 |
+
initialized_logger[logger_name] = True
|
| 181 |
+
return logger
|
| 182 |
+
|
| 183 |
+
|
| 184 |
+
def get_env_info():
|
| 185 |
+
"""Get environment information.
|
| 186 |
+
|
| 187 |
+
Currently, only log the software version.
|
| 188 |
+
"""
|
| 189 |
+
import torch
|
| 190 |
+
import torchvision
|
| 191 |
+
|
| 192 |
+
from basicsr.version import __version__
|
| 193 |
+
msg = r"""
|
| 194 |
+
____ _ _____ ____
|
| 195 |
+
/ __ ) ____ _ _____ (_)_____/ ___/ / __ \
|
| 196 |
+
/ __ |/ __ `// ___// // ___/\__ \ / /_/ /
|
| 197 |
+
/ /_/ // /_/ /(__ )/ // /__ ___/ // _, _/
|
| 198 |
+
/_____/ \__,_//____//_/ \___//____//_/ |_|
|
| 199 |
+
______ __ __ __ __
|
| 200 |
+
/ ____/____ ____ ____/ / / / __ __ _____ / /__ / /
|
| 201 |
+
/ / __ / __ \ / __ \ / __ / / / / / / // ___// //_/ / /
|
| 202 |
+
/ /_/ // /_/ // /_/ // /_/ / / /___/ /_/ // /__ / /< /_/
|
| 203 |
+
\____/ \____/ \____/ \____/ /_____/\____/ \___//_/|_| (_)
|
| 204 |
+
"""
|
| 205 |
+
msg += ('\nVersion Information: '
|
| 206 |
+
f'\n\tBasicSR: {__version__}'
|
| 207 |
+
f'\n\tPyTorch: {torch.__version__}'
|
| 208 |
+
f'\n\tTorchVision: {torchvision.__version__}')
|
| 209 |
+
return msg
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/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 siz.
|
| 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}'
|
experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/registry.py
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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):
|
| 39 |
+
assert (name not in self._obj_map), (f"An object named '{name}' was already registered "
|
| 40 |
+
f"in '{self._name}' registry!")
|
| 41 |
+
self._obj_map[name] = obj
|
| 42 |
+
|
| 43 |
+
def register(self, obj=None):
|
| 44 |
+
"""
|
| 45 |
+
Register the given object under the the name `obj.__name__`.
|
| 46 |
+
Can be used as either a decorator or not.
|
| 47 |
+
See docstring of this class for usage.
|
| 48 |
+
"""
|
| 49 |
+
if obj is None:
|
| 50 |
+
# used as a decorator
|
| 51 |
+
def deco(func_or_class):
|
| 52 |
+
name = func_or_class.__name__
|
| 53 |
+
self._do_register(name, func_or_class)
|
| 54 |
+
return func_or_class
|
| 55 |
+
|
| 56 |
+
return deco
|
| 57 |
+
|
| 58 |
+
# used as a function call
|
| 59 |
+
name = obj.__name__
|
| 60 |
+
self._do_register(name, obj)
|
| 61 |
+
|
| 62 |
+
def get(self, name):
|
| 63 |
+
ret = self._obj_map.get(name)
|
| 64 |
+
if ret is None:
|
| 65 |
+
raise KeyError(f"No object named '{name}' found in '{self._name}' registry!")
|
| 66 |
+
return ret
|
| 67 |
+
|
| 68 |
+
def __contains__(self, name):
|
| 69 |
+
return name in self._obj_map
|
| 70 |
+
|
| 71 |
+
def __iter__(self):
|
| 72 |
+
return iter(self._obj_map.items())
|
| 73 |
+
|
| 74 |
+
def keys(self):
|
| 75 |
+
return self._obj_map.keys()
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
DATASET_REGISTRY = Registry('dataset')
|
| 79 |
+
ARCH_REGISTRY = Registry('arch')
|
| 80 |
+
MODEL_REGISTRY = Registry('model')
|
| 81 |
+
LOSS_REGISTRY = Registry('loss')
|
| 82 |
+
METRIC_REGISTRY = Registry('metric')
|
experiments/round5-20260927/source/vendor_ddcolor/ddcolor/__init__.py
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .model import DDColor
|
| 2 |
+
from .pipeline import ColorizationPipeline, build_ddcolor_model, load_checkpoint_state_dict
|
| 3 |
+
|
| 4 |
+
__all__ = [
|
| 5 |
+
"DDColor",
|
| 6 |
+
"ColorizationPipeline",
|
| 7 |
+
"build_ddcolor_model",
|
| 8 |
+
"load_checkpoint_state_dict",
|
| 9 |
+
]
|
experiments/round5-20260927/source/vendor_ddcolor/ddcolor/model.py
ADDED
|
@@ -0,0 +1,278 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
|
| 4 |
+
from basicsr.archs.ddcolor_arch_utils.unet import Hook, CustomPixelShuffle_ICNR, UnetBlockWide, NormType, custom_conv_layer
|
| 5 |
+
from basicsr.archs.ddcolor_arch_utils.convnext import ConvNeXt
|
| 6 |
+
from basicsr.archs.ddcolor_arch_utils.transformer_utils import SelfAttentionLayer, CrossAttentionLayer, FFNLayer, MLP
|
| 7 |
+
from basicsr.archs.ddcolor_arch_utils.position_encoding import PositionEmbeddingSine
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class DDColor(nn.Module):
|
| 11 |
+
def __init__(
|
| 12 |
+
self,
|
| 13 |
+
encoder_name='convnext-l',
|
| 14 |
+
decoder_name='MultiScaleColorDecoder',
|
| 15 |
+
num_input_channels=3,
|
| 16 |
+
input_size=(256, 256),
|
| 17 |
+
nf=512,
|
| 18 |
+
num_output_channels=3,
|
| 19 |
+
last_norm='Weight',
|
| 20 |
+
do_normalize=False,
|
| 21 |
+
num_queries=256,
|
| 22 |
+
num_scales=3,
|
| 23 |
+
dec_layers=9,
|
| 24 |
+
):
|
| 25 |
+
super().__init__()
|
| 26 |
+
|
| 27 |
+
self.encoder = ImageEncoder(encoder_name, ['norm0', 'norm1', 'norm2', 'norm3'])
|
| 28 |
+
self.encoder.eval()
|
| 29 |
+
test_input = torch.randn(1, num_input_channels, *input_size)
|
| 30 |
+
|
| 31 |
+
with torch.no_grad():
|
| 32 |
+
self.encoder(test_input)
|
| 33 |
+
|
| 34 |
+
self.decoder = DuelDecoder(
|
| 35 |
+
self.encoder.hooks,
|
| 36 |
+
nf=nf,
|
| 37 |
+
last_norm=last_norm,
|
| 38 |
+
num_queries=num_queries,
|
| 39 |
+
num_scales=num_scales,
|
| 40 |
+
dec_layers=dec_layers,
|
| 41 |
+
decoder_name=decoder_name
|
| 42 |
+
)
|
| 43 |
+
|
| 44 |
+
self.refine_net = nn.Sequential(
|
| 45 |
+
custom_conv_layer(num_queries + 3, num_output_channels, ks=1, use_activ=False, norm_type=NormType.Spectral)
|
| 46 |
+
)
|
| 47 |
+
|
| 48 |
+
self.do_normalize = do_normalize
|
| 49 |
+
self.register_buffer('mean', torch.Tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1))
|
| 50 |
+
self.register_buffer('std', torch.Tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1))
|
| 51 |
+
|
| 52 |
+
def normalize(self, img):
|
| 53 |
+
return (img - self.mean) / self.std
|
| 54 |
+
|
| 55 |
+
def denormalize(self, img):
|
| 56 |
+
return img * self.std + self.mean
|
| 57 |
+
|
| 58 |
+
def forward(self, x):
|
| 59 |
+
if x.shape[1] == 3:
|
| 60 |
+
x = self.normalize(x)
|
| 61 |
+
|
| 62 |
+
self.encoder(x)
|
| 63 |
+
out_feat = self.decoder()
|
| 64 |
+
coarse_input = torch.cat([out_feat, x], dim=1)
|
| 65 |
+
out = self.refine_net(coarse_input)
|
| 66 |
+
|
| 67 |
+
if self.do_normalize:
|
| 68 |
+
out = self.denormalize(out)
|
| 69 |
+
return out
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
class ImageEncoder(nn.Module):
|
| 73 |
+
def __init__(self, encoder_name, hook_names):
|
| 74 |
+
super().__init__()
|
| 75 |
+
|
| 76 |
+
assert encoder_name == 'convnext-t' or encoder_name == 'convnext-l'
|
| 77 |
+
if encoder_name == 'convnext-t':
|
| 78 |
+
self.arch = ConvNeXt(depths=[3, 3, 9, 3], dims=[96, 192, 384, 768])
|
| 79 |
+
elif encoder_name == 'convnext-l':
|
| 80 |
+
self.arch = ConvNeXt(depths=[3, 3, 27, 3], dims=[192, 384, 768, 1536])
|
| 81 |
+
else:
|
| 82 |
+
raise NotImplementedError
|
| 83 |
+
|
| 84 |
+
self.encoder_name = encoder_name
|
| 85 |
+
self.hook_names = hook_names
|
| 86 |
+
self.hooks = self.setup_hooks()
|
| 87 |
+
|
| 88 |
+
def setup_hooks(self):
|
| 89 |
+
hooks = [Hook(self.arch._modules[name]) for name in self.hook_names]
|
| 90 |
+
return hooks
|
| 91 |
+
|
| 92 |
+
def forward(self, x):
|
| 93 |
+
return self.arch(x)
|
| 94 |
+
|
| 95 |
+
|
| 96 |
+
class DuelDecoder(nn.Module):
|
| 97 |
+
def __init__(
|
| 98 |
+
self,
|
| 99 |
+
hooks,
|
| 100 |
+
nf=512,
|
| 101 |
+
blur=True,
|
| 102 |
+
last_norm='Weight',
|
| 103 |
+
num_queries=256,
|
| 104 |
+
num_scales=3,
|
| 105 |
+
dec_layers=9,
|
| 106 |
+
decoder_name='MultiScaleColorDecoder',
|
| 107 |
+
):
|
| 108 |
+
super().__init__()
|
| 109 |
+
self.hooks = hooks
|
| 110 |
+
self.nf = nf
|
| 111 |
+
self.blur = blur
|
| 112 |
+
self.last_norm = getattr(NormType, last_norm)
|
| 113 |
+
self.decoder_name = decoder_name
|
| 114 |
+
|
| 115 |
+
self.layers = self.make_layers()
|
| 116 |
+
embed_dim = nf // 2
|
| 117 |
+
self.last_shuf = CustomPixelShuffle_ICNR(embed_dim, embed_dim, blur=self.blur, norm_type=self.last_norm, scale=4)
|
| 118 |
+
|
| 119 |
+
assert decoder_name == 'MultiScaleColorDecoder'
|
| 120 |
+
self.color_decoder = MultiScaleColorDecoder(
|
| 121 |
+
in_channels=[512, 512, 256],
|
| 122 |
+
num_queries=num_queries,
|
| 123 |
+
num_scales=num_scales,
|
| 124 |
+
dec_layers=dec_layers,
|
| 125 |
+
)
|
| 126 |
+
|
| 127 |
+
def make_layers(self):
|
| 128 |
+
decoder_layers = []
|
| 129 |
+
in_c = self.hooks[-1].feature.shape[1]
|
| 130 |
+
out_c = self.nf
|
| 131 |
+
|
| 132 |
+
setup_hooks = self.hooks[-2::-1]
|
| 133 |
+
for layer_index, hook in enumerate(setup_hooks):
|
| 134 |
+
feature_c = hook.feature.shape[1]
|
| 135 |
+
if layer_index == len(setup_hooks) - 1:
|
| 136 |
+
out_c = out_c // 2
|
| 137 |
+
decoder_layers.append(
|
| 138 |
+
UnetBlockWide(
|
| 139 |
+
in_c, feature_c, out_c, hook, blur=self.blur, self_attention=False, norm_type=NormType.Spectral))
|
| 140 |
+
in_c = out_c
|
| 141 |
+
|
| 142 |
+
return nn.Sequential(*decoder_layers)
|
| 143 |
+
|
| 144 |
+
def forward(self):
|
| 145 |
+
encode_feat = self.hooks[-1].feature
|
| 146 |
+
out0 = self.layers[0](encode_feat)
|
| 147 |
+
out1 = self.layers[1](out0)
|
| 148 |
+
out2 = self.layers[2](out1)
|
| 149 |
+
out3 = self.last_shuf(out2)
|
| 150 |
+
|
| 151 |
+
return self.color_decoder([out0, out1, out2], out3)
|
| 152 |
+
|
| 153 |
+
|
| 154 |
+
class MultiScaleColorDecoder(nn.Module):
|
| 155 |
+
def __init__(
|
| 156 |
+
self,
|
| 157 |
+
in_channels,
|
| 158 |
+
hidden_dim=256,
|
| 159 |
+
num_queries=100,
|
| 160 |
+
nheads=8,
|
| 161 |
+
dim_feedforward=2048,
|
| 162 |
+
dec_layers=9,
|
| 163 |
+
pre_norm=False,
|
| 164 |
+
color_embed_dim=256,
|
| 165 |
+
enforce_input_project=True,
|
| 166 |
+
num_scales=3,
|
| 167 |
+
):
|
| 168 |
+
super().__init__()
|
| 169 |
+
|
| 170 |
+
self.hidden_dim = hidden_dim
|
| 171 |
+
self.num_queries = num_queries
|
| 172 |
+
self.num_layers = dec_layers
|
| 173 |
+
self.num_feature_levels = num_scales
|
| 174 |
+
|
| 175 |
+
# Positional encoding layer
|
| 176 |
+
self.pe_layer = PositionEmbeddingSine(hidden_dim // 2, normalize=True)
|
| 177 |
+
|
| 178 |
+
# Learnable query features and embeddings
|
| 179 |
+
self.query_feat = nn.Embedding(num_queries, hidden_dim)
|
| 180 |
+
self.query_embed = nn.Embedding(num_queries, hidden_dim)
|
| 181 |
+
|
| 182 |
+
# Learnable level embeddings
|
| 183 |
+
self.level_embed = nn.Embedding(num_scales, hidden_dim)
|
| 184 |
+
|
| 185 |
+
# Input projection layers
|
| 186 |
+
self.input_proj = nn.ModuleList(
|
| 187 |
+
[self._make_input_proj(in_ch, hidden_dim, enforce_input_project) for in_ch in in_channels]
|
| 188 |
+
)
|
| 189 |
+
|
| 190 |
+
# Transformer layers
|
| 191 |
+
self.transformer_self_attention_layers = nn.ModuleList()
|
| 192 |
+
self.transformer_cross_attention_layers = nn.ModuleList()
|
| 193 |
+
self.transformer_ffn_layers = nn.ModuleList()
|
| 194 |
+
|
| 195 |
+
for _ in range(dec_layers):
|
| 196 |
+
self.transformer_self_attention_layers.append(
|
| 197 |
+
SelfAttentionLayer(
|
| 198 |
+
d_model=hidden_dim,
|
| 199 |
+
nhead=nheads,
|
| 200 |
+
dropout=0.0,
|
| 201 |
+
normalize_before=pre_norm,
|
| 202 |
+
)
|
| 203 |
+
)
|
| 204 |
+
self.transformer_cross_attention_layers.append(
|
| 205 |
+
CrossAttentionLayer(
|
| 206 |
+
d_model=hidden_dim,
|
| 207 |
+
nhead=nheads,
|
| 208 |
+
dropout=0.0,
|
| 209 |
+
normalize_before=pre_norm,
|
| 210 |
+
)
|
| 211 |
+
)
|
| 212 |
+
self.transformer_ffn_layers.append(
|
| 213 |
+
FFNLayer(
|
| 214 |
+
d_model=hidden_dim,
|
| 215 |
+
dim_feedforward=dim_feedforward,
|
| 216 |
+
dropout=0.0,
|
| 217 |
+
normalize_before=pre_norm,
|
| 218 |
+
)
|
| 219 |
+
)
|
| 220 |
+
|
| 221 |
+
# Layer normalization for the decoder output
|
| 222 |
+
self.decoder_norm = nn.LayerNorm(hidden_dim)
|
| 223 |
+
|
| 224 |
+
# Output embedding layer
|
| 225 |
+
self.color_embed = MLP(hidden_dim, hidden_dim, color_embed_dim, 3)
|
| 226 |
+
|
| 227 |
+
def forward(self, x, img_features):
|
| 228 |
+
assert len(x) == self.num_feature_levels
|
| 229 |
+
|
| 230 |
+
src, pos = self._get_src_and_pos(x)
|
| 231 |
+
|
| 232 |
+
bs = src[0].shape[1]
|
| 233 |
+
|
| 234 |
+
# Prepare query embeddings (QxNxC)
|
| 235 |
+
query_embed = self.query_embed.weight.unsqueeze(1).repeat(1, bs, 1)
|
| 236 |
+
output = self.query_feat.weight.unsqueeze(1).repeat(1, bs, 1)
|
| 237 |
+
|
| 238 |
+
for i in range(self.num_layers):
|
| 239 |
+
level_index = i % self.num_feature_levels
|
| 240 |
+
# attention: cross-attention first
|
| 241 |
+
output = self.transformer_cross_attention_layers[i](
|
| 242 |
+
output, src[level_index],
|
| 243 |
+
memory_mask=None,
|
| 244 |
+
memory_key_padding_mask=None,
|
| 245 |
+
pos=pos[level_index], query_pos=query_embed
|
| 246 |
+
)
|
| 247 |
+
output = self.transformer_self_attention_layers[i](
|
| 248 |
+
output, tgt_mask=None,
|
| 249 |
+
tgt_key_padding_mask=None,
|
| 250 |
+
query_pos=query_embed
|
| 251 |
+
)
|
| 252 |
+
# FFN
|
| 253 |
+
output = self.transformer_ffn_layers[i](
|
| 254 |
+
output
|
| 255 |
+
)
|
| 256 |
+
|
| 257 |
+
decoder_output = self.decoder_norm(output).transpose(0, 1)
|
| 258 |
+
color_embed = self.color_embed(decoder_output)
|
| 259 |
+
|
| 260 |
+
out = torch.einsum("bqc,bchw->bqhw", color_embed, img_features)
|
| 261 |
+
|
| 262 |
+
return out
|
| 263 |
+
|
| 264 |
+
def _make_input_proj(self, in_ch, hidden_dim, enforce):
|
| 265 |
+
if in_ch != hidden_dim or enforce:
|
| 266 |
+
proj = nn.Conv2d(in_ch, hidden_dim, kernel_size=1)
|
| 267 |
+
nn.init.kaiming_uniform_(proj.weight, a=1)
|
| 268 |
+
if proj.bias is not None:
|
| 269 |
+
nn.init.constant_(proj.bias, 0)
|
| 270 |
+
return proj
|
| 271 |
+
return nn.Sequential()
|
| 272 |
+
|
| 273 |
+
def _get_src_and_pos(self, x):
|
| 274 |
+
src, pos = [], []
|
| 275 |
+
for i, feature in enumerate(x):
|
| 276 |
+
pos.append(self.pe_layer(feature).flatten(2).permute(2, 0, 1)) # flatten NxCxHxW to HWxNxC
|
| 277 |
+
src.append((self.input_proj[i](feature).flatten(2) + self.level_embed.weight[i][None, :, None]).permute(2, 0, 1))
|
| 278 |
+
return src, pos
|
experiments/round5-20260927/source/vendor_ddcolor/ddcolor/pipeline.py
ADDED
|
@@ -0,0 +1,127 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import cv2
|
| 2 |
+
import numpy as np
|
| 3 |
+
import torch
|
| 4 |
+
import torch.nn.functional as F
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
def load_checkpoint_state_dict(model_path: str, map_location="cpu"):
|
| 8 |
+
"""Load a checkpoint and return a state_dict.
|
| 9 |
+
|
| 10 |
+
Supports both:
|
| 11 |
+
- {'params': state_dict, ...} (common in this repo)
|
| 12 |
+
- raw state_dict
|
| 13 |
+
"""
|
| 14 |
+
ckpt = torch.load(model_path, map_location=map_location)
|
| 15 |
+
if isinstance(ckpt, dict) and "params" in ckpt:
|
| 16 |
+
return ckpt["params"]
|
| 17 |
+
return ckpt
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def build_ddcolor_model(
|
| 21 |
+
model_cls,
|
| 22 |
+
*,
|
| 23 |
+
model_path: str,
|
| 24 |
+
input_size: int = 512,
|
| 25 |
+
model_size: str = "large",
|
| 26 |
+
decoder_type: str = "MultiScaleColorDecoder",
|
| 27 |
+
device=None,
|
| 28 |
+
**kwargs,
|
| 29 |
+
):
|
| 30 |
+
"""Build a DDColor model and load weights.
|
| 31 |
+
|
| 32 |
+
This helper is intentionally backend-agnostic: `model_cls` can be
|
| 33 |
+
`ddcolor.DDColor` or `basicsr.archs.ddcolor_arch.DDColor` as long as
|
| 34 |
+
it supports the common constructor args used below.
|
| 35 |
+
"""
|
| 36 |
+
if device is None:
|
| 37 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 38 |
+
|
| 39 |
+
if model_size not in ("tiny", "large"):
|
| 40 |
+
raise ValueError(f"model_size must be 'tiny' or 'large', got: {model_size}")
|
| 41 |
+
encoder_name = "convnext-t" if model_size == "tiny" else "convnext-l"
|
| 42 |
+
|
| 43 |
+
if decoder_type == "MultiScaleColorDecoder":
|
| 44 |
+
# keep default consistent with existing scripts
|
| 45 |
+
kwargs.setdefault("num_queries", 100)
|
| 46 |
+
kwargs.setdefault("num_scales", 3)
|
| 47 |
+
kwargs.setdefault("dec_layers", 9)
|
| 48 |
+
elif decoder_type == "SingleColorDecoder":
|
| 49 |
+
kwargs.setdefault("num_queries", 256)
|
| 50 |
+
else:
|
| 51 |
+
raise NotImplementedError(f"decoder_type not implemented: {decoder_type}")
|
| 52 |
+
|
| 53 |
+
model = model_cls(
|
| 54 |
+
encoder_name=encoder_name,
|
| 55 |
+
decoder_name=decoder_type,
|
| 56 |
+
input_size=[input_size, input_size],
|
| 57 |
+
num_output_channels=2,
|
| 58 |
+
last_norm="Spectral",
|
| 59 |
+
do_normalize=False,
|
| 60 |
+
**kwargs,
|
| 61 |
+
)
|
| 62 |
+
|
| 63 |
+
state_dict = load_checkpoint_state_dict(model_path, map_location="cpu")
|
| 64 |
+
model.load_state_dict(state_dict, strict=False)
|
| 65 |
+
model = model.to(device)
|
| 66 |
+
model.eval()
|
| 67 |
+
return model
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
class ColorizationPipeline:
|
| 71 |
+
"""Shared image colorization pipeline used by CLI/Gradio/Cog.
|
| 72 |
+
|
| 73 |
+
- input: BGR uint8 image (OpenCV)
|
| 74 |
+
- output: BGR uint8 image (OpenCV)
|
| 75 |
+
"""
|
| 76 |
+
|
| 77 |
+
def __init__(self, model, *, input_size: int = 512, device=None):
|
| 78 |
+
self.input_size = int(input_size)
|
| 79 |
+
if device is None:
|
| 80 |
+
try:
|
| 81 |
+
device = next(model.parameters()).device
|
| 82 |
+
except StopIteration:
|
| 83 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 84 |
+
self.device = device
|
| 85 |
+
self.model = model.to(self.device)
|
| 86 |
+
self.model.eval()
|
| 87 |
+
|
| 88 |
+
def process(self, img_bgr: np.ndarray) -> np.ndarray:
|
| 89 |
+
ctx = torch.inference_mode if hasattr(torch, "inference_mode") else torch.no_grad
|
| 90 |
+
with ctx():
|
| 91 |
+
if img_bgr is None:
|
| 92 |
+
raise ValueError("img is None (cv2.imread failed?)")
|
| 93 |
+
|
| 94 |
+
height, width = img_bgr.shape[:2]
|
| 95 |
+
|
| 96 |
+
img = (img_bgr / 255.0).astype(np.float32)
|
| 97 |
+
orig_l = cv2.cvtColor(img, cv2.COLOR_BGR2Lab)[:, :, :1] # (h, w, 1)
|
| 98 |
+
|
| 99 |
+
# resize rgb image -> lab -> get grey -> rgb
|
| 100 |
+
img_resized = cv2.resize(img, (self.input_size, self.input_size))
|
| 101 |
+
img_l = cv2.cvtColor(img_resized, cv2.COLOR_BGR2Lab)[:, :, :1]
|
| 102 |
+
img_gray_lab = np.concatenate(
|
| 103 |
+
(img_l, np.zeros_like(img_l), np.zeros_like(img_l)), axis=-1
|
| 104 |
+
)
|
| 105 |
+
img_gray_rgb = cv2.cvtColor(img_gray_lab, cv2.COLOR_LAB2RGB)
|
| 106 |
+
|
| 107 |
+
tensor_gray_rgb = (
|
| 108 |
+
torch.from_numpy(img_gray_rgb.transpose((2, 0, 1)))
|
| 109 |
+
.float()
|
| 110 |
+
.unsqueeze(0)
|
| 111 |
+
.to(self.device)
|
| 112 |
+
)
|
| 113 |
+
|
| 114 |
+
output_ab = self.model(tensor_gray_rgb).cpu() # (1, 2, input_size, input_size)
|
| 115 |
+
|
| 116 |
+
# resize ab -> concat original l -> bgr
|
| 117 |
+
output_ab_resized = (
|
| 118 |
+
F.interpolate(output_ab, size=(height, width))[0]
|
| 119 |
+
.float()
|
| 120 |
+
.numpy()
|
| 121 |
+
.transpose(1, 2, 0)
|
| 122 |
+
)
|
| 123 |
+
output_lab = np.concatenate((orig_l, output_ab_resized), axis=-1)
|
| 124 |
+
output_bgr = cv2.cvtColor(output_lab, cv2.COLOR_LAB2BGR)
|
| 125 |
+
|
| 126 |
+
output_img = (output_bgr * 255.0).round().astype(np.uint8)
|
| 127 |
+
return output_img
|
experiments/round5-20260927/status.json
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"status": "initializing"
|
| 3 |
+
}
|