User-2468 commited on
Commit
955113e
·
verified ·
1 Parent(s): 078af3e

Colorizer round5: baseline checkpoint and source before GPU work

Browse files
Files changed (37) hide show
  1. experiments/round5-20260927/SHA256SUMS.json +36 -3
  2. experiments/round5-20260927/best/README.md +16 -0
  3. experiments/round5-20260927/best/config.json +956 -0
  4. experiments/round5-20260927/best/model.safetensors +3 -0
  5. experiments/round5-20260927/selection.json +5 -0
  6. experiments/round5-20260927/source/PROTOCOL.md +36 -0
  7. experiments/round5-20260927/source/data.py +116 -0
  8. experiments/round5-20260927/source/inference.py +64 -0
  9. experiments/round5-20260927/source/metrics.py +35 -0
  10. experiments/round5-20260927/source/model.py +122 -0
  11. experiments/round5-20260927/source/persistence.py +42 -0
  12. experiments/round5-20260927/source/previous_manifest.json +0 -0
  13. experiments/round5-20260927/source/spatial.py +34 -0
  14. experiments/round5-20260927/source/train.py +259 -0
  15. experiments/round5-20260927/source/vendor_ddcolor/LICENSE +201 -0
  16. experiments/round5-20260927/source/vendor_ddcolor/basicsr/__init__.py +16 -0
  17. experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/__init__.py +41 -0
  18. experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/__init__.py +0 -0
  19. experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/convnext.py +206 -0
  20. experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/position_encoding.py +52 -0
  21. experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/transformer.py +368 -0
  22. experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/transformer_utils.py +192 -0
  23. experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/unet.py +208 -0
  24. experiments/round5-20260927/source/vendor_ddcolor/basicsr/archs/ddcolor_arch_utils/util.py +63 -0
  25. experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/__init__.py +37 -0
  26. experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/diffjpeg.py +515 -0
  27. experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/dist_util.py +82 -0
  28. experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/file_client.py +167 -0
  29. experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/img_process_util.py +83 -0
  30. experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/img_util.py +227 -0
  31. experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/logger.py +209 -0
  32. experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/misc.py +141 -0
  33. experiments/round5-20260927/source/vendor_ddcolor/basicsr/utils/registry.py +82 -0
  34. experiments/round5-20260927/source/vendor_ddcolor/ddcolor/__init__.py +9 -0
  35. experiments/round5-20260927/source/vendor_ddcolor/ddcolor/model.py +278 -0
  36. experiments/round5-20260927/source/vendor_ddcolor/ddcolor/pipeline.py +127 -0
  37. experiments/round5-20260927/status.json +3 -0
experiments/round5-20260927/SHA256SUMS.json CHANGED
@@ -1,5 +1,38 @@
1
  {
2
- "initial/config.json": "90299ee25ee3c3e645b362ddbd5a859b9e739c2a0a1c8a5e9ea79df171288bff",
3
- "initial/model.safetensors": "0e4c417375684a044860f8af3ac3a2fb44e1a5729254ca33abda758ea71aea6e",
4
- "source.zip": "e5e91e575dfe5a5d4b50c60a3e8f831be5c09693d2d6256044f31d853114618f"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ }