User-2468 commited on
Commit
4ae53b1
·
verified ·
1 Parent(s): c3df882

Colorizer round6: preflight passed before data preparation

Browse files
Files changed (25) hide show
  1. experiments/medium-critic-20260928/SHA256SUMS.json +26 -0
  2. experiments/medium-critic-20260928/data_and_eval.py +234 -0
  3. experiments/medium-critic-20260928/initial/README.md +5 -0
  4. experiments/medium-critic-20260928/initial/config.json +7 -0
  5. experiments/medium-critic-20260928/initial/model.safetensors +3 -0
  6. experiments/medium-critic-20260928/protocol.json +25 -0
  7. experiments/medium-critic-20260928/smoke.json +7 -0
  8. experiments/medium-critic-20260928/source/__pycache__/data.cpython-311.pyc +0 -0
  9. experiments/medium-critic-20260928/source/__pycache__/legacy_eval.cpython-311.pyc +0 -0
  10. experiments/medium-critic-20260928/source/__pycache__/metrics.cpython-311.pyc +0 -0
  11. experiments/medium-critic-20260928/source/__pycache__/model.cpython-311.pyc +0 -0
  12. experiments/medium-critic-20260928/source/__pycache__/objectives.cpython-311.pyc +0 -0
  13. experiments/medium-critic-20260928/source/__pycache__/persistence.cpython-311.pyc +0 -0
  14. experiments/medium-critic-20260928/source/__pycache__/semantic_model.cpython-311.pyc +0 -0
  15. experiments/medium-critic-20260928/source/__pycache__/spatial.cpython-311.pyc +0 -0
  16. experiments/medium-critic-20260928/source/data.py +116 -0
  17. experiments/medium-critic-20260928/source/legacy_eval.py +150 -0
  18. experiments/medium-critic-20260928/source/metrics.py +35 -0
  19. experiments/medium-critic-20260928/source/model.py +122 -0
  20. experiments/medium-critic-20260928/source/objectives.py +17 -0
  21. experiments/medium-critic-20260928/source/persistence.py +42 -0
  22. experiments/medium-critic-20260928/source/previous_manifest.json +0 -0
  23. experiments/medium-critic-20260928/source/semantic_model.py +94 -0
  24. experiments/medium-critic-20260928/source/spatial.py +34 -0
  25. experiments/medium-critic-20260928/train.py +261 -0
experiments/medium-critic-20260928/SHA256SUMS.json ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "data_and_eval.py": "019ea992052f0233954b86e06e31c85c74ca313e67a55ff7fb03248f8874aaf5",
3
+ "initial/README.md": "2508e51f14898f651c49adb184d1570b8adad3702313dee9bdac6fa0390d2bec",
4
+ "initial/config.json": "eaeb49b39f9116bfed0d954e81852dc93476298752787dba88a63c0eec97ccec",
5
+ "initial/model.safetensors": "ec1f27d74533adc83f7ab3639a091fc4d8738a434dafc7d172c7873c28a9e715",
6
+ "protocol.json": "6e08c212bb5cf106047fee69cca20b64d119055f8e871fdf62b349b70fd594d1",
7
+ "smoke.json": "390beb83fcbfe9b00bd1159092e05c066026172bf147bc382e20461cda45993e",
8
+ "source/__pycache__/data.cpython-311.pyc": "526112ecc37d2691497f7ef76b527361af7b49f1771f4dfb9f3c7c9c751b237c",
9
+ "source/__pycache__/legacy_eval.cpython-311.pyc": "25bc477f5fa66a8da0cf57fbc0c10f56047360f461056b863b1ea458680bb610",
10
+ "source/__pycache__/metrics.cpython-311.pyc": "de27195e97d9c9736879f6d71eb5003d39b87fe32921f6032dacbc27d877bde9",
11
+ "source/__pycache__/model.cpython-311.pyc": "bba35c57917ce7694d9cbd7453fdd732bf8bf188637ed2344744a29b46eb8704",
12
+ "source/__pycache__/objectives.cpython-311.pyc": "93c9a025eb4ef7b90b7ab5e9afb0f3a922151ea460426eb2d89c5a8ca953a7fe",
13
+ "source/__pycache__/persistence.cpython-311.pyc": "e2faff82ba28c5f1d11dbbffc980d44666b0ab24fe66bc31412ded9a9461e608",
14
+ "source/__pycache__/semantic_model.cpython-311.pyc": "2d4ee1e0549f5f1f4a2665b4ef5d5e169231e044fe47abeec4af7cddffc506e0",
15
+ "source/__pycache__/spatial.cpython-311.pyc": "191dfa051e73c71512eed82ea9b992242cc4ad4ce34568a88ea1bff0c602876a",
16
+ "source/data.py": "eb43b886fe45435841439c36646985f914386a48faa712fb6806accac2ad48e0",
17
+ "source/legacy_eval.py": "dd7e8bd5f961dfb6f1c2ce98fac41824c5009188e306000b195e922e4ed03e58",
18
+ "source/metrics.py": "af0ac996f6be0a351a55e372fab637a0d06740afe4a744cb80d0eae56ec97f66",
19
+ "source/model.py": "4cc57f82ffd6378bdf23a088ce7a1ed56c09e5673de3ee5232b8ee1cacc9be0a",
20
+ "source/objectives.py": "133e85b6f2914c0e3803cc26a26fa29b779a49036219dd245c2ecfe0ba411e20",
21
+ "source/persistence.py": "8b26dd8acdc135ecbfcdd30b7f7723feb6008ac820c083403da1e820bdc15d33",
22
+ "source/previous_manifest.json": "24bc081bf48f1109d9dfa217ae0d8593c0c06af6785a5cbfc5dc5d9e25461f5c",
23
+ "source/semantic_model.py": "b019358cbcaa214743cbd244d20eb6e59eeff0ab5264093ec0927be8529b25a2",
24
+ "source/spatial.py": "d13abe399ef049e21a6459a7003461afaae0c00e0c560262c6cf375de4c9884a",
25
+ "train.py": "44957be87bff1708f832ef5d8cb0fc8314ac7d727777c39b82bb50ea37fb8037"
26
+ }
experiments/medium-critic-20260928/data_and_eval.py ADDED
@@ -0,0 +1,234 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os,sys,json,random,time,hashlib,shutil,math,traceback,io
2
+ from pathlib import Path
3
+ import numpy as np
4
+ import torch
5
+ import torch.nn.functional as F
6
+ from torch.utils.data import Dataset,DataLoader
7
+ from huggingface_hub import hf_hub_download
8
+ import pyarrow as pa
9
+ import pyarrow.parquet as pq
10
+ from PIL import Image,ImageDraw
11
+ from skimage.color import rgb2lab,lab2rgb
12
+ from skimage.segmentation import slic
13
+ import cv2
14
+ import persistence
15
+ from semantic_model import load_semantic,save_semantic
16
+ from objectives import color_loss
17
+ from data import decode_image
18
+ import legacy_eval as ev
19
+ REPO='User-2468/mini-unet-colorizer'
20
+ REV='7b10485ee9d205ce2f7b660dfc6a90012d78b1cb'
21
+ OLD='experiments/round6-20260928'
22
+ OUT=Path('/work/results');OUT.mkdir(parents=True,exist_ok=True)
23
+ persistence.PREFIX='experiments/round7-20260928'
24
+ ev.OUT=OUT
25
+ START=time.monotonic()
26
+ SEED=709
27
+ STEPS=2000
28
+ DEVICE='cpu' if os.environ.get('SMOKE')=='1' else 'cuda'
29
+ ev.DEVICE=DEVICE
30
+ DURABLE=None
31
+ def write(name,obj):
32
+ p=OUT/name;p.parent.mkdir(parents=True,exist_ok=True);p.write_text(json.dumps(obj,indent=2));return obj
33
+ def deadline():
34
+ if time.monotonic()-START>9000:raise TimeoutError('Graceful deadline')
35
+ ev.deadline=deadline
36
+ def losses(p,t):
37
+ p=F.adaptive_avg_pool2d(p.float(),(64,64));t=t.float()
38
+ pixel=(p-t).abs().mean((1,2,3))/20
39
+ a=F.avg_pool2d(p,2).flatten(2);b=F.avg_pool2d(t,2).flatten(2)
40
+ angles=torch.arange(8,device=p.device)*math.pi/8
41
+ directions=torch.stack([angles.cos(),angles.sin()],1)
42
+ ap=torch.einsum('kc,bcn->bkn',directions,a).sort(-1).values
43
+ bp=torch.einsum('kc,bcn->bkn',directions,b).sort(-1).values
44
+ distribution=(ap-bp).abs().mean((1,2))/20
45
+ edge=((torch.diff(p,dim=-1)-torch.diff(t,dim=-1)).abs().mean((1,2,3))+(torch.diff(p,dim=-2)-torch.diff(t,dim=-2)).abs().mean((1,2,3)))/40
46
+ return pixel+.25*distribution+.1*edge
47
+ def select_loss(p,teacher,original,mode):
48
+ a=losses(p,teacher)
49
+ if mode=='teacher':return a.mean(),torch.zeros_like(a,dtype=torch.bool)
50
+ b=losses(p,original);choose=b<a
51
+ return torch.where(choose,b,a).mean(),choose
52
+ def smoke():
53
+ torch.set_num_threads(4);torch.manual_seed(SEED)
54
+ p=torch.randn(2,2,64,64,requires_grad=True);t=torch.randn_like(p)
55
+ assert torch.allclose(losses(p,t).mean(),color_loss(p,t)[0],atol=1e-6)
56
+ # Whole-image choice: exact red or blue is better than an average or a patchwork.
57
+ red=torch.zeros(1,2,64,64);red[:,0]=30
58
+ blue=-red;mix=red.clone();mix[:,:,:,32:]=blue[:,:,:,32:]
59
+ assert select_loss(red,red,blue,'set')[0]==0
60
+ assert select_loss(torch.zeros_like(red),red,blue,'set')[0]>0
61
+ assert select_loss(mix,red,blue,'set')[0]>0
62
+ m=load_semantic('/work/initial',DEVICE)
63
+ assert sum(p.numel() for p in m.parameters())==3994676
64
+ x=torch.rand(2,1,96,128,device=DEVICE)*2-1
65
+ pred=m(x);loss,_=select_loss(pred,t.to(DEVICE),-t.to(DEVICE),'set')
66
+ loss.backward()
67
+ assert torch.isfinite(loss) and all(torch.isfinite(p.grad).all() for p in m.parameters() if p.grad is not None)
68
+ save_semantic(m,OUT/'initial')
69
+ re=load_semantic(OUT/'initial',DEVICE)
70
+ with torch.no_grad():assert torch.allclose(m.eval()(x),re(x),atol=1e-5)
71
+ write('smoke.json',{'passed':True,'parameters':3994676,'loss_equivalence':True,'whole_image_choice':True,'finite_backward':True,'reload':True})
72
+ del m,re
73
+ class Photos(Dataset):
74
+ def __init__(self,table,ids,cache,archive):
75
+ self.table,self.ids,self.cache,self.archive=table,ids,cache,archive;self.epoch=0
76
+ def __len__(self):return len(self.ids)
77
+ def __getitem__(self,k):
78
+ im=decode_image(self.table,self.ids[k]);L,ab,gray=ev.gray_arrays(im)
79
+ # Independent per-example RNG; augmentation differences never alter image ordering/flips.
80
+ rng=np.random.default_rng(np.random.SeedSequence([SEED,self.epoch,k]))
81
+ target=np.array(self.cache[k],np.float32)
82
+ if self.archive and rng.random()<.5:
83
+ rgb=np.asarray(im.resize((256,256),Image.Resampling.BILINEAR),np.float32)/255
84
+ weights=rng.dirichlet(np.array([3.,5.,2.]))
85
+ g=(rgb*weights).sum(-1)
86
+ g=np.clip((g**rng.uniform(.7,1.5)-.5)*rng.uniform(.6,1.1)+.5+rng.uniform(-.08,.08),0,1)
87
+ if rng.random()<.5:g=cv2.GaussianBlur(g,(5,5),float(rng.uniform(.3,1.2)))
88
+ if rng.random()<.6:g=np.clip(g+rng.normal(0,rng.uniform(.003,.025),g.shape),0,1)
89
+ L=(rgb2lab(np.repeat(g[...,None],3,-1).astype(np.float32))[...,0][None]/50-1).astype(np.float32)
90
+ else:
91
+ if rng.random()<.5:L=np.clip(L*rng.uniform(.7,1.15)+rng.uniform(-.12,.12),-1,1)
92
+ flip=np.random.default_rng(np.random.SeedSequence([SEED+1,self.epoch,k])).random()<.5
93
+ if flip:L=L[...,::-1].copy();ab=ab[...,::-1].copy();target=target[...,::-1].copy()
94
+ original=cv2.resize(ab.transpose(1,2,0),(64,64),interpolation=cv2.INTER_AREA).transpose(2,0,1).copy()
95
+ return torch.from_numpy(L),torch.from_numpy(target),torch.from_numpy(original)
96
+ @torch.inference_mode()
97
+ def predict(m,L):
98
+ x=torch.from_numpy(L)[None].to(DEVICE);raw=m(x)
99
+ guided=ev.guided_chroma(x,raw,8)
100
+ return raw,guided
101
+ def toimage(L,ab):
102
+ rgb=lab2rgb(np.concatenate([(L[0]*50+50)[...,None],ab],-1))
103
+ return Image.fromarray(np.uint8(np.clip(rgb,0,1)*255))
104
+ @torch.inference_mode()
105
+ def review(m,table,ids,name):
106
+ m.eval();rows=[]
107
+ # Reference-derived surface regions for diagnostics only, never inference.
108
+ for i in ids:
109
+ L,target,_=ev.gray_arrays(decode_image(table,i))
110
+ raw,pred=predict(m,L);x=torch.from_numpy(L)[None].to(DEVICE)
111
+ resize=F.interpolate(x,(224,224),mode='bilinear',align_corners=False)
112
+ resized=F.interpolate(m(resize),(256,256),mode='bilinear',align_corners=False)
113
+ contrast=m((x*.95+.02).clamp(-1,1))
114
+ pab=pred[0].cpu().numpy().transpose(1,2,0)
115
+ rgb=np.asarray(decode_image(table,i).resize((256,256)),np.float32)/255
116
+ seg=slic(rgb,n_segments=100,compactness=10,start_label=0)
117
+ tab=target.transpose(1,2,0);variations=[]
118
+ for label in np.unique(seg):
119
+ mask=seg==label
120
+ if mask.sum()<64:continue
121
+ ground=tab[mask];center=np.median(ground,axis=0)
122
+ if np.percentile(np.linalg.norm(ground-center,axis=1),90)>5:continue
123
+ col=pab[mask];variations.append(float(np.mean(np.linalg.norm(col-np.median(col,axis=0),axis=1))))
124
+ rows.append({'id':int(i),'resize_ab_change':float((raw-resized).abs().mean()),'contrast_ab_change':float((raw-contrast).abs().mean()),'surface_chroma_variation':float(np.mean(variations)) if variations else None,'chroma_coverage':float((pred.norm(dim=1)>10).float().mean()),'mean_chroma':float(pred.norm(dim=1).mean())})
125
+ write(name+'_diagnostics.json',{'rows':rows,'summary':{k:float(np.mean([r[k] for r in rows if r[k] is not None])) for k in rows[0] if k!='id'},'warning':'Surface uniformity alone rewards grey collapse; never use alone to select model. Color differences from original are diagnostic only.'})
126
+ return rows
127
+ @torch.inference_mode()
128
+ def grid(models,table,ids,label):
129
+ order=np.random.default_rng(SEED+2).permutation(len(models))
130
+ names=list(models);key={chr(65+j):names[int(i)] for j,i in enumerate(order)}
131
+ write(label+'_key.json',key)
132
+ for page in range(0,min(len(ids),32),8):
133
+ chosen=ids[page:page+8];canvas=Image.new('RGB',(256*(len(models)+1),274*len(chosen)),'white');draw=ImageDraw.Draw(canvas)
134
+ for row,index in enumerate(chosen):
135
+ L,_,gray=ev.gray_arrays(decode_image(table,index));draw.text((2,row*274),f'{index}: gray | '+' | '.join(key),fill='black')
136
+ canvas.paste(Image.fromarray(np.uint8(gray*255)),(0,row*274+18))
137
+ for j,i in enumerate(order):
138
+ _,p=predict(models[names[int(i)]],L)
139
+ canvas.paste(toimage(L,p[0].cpu().numpy().transpose(1,2,0)),((j+1)*256,row*274+18))
140
+ canvas.save(OUT/f'{label}_{page//8}.jpg',quality=92)
141
+ def main():
142
+ global DURABLE
143
+ torch.set_num_threads(6);torch.manual_seed(SEED);random.seed(SEED);np.random.seed(SEED)
144
+ DURABLE=persistence.DurableRun(OUT)
145
+ shutil.copy2('/work/runner.py',OUT/'runner.py');shutil.copy2('/work/PROTOCOL.md',OUT/'PROTOCOL.md')
146
+ shutil.copytree('/work/source',OUT/'source',ignore=shutil.ignore_patterns('__pycache__'),dirs_exist_ok=True)
147
+ smoke();DURABLE.sync('round7 initial checkpoint and smoke verified')
148
+ if DEVICE=='cpu':return
149
+ torch.backends.cudnn.benchmark=True
150
+ tables=[]
151
+ 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'])]:
152
+ for f in files:tables.append(pq.read_table(hf_hub_download(repo,f,repo_type='dataset',revision=rev),columns=['image']))
153
+ table=pa.concat_tables(tables);manifest=json.loads(Path('/work/source/previous_manifest.json').read_text());ids=manifest['train']
154
+ cache=np.empty((len(ids),2,64,64),np.float16)
155
+ sums=json.loads(Path(hf_hub_download(REPO,OLD+'/SHA256SUMS.json',revision=REV)).read_text())
156
+ for start in range(0,len(ids),2048):
157
+ f=f'teacher_cache/{start:05d}.npz';path=Path(hf_hub_download(REPO,OLD+'/'+f,revision=REV))
158
+ assert hashlib.sha256(path.read_bytes()).hexdigest()==sums[f]
159
+ with np.load(path) as z:
160
+ assert z['indices'].tolist()==ids[start:start+len(z['indices'])]
161
+ cache[start:start+len(z['target'])]=z['target']
162
+ fresh=pq.read_table(hf_hub_download('detection-datasets/coco','default/val/0001.parquet',repo_type='dataset',revision='26ddc382fe75dfc2a0655b5977e296ea10efebce'),columns=['image'])
163
+ hashes={hashlib.sha256(table['image'][int(i)].as_py()['bytes']).digest() for i in ids}
164
+ freshids=[]
165
+ for i in np.random.default_rng(SEED).permutation(len(fresh)):
166
+ if hashlib.sha256(fresh['image'][int(i)].as_py()['bytes']).digest() not in hashes:freshids.append(int(i))
167
+ if len(freshids)==160:break
168
+ write('manifest.json',{'train':ids,'development':freshids[:80],'heldout':freshids[80:],'fresh_file':'default/val/0001.parquet','revision':'26ddc382fe75dfc2a0655b5977e296ea10efebce','prior_model_revision':REV,'upstream_pretraining_overlap_unknown':True,'near_duplicates_not_excluded':True})
169
+ DURABLE.sync('data cache verified and split fixed')
170
+ m=load_semantic('/work/initial',DEVICE)
171
+ write('initial_validation.json',ev.evaluate(m,table,manifest['validation']))
172
+ review(m,fresh,freshids[:80],'initial_development')
173
+ review(m,fresh,freshids[80:],'initial_heldout')
174
+ write('initial_heldout_reference_metrics.json',ev.evaluate(m,fresh,freshids[80:]))
175
+ del m
176
+ arms=[('teacher_clean','teacher',False),('set_clean','set',False),('teacher_archive','teacher',True),('set_archive','set',True)]
177
+ for name,mode,archive in arms:
178
+ torch.manual_seed(SEED);np.random.seed(SEED);random.seed(SEED)
179
+ m=load_semantic('/work/initial',DEVICE)
180
+ opt=torch.optim.AdamW([{'params':m.encoder.parameters(),'lr':5e-6},{'params':[p for n,p in m.named_parameters() if not n.startswith('encoder.')],'lr':5e-5}],weight_decay=.01)
181
+ sched=torch.optim.lr_scheduler.LambdaLR(opt,lambda step:min((step+1)/100,1)*(.2+.8*(1+math.cos(math.pi*min(step,STEPS)/STEPS))/2))
182
+ ds=Photos(table,ids,cache,archive)
183
+ loader=DataLoader(ds,batch_size=32,shuffle=True,num_workers=4,pin_memory=True,generator=torch.Generator().manual_seed(SEED))
184
+ it=iter(loader);history=[];choices=0;seen=0
185
+ for step in range(1,STEPS+1):
186
+ deadline();m.train()
187
+ for layer in m.encoder.modules():
188
+ if isinstance(layer,torch.nn.BatchNorm2d):layer.eval()
189
+ try:L,t,o=next(it)
190
+ except StopIteration:ds.epoch+=1;it=iter(loader);L,t,o=next(it)
191
+ L,t,o=L.to(DEVICE),t.to(DEVICE),o.to(DEVICE);opt.zero_grad(set_to_none=True)
192
+ with torch.autocast('cuda',dtype=torch.bfloat16):p=m(L)
193
+ loss,choice=select_loss(p,t,o,mode)
194
+ assert torch.isfinite(loss)
195
+ loss.backward();torch.nn.utils.clip_grad_norm_(m.parameters(),1,error_if_nonfinite=True);opt.step();sched.step()
196
+ choices+=int(choice.sum());seen+=len(L)
197
+ if step%100==0:
198
+ history.append({'step':step,'loss':float(loss),'original_choice_fraction':choices/seen})
199
+ print('STEP',name,step,round(float(loss),4),flush=True)
200
+ if step%500==0:
201
+ path=OUT/name/f'step_{step:05d}';save_semantic(m,path)
202
+ torch.save({'optimizer':opt.state_dict(),'scheduler':sched.state_dict(),'step':step,'epoch':ds.epoch,'torch_rng':torch.get_rng_state(),'cuda_rng':torch.cuda.get_rng_state_all(),'numpy_rng':np.random.get_state(),'python_rng':random.getstate()},path/'training_state.pt')
203
+ write(name+'/history.json',history);write('status.json',{'status':'training','arm':name,'step':step,'elapsed_seconds':time.monotonic()-START})
204
+ DURABLE.sync(name+' checkpoint '+str(step))
205
+ write(name+'_validation.json',ev.evaluate(m,table,manifest['validation']))
206
+ review(m,fresh,freshids[:80],name+'_development')
207
+ review(m,fresh,freshids[80:],name+'_heldout')
208
+ write(name+'_heldout_reference_metrics.json',ev.evaluate(m,fresh,freshids[80:]))
209
+ ev.probes(m,name)
210
+ DURABLE.sync(name+' complete diagnostics')
211
+ del m,opt,sched,it,loader;torch.cuda.empty_cache()
212
+ models={'round6':load_semantic('/work/initial',DEVICE)}
213
+ for name,_,_ in arms:models[name]=load_semantic(OUT/name/f'step_{STEPS:05d}',DEVICE)
214
+ grid(models,fresh,freshids[:80],'blinded_development')
215
+ grid(models,fresh,freshids[80:],'blinded_heldout')
216
+ # Familiar historical development probes; never claim these are a held-out archival benchmark.
217
+ import requests
218
+ sources={'migrant_mother':'https://cdn.loc.gov/service/pnp/ppmsca/50200/50236v.jpg','power_house_mechanic':'https://www.archives.gov/exhibits/picturing_the_century/images/port_hine_022_v36.jpg'}
219
+ for name,url in sources.items():
220
+ try:
221
+ response=requests.get(url,timeout=45);response.raise_for_status()
222
+ im=Image.open(io.BytesIO(response.content)).convert('RGB')
223
+ tb=pa.table({'image':[{'bytes':response.content,'path':None}]})
224
+ grid(models,tb,[0],name)
225
+ write(name+'_source.json',{'url':url,'sha256':hashlib.sha256(response.content).hexdigest(),'use':'previously seen archival development probe; original colors unknown'})
226
+ except Exception as e:write(name+'_error.json',{'type':type(e).__name__,'message':str(e)})
227
+ write('status.json',{'status':'completed','elapsed_seconds':time.monotonic()-START,'production_approved':False,'next':'Review blinded grids for coherent plausible color and gray/sepia collapse. This pilot does not test a stochastic conditioned decoder.'})
228
+ DURABLE.sync('round7 four-arm experiment completed')
229
+ if __name__=='__main__':
230
+ try:main()
231
+ except BaseException as e:
232
+ write('failure.json',{'type':type(e).__name__,'traceback':traceback.format_exc().replace(os.environ.get('HF_TOKEN','__EMPTY__'),'[REDACTED]')})
233
+ if DURABLE:DURABLE.sync('round7 failure report and latest saved checkpoints')
234
+ raise
experiments/medium-critic-20260928/initial/README.md ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ # Experimental semantic colorizer
2
+
3
+ Head: palette. Parameters: 3994676.
4
+
5
+ Not approved for production. Requires semantic_model.py; incompatible with the old U-Net loader. Input is Lab lightness normalized to [-1,1]; output is Lab ab. See the run protocol, provenance, selection and visual comparisons. Predictions are plausible colors, not recovered historical truth.
experiments/medium-critic-20260928/initial/config.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "architecture": "SemanticColorizer",
3
+ "head": "palette",
4
+ "width": 128,
5
+ "queries": 16,
6
+ "format_version": 1
7
+ }
experiments/medium-critic-20260928/initial/model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:ec1f27d74533adc83f7ab3639a091fc4d8738a434dafc7d172c7873c28a9e715
3
+ size 16112432
experiments/medium-critic-20260928/protocol.json ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "generator_parameters": 3994676,
3
+ "critic_parameters": 11623043,
4
+ "base_revision": "1a9eb8af2754ad2329a24cfe50d388cb559441d0",
5
+ "steps": 8000,
6
+ "batch_size": 16,
7
+ "seed": 92815,
8
+ "warmup_critic_updates": 500,
9
+ "generator_objective": "ONLY non-saturating conditional adversarial logistic loss, equally weighted global and two patch heads; no reconstruction, teacher, perceptual, feature matching, hue-target, or smoothing loss",
10
+ "critic": "ImageNet pretrained ResNet18, RGB+L conditioning, global plus layer2/layer3 patch heads; all weights trainable with frozen BatchNorm statistics",
11
+ "regularization": "D-only lazy R1 gamma=10 every16 updates on8 samples; adaptive translation augmentation target real-positive fraction0.6; identical transforms for real/fake; EMA0.999",
12
+ "generator_optimizer": "Adam encoder1e-6 decoder2e-5; warmup500 then cosine floor0.2, no weight decay",
13
+ "critic_optimizer": "Adam pretrained2e-5 newheads1e-4; betas0,0.99",
14
+ "data": "up to60000 real train photographs; pinned COCO/Imagenette; original fold excluded; exact-byte train/eval duplicates removed; no teacher outputs",
15
+ "checkpoint_interval": 250,
16
+ "evaluation_interval": 1000,
17
+ "hard_timeout_hours": 1.1,
18
+ "graceful_training_budget_hours": 0.85,
19
+ "release_policy": "experiment folder only; compare grids with v3 and previous 27M candidate; no automatic replacement of v3; critic score not used to approve production",
20
+ "research": [
21
+ "https://arxiv.org/abs/2006.06676",
22
+ "https://arxiv.org/abs/1703.10593"
23
+ ],
24
+ "limitations": "deterministic colouriser; one seed; adversarial loss can still collapse or exploit critic; grayscale is sometimes correct; no universal quality guarantee"
25
+ }
experiments/medium-critic-20260928/smoke.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "backward_finite": true,
3
+ "R1_second_derivative": true,
4
+ "checkpoint_roundtrip": true,
5
+ "generator_parameters": 3994676,
6
+ "critic_parameters": 11623043
7
+ }
experiments/medium-critic-20260928/source/__pycache__/data.cpython-311.pyc ADDED
Binary file (13.7 kB). View file
 
experiments/medium-critic-20260928/source/__pycache__/legacy_eval.cpython-311.pyc ADDED
Binary file (20 kB). View file
 
experiments/medium-critic-20260928/source/__pycache__/metrics.cpython-311.pyc ADDED
Binary file (4.04 kB). View file
 
experiments/medium-critic-20260928/source/__pycache__/model.cpython-311.pyc ADDED
Binary file (10.3 kB). View file
 
experiments/medium-critic-20260928/source/__pycache__/objectives.cpython-311.pyc ADDED
Binary file (2.86 kB). View file
 
experiments/medium-critic-20260928/source/__pycache__/persistence.cpython-311.pyc ADDED
Binary file (5.95 kB). View file
 
experiments/medium-critic-20260928/source/__pycache__/semantic_model.cpython-311.pyc ADDED
Binary file (15.9 kB). View file
 
experiments/medium-critic-20260928/source/__pycache__/spatial.cpython-311.pyc ADDED
Binary file (2.89 kB). View file
 
experiments/medium-critic-20260928/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/medium-critic-20260928/source/legacy_eval.py ADDED
@@ -0,0 +1,150 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ OUT=Path('round6');OUT.mkdir(exist_ok=True)
21
+ DURABLE=None
22
+ START=time.monotonic();DEADLINE=START+3.5*3600
23
+ DEVICE='cuda';SEED=409;STEPS=3000
24
+ BASE_REV='704fa80d792c3d759db91daa00b2dcfe6f0f6412'
25
+ BASE_HASH='0e4c417375684a044860f8af3ac3a2fb44e1a5729254ca33abda758ea71aea6e'
26
+ TEACHER_ID='piddnad/ddcolor_artistic'
27
+ torch.set_num_threads(6)
28
+
29
+ def deadline():
30
+ if time.monotonic()>DEADLINE:raise TimeoutError('Graceful export before remote hard timeout')
31
+
32
+ def gray_arrays(im,style='gray'):
33
+ im=im.resize((256,256),Image.Resampling.BILINEAR)
34
+ rgb=np.asarray(im,np.float32)/255
35
+ lab=rgb2lab(rgb).astype(np.float32)
36
+ if style=='film':
37
+ g=(rgb*np.array([.42,.45,.13],np.float32)).sum(-1)
38
+ g=np.clip((g**1.15-.5)*.8+.5,0,1)
39
+ else:g=np.asarray(im.convert('L'),np.float32)/255
40
+ # Feed the same neutral sRGB input to teacher and student.
41
+ gray=np.repeat(g[...,None],3,-1).astype(np.float32)
42
+ L=rgb2lab(gray)[...,0].astype(np.float32)
43
+ return (L[None]/50-1).copy(),lab[...,1:].transpose(2,0,1).copy(),gray
44
+
45
+ class Photos(Dataset):
46
+ def __init__(self,table,ids,bins=None,cache=None,augment=False,style='gray'):
47
+ self.table,self.ids,self.bins,self.cache,self.augment,self.style=table,ids,bins,cache,augment,style
48
+ def __len__(self):return len(self.ids)
49
+ def __getitem__(self,k):
50
+ L,ab,gray=gray_arrays(decode_image(self.table,self.ids[k]),self.style)
51
+ teach=np.array(self.cache[k],np.float32) if self.cache is not None else np.zeros((2,64,64),np.float32)
52
+ if self.augment:
53
+ if random.random()<.5:L=L[...,::-1].copy();ab=ab[...,::-1].copy();teach=teach[...,::-1].copy()
54
+ if random.random()<.3:L=np.clip(L*random.uniform(.85,1.15)+random.uniform(-.08,.08),-1,1)
55
+ if self.bins is not None:
56
+ 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)
57
+ return torch.from_numpy(L),torch.from_numpy(ab),torch.from_numpy(gray.transpose(2,0,1)),self.ids[k]
58
+
59
+ def patch_excess(p,t):
60
+ """Excess color differences across 4–64px regions, in target-flat areas."""
61
+ values=[]
62
+ for scale in [4,16]:
63
+ a=F.avg_pool2d(p,scale);b=F.avg_pool2d(t,scale)
64
+ for offset in [1,4]:
65
+ for dim in [-1,-2]:
66
+ x=a.narrow(dim,offset,a.shape[dim]-offset)-a.narrow(dim,0,a.shape[dim]-offset)
67
+ y=b.narrow(dim,offset,b.shape[dim]-offset)-b.narrow(dim,0,b.shape[dim]-offset)
68
+ dx=x.norm(dim=1);dy=y.norm(dim=1);mask=(dy<3).float()
69
+ values.append(((dx-dy-1).relu()*mask).sum((1,2))/mask.sum((1,2)).clamp_min(1))
70
+ return sum(values)/len(values)
71
+
72
+ def structural_loss(p,t):
73
+ terms=[]
74
+ for scale in [4,16]:
75
+ a=F.avg_pool2d(p,scale);b=F.avg_pool2d(t,scale)
76
+ for offset in [1,4]:
77
+ for dim in [-1,-2]:
78
+ x=a.narrow(dim,offset,a.shape[dim]-offset)-a.narrow(dim,0,a.shape[dim]-offset)
79
+ y=b.narrow(dim,offset,b.shape[dim]-offset)-b.narrow(dim,0,b.shape[dim]-offset)
80
+ terms.append(F.smooth_l1_loss(x/10,y/10,beta=.3))
81
+ return sum(terms)/len(terms)
82
+
83
+ def additional_losses(p,t,teacher):
84
+ small=F.avg_pool2d(p,4);target=F.avg_pool2d(t,4)
85
+ # Teacher is fallible; reduce guidance when it conflicts strongly with ground truth.
86
+ confidence=torch.exp(-(teacher-target).norm(dim=1,keepdim=True)/30).detach()
87
+ kd=(F.smooth_l1_loss(small/20,teacher/20,beta=.5,reduction='none')*confidence).mean()
88
+ neutral=(t.norm(dim=1)<3).float()
89
+ neutral_loss=(p.norm(dim=1)*neutral).sum()/neutral.sum().clamp_min(1)/20
90
+ return structural_loss(p,t),kd,neutral_loss
91
+
92
+ def measure(L,p,t,raw=None):
93
+ r=per_image_metrics(L,p,t);C=p.norm(dim=1);T=t.norm(dim=1)
94
+ for name,mask,event in [('missed_color',T>12,C<5),('neutral_spill',T<3,C>10)]:
95
+ r[name]=(mask&event).sum((1,2))/mask.sum((1,2)).clamp_min(1)
96
+ r['color_coverage']=(C>10).float().mean((1,2));r['patch_excess']=patch_excess(p,t)
97
+ r['raw_patch_excess']=patch_excess(p if raw is None else raw,t)
98
+ return r
99
+
100
+ @torch.inference_mode()
101
+ def teacher_predict(teacher,gray):
102
+ gray=F.interpolate(gray.to(DEVICE),size=(512,512),mode='bilinear',align_corners=False)
103
+ # Full precision is deliberate: spectral normalization / attention are not assumed BF16-safe.
104
+ out=teacher(gray).float()
105
+ assert out.shape[1]==2 and torch.isfinite(out).all()
106
+ return F.interpolate(out,size=(256,256),mode='bilinear',align_corners=False)
107
+
108
+ @torch.inference_mode()
109
+ def evaluate(model,table,ids,is_teacher=False):
110
+ result={};model.eval()
111
+ for style in ['gray','film']:
112
+ rows=[]
113
+ for L,t,gray,indices in DataLoader(Photos(table,ids,style=style),batch_size=2 if is_teacher else 8,num_workers=4):
114
+ deadline();L,t=L.to(DEVICE),t.to(DEVICE)
115
+ raw=teacher_predict(model,gray) if is_teacher else model.decode(model(L),.38)
116
+ pred=raw if is_teacher else guided_chroma(L,raw,8)
117
+ metrics=measure(L,pred,t,raw)
118
+ for j,index in enumerate(indices):rows.append({'index':int(index)}|{k:float(v[j]) for k,v in metrics.items()})
119
+ keep=[x for x in rows if x['target_chroma']>=5]
120
+ 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}
121
+ return result
122
+
123
+ def selection(result,baseline):
124
+ scores=[];eligible=True
125
+ for style in ['gray','film']:
126
+ s=result[style]['summary'];b=baseline[style]['summary']
127
+ eligible &= (s['ab_error']<=b['ab_error']*1.05 and s['neutral_spill']<=b['neutral_spill']+.01
128
+ and s['missed_color']<=b['missed_color']+.01 and s['color_coverage']>=.95*b['color_coverage']
129
+ and s['patch_excess']<.85*b['patch_excess'] and s['raw_patch_excess']<.9*b['raw_patch_excess'])
130
+ scores.append(s['patch_excess']/max(b['patch_excess'],1e-6)+.3*s['neutral_spill']+.2*s['ab_error']/b['ab_error'])
131
+ return float(np.mean(scores)),bool(eligible)
132
+
133
+ @torch.inference_mode()
134
+ def probes(model,tag,teacher=False):
135
+ names=['astronaut','coffee','chelsea','rocket','camera','coins','moon'];images=[]
136
+ for name in names:
137
+ original=Image.fromarray(getattr(sample_data,name)()).convert('RGB')
138
+ L,_,gray=gray_arrays(original);x=torch.from_numpy(L)[None].to(DEVICE)
139
+ 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)
140
+ lab=np.concatenate([(L[0]*50+50)[...,None],ab[0].cpu().numpy().transpose(1,2,0)],-1)
141
+ rgb=np.uint8(np.clip(lab2rgb(lab),0,1)*255)
142
+ images.append((name,Image.fromarray(np.uint8(gray*255)),Image.fromarray(rgb)))
143
+ canvas=Image.new('RGB',(512,len(names)*280),'white');d=ImageDraw.Draw(canvas)
144
+ for i,(name,gray,col) in enumerate(images):
145
+ 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')
146
+ canvas.save(OUT/(tag+'.jpg'),quality=90)
147
+
148
+ def write(name,obj):
149
+ (OUT/name).write_text(json.dumps(obj,indent=2));return obj
150
+
experiments/medium-critic-20260928/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/medium-critic-20260928/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/medium-critic-20260928/source/objectives.py ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ import torch
3
+ import torch.nn.functional as F
4
+
5
+ def color_loss(pred,target):
6
+ """Fit one coherent target; distribution matching discourages gray averages."""
7
+ pred=F.adaptive_avg_pool2d(pred,(64,64)).float();target=target.float()
8
+ pixel=F.l1_loss(pred/20,target/20)
9
+ p=F.avg_pool2d(pred,2).flatten(2);t=F.avg_pool2d(target,2).flatten(2)
10
+ angle=torch.arange(8,device=p.device,dtype=p.dtype)*math.pi/8
11
+ directions=torch.stack([angle.cos(),angle.sin()],1)
12
+ a=torch.einsum('kc,bcn->bkn',directions,p).sort(-1).values
13
+ b=torch.einsum('kc,bcn->bkn',directions,t).sort(-1).values
14
+ distribution=F.l1_loss(a/20,b/20)
15
+ # Match boundaries in the teacher target, never minimize gradients toward zero.
16
+ edge=(F.l1_loss((pred[:,:,:,1:]-pred[:,:,:,:-1])/20,(target[:,:,:,1:]-target[:,:,:,:-1])/20)+F.l1_loss((pred[:,:,1:]-pred[:,:,:-1])/20,(target[:,:,1:]-target[:,:,:-1])/20))/2
17
+ return pixel+.25*distribution+.1*edge,{'pixel':pixel,'distribution':distribution,'edge':edge}
experiments/medium-critic-20260928/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/round6-20260928'
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 round6: '+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/medium-critic-20260928/source/previous_manifest.json ADDED
The diff for this file is too large to render. See raw diff
 
experiments/medium-critic-20260928/source/semantic_model.py ADDED
@@ -0,0 +1,94 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Compact pretrained semantic colorizer; dense and shared-palette variants."""
2
+ import json
3
+ from pathlib import Path
4
+ import torch
5
+ from torch import nn
6
+ import torch.nn.functional as F
7
+ from torchvision.models import mobilenet_v3_large, MobileNet_V3_Large_Weights
8
+ from safetensors.torch import save_file, load_file
9
+
10
+ class QueryBlock(nn.Module):
11
+ def __init__(self,d=96):
12
+ super().__init__()
13
+ self.self_attn=nn.MultiheadAttention(d,4,batch_first=True,dropout=0)
14
+ self.cross_attn=nn.MultiheadAttention(d,4,batch_first=True,dropout=0)
15
+ self.norms=nn.ModuleList([nn.LayerNorm(d) for _ in range(3)])
16
+ self.ff=nn.Sequential(nn.Linear(d,2*d),nn.GELU(),nn.Linear(2*d,d))
17
+ def forward(self,q,memory):
18
+ x=self.norms[0](q);q=q+self.self_attn(x,x,x,need_weights=False)[0]
19
+ x=self.norms[1](q);q=q+self.cross_attn(x,memory,memory,need_weights=False)[0]
20
+ return q+self.ff(self.norms[2](q))
21
+
22
+ def refine(d):
23
+ return nn.Sequential(nn.Conv2d(d,d,3,padding=1,bias=False),nn.GroupNorm(8,d),nn.SiLU())
24
+
25
+ class SemanticColorizer(nn.Module):
26
+ def __init__(self,head='palette',pretrained=False,width=128,queries=16):
27
+ super().__init__()
28
+ if head not in ['palette','dense']:raise ValueError(head)
29
+ self.config={'architecture':'SemanticColorizer','head':head,'width':width,'queries':queries,'format_version':1}
30
+ self.encoder=mobilenet_v3_large(weights=MobileNet_V3_Large_Weights.IMAGENET1K_V2 if pretrained else None,progress=False).features
31
+ self.lateral=nn.ModuleList([nn.Conv2d(c,width,1) for c in [24,40,112,960]])
32
+ self.refine=nn.ModuleList([refine(width) for _ in range(3)])
33
+ self.register_buffer('rgb_mean',torch.tensor([.485,.456,.406]).view(1,3,1,1))
34
+ self.register_buffer('rgb_std',torch.tensor([.229,.224,.225]).view(1,3,1,1))
35
+ if head=='palette':
36
+ self.queries=nn.Parameter(torch.randn(queries,width)*.2)
37
+ self.query_blocks=nn.ModuleList([QueryBlock(width) for _ in range(2)])
38
+ self.memory_norm=nn.LayerNorm(width)
39
+ self.query_norm=nn.LayerNorm(width)
40
+ self.pixel=nn.Conv2d(width,width,1)
41
+ self.palette=nn.Sequential(nn.Linear(width,width),nn.GELU(),nn.Linear(width,2))
42
+ self.residual=nn.Conv2d(width,2,1)
43
+ nn.init.normal_(self.palette[-1].weight,std=.01);nn.init.zeros_(self.palette[-1].bias)
44
+ nn.init.zeros_(self.residual.weight);nn.init.zeros_(self.residual.bias)
45
+ else:
46
+ self.dense=nn.Sequential(refine(width),nn.Conv2d(width,2,1))
47
+ nn.init.normal_(self.dense[-1].weight,std=.01);nn.init.zeros_(self.dense[-1].bias)
48
+ count=sum(p.numel() for p in self.parameters())
49
+ if count>=4_000_000:raise ValueError(f'Parameter budget exceeded: {count}')
50
+
51
+ @staticmethod
52
+ def neutral_rgb(L):
53
+ light=(L.float()*50+50).clamp(0,100)
54
+ y=torch.where(light>8,((light+16)/116)**3,light/903.296296)
55
+ g=torch.where(y<=.0031308,12.92*y,1.055*y.clamp_min(1e-8).pow(1/2.4)-.055)
56
+ return g.expand(-1,3,-1,-1)
57
+
58
+ def forward(self,L):
59
+ h,w=L.shape[-2:]
60
+ x=F.pad(self.neutral_rgb(L),(0,(-w)%32,0,(-h)%32),mode='replicate')
61
+ x=(x-self.rgb_mean)/self.rgb_std
62
+ features=[]
63
+ for i,layer in enumerate(self.encoder):
64
+ x=layer(x)
65
+ if i in [3,6,12,16]:features.append(x)
66
+ projected=[layer(f) for layer,f in zip(self.lateral,features)]
67
+ x=projected[-1]
68
+ for i in range(2,-1,-1):
69
+ x=self.refine[2-i](F.interpolate(x,size=projected[i].shape[-2:],mode='bilinear',align_corners=False)+projected[i])
70
+ if self.config['head']=='palette':
71
+ memory=torch.cat([F.adaptive_avg_pool2d(f,(8,8)).flatten(2).transpose(1,2) for f in projected[1:]],1)
72
+ memory=self.memory_norm(memory)
73
+ q=self.queries[None].expand(L.shape[0],-1,-1)
74
+ for block in self.query_blocks:q=block(q,memory)
75
+ q=self.query_norm(q)
76
+ palette=80*torch.tanh(self.palette(q))
77
+ masks=torch.einsum('bqd,bdhw->bqhw',q,self.pixel(x))/(self.config['width']**.5)
78
+ weights=F.softmax(masks.float(),dim=1)
79
+ ab=torch.einsum('bqhw,bqc->bchw',weights,palette.float())+2*torch.tanh(self.residual(x).float())
80
+ else:ab=80*torch.tanh(self.dense(x).float())
81
+ return F.interpolate(ab,size=(x.shape[-2]*4,x.shape[-1]*4),mode='bilinear',align_corners=False)[...,:h,:w]
82
+ def decode(self,z,temperature=.38):return z.float()
83
+
84
+ def save_semantic(model,path):
85
+ path=Path(path);path.mkdir(parents=True,exist_ok=True)
86
+ save_file({k:v.detach().cpu().contiguous() for k,v in model.state_dict().items()},str(path/'model.safetensors'))
87
+ (path/'config.json').write_text(json.dumps(model.config,indent=2))
88
+ (path/'README.md').write_text('# Experimental semantic colorizer\n\nHead: '+model.config['head']+'. Parameters: '+str(sum(p.numel() for p in model.parameters()))+'.\n\nNot approved for production. Requires semantic_model.py; incompatible with the old U-Net loader. Input is Lab lightness normalized to [-1,1]; output is Lab ab. See the run protocol, provenance, selection and visual comparisons. Predictions are plausible colors, not recovered historical truth.\n')
89
+
90
+ def load_semantic(path,device='cpu'):
91
+ path=Path(path);cfg=json.loads((path/'config.json').read_text())
92
+ model=SemanticColorizer(**{k:cfg[k] for k in ['head','width','queries']})
93
+ model.load_state_dict(load_file(str(path/'model.safetensors')),strict=True)
94
+ return model.to(device).eval()
experiments/medium-critic-20260928/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/medium-critic-20260928/train.py ADDED
@@ -0,0 +1,261 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import os,sys,time,json,hashlib,shutil,random,copy,math,traceback,io
3
+ from pathlib import Path
4
+ import numpy as np
5
+ import torch
6
+ from torch import nn
7
+ import torch.nn.functional as F
8
+ from torch.utils.data import Dataset,DataLoader
9
+ from torchvision.models import resnet18,ResNet18_Weights
10
+ from torchvision.transforms import RandomResizedCrop
11
+ from torchvision.transforms import functional as TF
12
+ from PIL import Image
13
+ from skimage.color import rgb2lab
14
+ import pyarrow as pa
15
+ import pyarrow.parquet as pq
16
+ from huggingface_hub import HfApi,hf_hub_download
17
+ REPO='User-2468/mini-unet-colorizer'
18
+ BASE='1a9eb8af2754ad2329a24cfe50d388cb559441d0'
19
+ PREFIX='experiments/medium-critic-20260928'
20
+ WORK=Path('/work');WORK.mkdir(exist_ok=True);os.chdir(WORK)
21
+ START=time.monotonic();MAX_SECONDS=.85*3600
22
+ api=HfApi(token=os.environ['HF_TOKEN']);assert api.whoami()['name']=='User-2468'
23
+ src='experiments/final-20260928/realism_low_v2'
24
+ source=WORK/'source';source.mkdir(exist_ok=True)
25
+ sums=json.loads(Path(hf_hub_download(REPO,src+'/SHA256SUMS.json',revision=BASE)).read_text())
26
+ for f in ['data.py','legacy_eval.py','metrics.py','model.py','objectives.py','persistence.py','previous_manifest.json','semantic_model.py','spatial.py']:
27
+ p=Path(hf_hub_download(REPO,src+'/source/'+f,revision=BASE))
28
+ assert hashlib.sha256(p.read_bytes()).hexdigest()==sums['source/'+f]
29
+ shutil.copy2(p,source/f)
30
+ shutil.copy2(hf_hub_download(REPO,src+'/data_and_eval.py',revision=BASE),WORK/'r7.py')
31
+ sys.path.insert(0,str(source));sys.path.insert(0,str(WORK))
32
+ import r7 as r
33
+ from semantic_model import load_semantic,save_semantic
34
+ from data import decode_image
35
+ r.persistence.PREFIX=PREFIX
36
+ OUT=r.OUT;durable=r.persistence.DurableRun(OUT)
37
+ initial=WORK/'initial';initial.mkdir(exist_ok=True)
38
+ for f in ['model.safetensors','config.json']:
39
+ shutil.copy2(hf_hub_download(REPO,f,revision=BASE),initial/f)
40
+ r.ev.deadline=lambda:None
41
+ SEED=92815;STEPS=8000;BATCH=16
42
+ torch.set_num_threads(8);torch.manual_seed(SEED);random.seed(SEED);np.random.seed(SEED)
43
+ torch.backends.cudnn.benchmark=True
44
+ torch.backends.cuda.matmul.allow_tf32=True
45
+ torch.backends.cudnn.allow_tf32=True
46
+
47
+ def lab_rgb(L,ab):
48
+ # Equal chroma bandwidth in real/fake critic inputs prevents a trivial resize shortcut.
49
+ ab=F.interpolate(F.avg_pool2d(ab.float(),4),L.shape[-2:],mode='bilinear',align_corners=False)
50
+ light=L.float()*50+50
51
+ y=(light+16)/116;x=y+ab[:,0:1]/500;z=y-ab[:,1:2]/200
52
+ xyz=torch.cat([x,y,z],1)
53
+ xyz=torch.where(xyz>6/29,xyz**3,(xyz-4/29)*3*(6/29)**2)
54
+ xyz=xyz*xyz.new_tensor([.95047,1,1.08883])[None,:,None,None]
55
+ matrix=xyz.new_tensor([[3.2404542,-1.5371385,-.4985314],[-.969266,1.8760108,.041556],[.0556434,-.2040259,1.0572252]])
56
+ lin=torch.einsum('ij,bjhw->bihw',matrix,xyz)
57
+ rgb=torch.where(lin>.0031308,1.055*lin.clamp_min(1e-8).pow(1/2.4)-.055,12.92*lin)
58
+ return rgb.clamp(0,1)
59
+
60
+ def pair(L,ab):
61
+ rgb=lab_rgb(L,ab)
62
+ return torch.cat([(rgb-rgb.new_tensor([.485,.456,.406])[None,:,None,None])/rgb.new_tensor([.229,.224,.225])[None,:,None,None],L.float()],1)
63
+
64
+ class Critic(nn.Module):
65
+ def __init__(self):
66
+ super().__init__()
67
+ self.net=resnet18(weights=ResNet18_Weights.IMAGENET1K_V1,progress=False)
68
+ old=self.net.conv1
69
+ self.net.conv1=nn.Conv2d(4,64,7,2,3,bias=False)
70
+ with torch.no_grad():
71
+ self.net.conv1.weight[:,:3].copy_(old.weight);self.net.conv1.weight[:,3:].zero_()
72
+ self.net.fc=nn.Linear(512,1)
73
+ self.patch2=nn.Sequential(nn.Conv2d(128,128,3,padding=1),nn.LeakyReLU(.2),nn.Conv2d(128,1,1))
74
+ self.patch3=nn.Sequential(nn.Conv2d(256,128,3,padding=1),nn.LeakyReLU(.2),nn.Conv2d(128,1,1))
75
+ def train(self,mode=True):
76
+ super().train(mode)
77
+ # Fixed population statistics prevent batch composition and tiny R1 batches becoming shortcuts.
78
+ for mod in self.modules():
79
+ if isinstance(mod,nn.BatchNorm2d):mod.eval()
80
+ return self
81
+ def forward(self,x):
82
+ n=self.net
83
+ x=n.maxpool(n.relu(n.bn1(n.conv1(x))))
84
+ x=n.layer1(x);x=n.layer2(x);p2=self.patch2(x).flatten(1)
85
+ x=n.layer3(x);p3=self.patch3(x).flatten(1)
86
+ x=n.layer4(x);g=n.fc(n.avgpool(x).flatten(1))
87
+ return [g,p2,p3]
88
+
89
+ def loss_real(scores):return sum(F.softplus(-s).mean() for s in scores)/len(scores)
90
+ def loss_fake(scores):return sum(F.softplus(s).mean() for s in scores)/len(scores)
91
+ def score_mean(scores):return torch.stack([s.mean(1) for s in scores]).mean(0)
92
+
93
+ def augment(x,p):
94
+ # Geometric only; no colour jitter that could teach the critic to accept wrong chroma.
95
+ if random.random()<p:
96
+ h,w=x.shape[-2:];dy=random.randint(-16,16);dx=random.randint(-16,16)
97
+ x=F.pad(x,(16,16,16,16),mode='reflect')[:,:,16+dy:16+dy+h,16+dx:16+dx+w]
98
+ return x
99
+
100
+ class Photos(Dataset):
101
+ def __init__(self,table,ids):self.table,self.ids=table,ids
102
+ def __len__(self):return len(self.ids)
103
+ def __getitem__(self,k):
104
+ im=decode_image(self.table,self.ids[k]).convert('RGB')
105
+ i,j,h,w=RandomResizedCrop.get_params(im,scale=(.65,1.),ratio=(.85,1.18))
106
+ im=TF.resized_crop(im,i,j,h,w,[256,256],interpolation=TF.InterpolationMode.BILINEAR)
107
+ if random.random()<.5:im=TF.hflip(im)
108
+ rgb=np.asarray(im,np.float32)/255
109
+ lab=rgb2lab(rgb).astype(np.float32)
110
+ # Half neutral L, half standard photographic RGB grayscale; equal real/fake L.
111
+ if random.random()<.5:L=lab[...,0]
112
+ else:
113
+ g=np.asarray(im.convert('L'),np.float32)/255
114
+ L=rgb2lab(np.repeat(g[...,None],3,-1))[...,0].astype(np.float32)
115
+ return torch.from_numpy((L[None]/50-1).copy()),torch.from_numpy(lab[...,1:].transpose(2,0,1).copy())
116
+
117
+ def save(step,model,ema,critic,opt,dopt,history,p,phase):
118
+ save_semantic(model,OUT/f'step_{step:05d}'/'raw')
119
+ save_semantic(ema,OUT/f'step_{step:05d}'/'ema')
120
+ torch.save({'step':step,'generator':model.state_dict(),'ema':ema.state_dict(),'critic':critic.state_dict(),'optimizer':opt.state_dict(),'critic_optimizer':dopt.state_dict(),'augmentation_probability':p,'torch_rng':torch.get_rng_state(),'cuda_rng':torch.cuda.get_rng_state_all(),'numpy_rng':np.random.get_state(),'python_rng':random.getstate(),'note':'Resume optimizer/model state; loader position is not serialized.'},OUT/'latest_training_state.pt')
121
+ r.write('history.json',history)
122
+ r.write('status.json',{'status':phase,'step':step,'elapsed_seconds':time.monotonic()-START,'production_approved':False})
123
+ durable.sync('large critic '+phase+' step '+str(step))
124
+
125
+ def main():
126
+ shutil.copytree(source,OUT/'source',dirs_exist_ok=True)
127
+ shutil.copy2('/work/train.py',OUT/'train.py')
128
+ shutil.copy2('/work/r7.py',OUT/'data_and_eval.py')
129
+ model=load_semantic(initial,'cuda');ema=copy.deepcopy(model).eval().requires_grad_(False)
130
+ critic=Critic().cuda()
131
+ gc=sum(p.numel() for p in model.parameters());dc=sum(p.numel() for p in critic.parameters())
132
+ assert gc==3994676 and 10_000_000<dc<15_000_000
133
+ protocol={'generator_parameters':gc,'critic_parameters':dc,'base_revision':BASE,'steps':STEPS,'batch_size':BATCH,'seed':SEED,'warmup_critic_updates':500,'generator_objective':'ONLY non-saturating conditional adversarial logistic loss, equally weighted global and two patch heads; no reconstruction, teacher, perceptual, feature matching, hue-target, or smoothing loss','critic':'ImageNet pretrained ResNet18, RGB+L conditioning, global plus layer2/layer3 patch heads; all weights trainable with frozen BatchNorm statistics','regularization':'D-only lazy R1 gamma=10 every16 updates on8 samples; adaptive translation augmentation target real-positive fraction0.6; identical transforms for real/fake; EMA0.999','generator_optimizer':'Adam encoder1e-6 decoder2e-5; warmup500 then cosine floor0.2, no weight decay','critic_optimizer':'Adam pretrained2e-5 newheads1e-4; betas0,0.99','data':'up to60000 real train photographs; pinned COCO/Imagenette; original fold excluded; exact-byte train/eval duplicates removed; no teacher outputs','checkpoint_interval':250,'evaluation_interval':1000,'hard_timeout_hours':1.1,'graceful_training_budget_hours':0.85,'release_policy':'experiment folder only; compare grids with v3 and previous 27M candidate; no automatic replacement of v3; critic score not used to approve production','research':['https://arxiv.org/abs/2006.06676','https://arxiv.org/abs/1703.10593'],'limitations':'deterministic colouriser; one seed; adversarial loss can still collapse or exploit critic; grayscale is sometimes correct; no universal quality guarantee'}
134
+ r.write('protocol.json',protocol)
135
+ opt=torch.optim.Adam([{'params':model.encoder.parameters(),'lr':1e-6,'base_lr':1e-6},{'params':[v for k,v in model.named_parameters() if not k.startswith('encoder.')],'lr':2e-5,'base_lr':2e-5}],betas=(0.0,.99))
136
+ pretrained=[v for k,v in critic.named_parameters() if k.startswith('net.') and not k.startswith('net.fc')]
137
+ heads=[v for k,v in critic.named_parameters() if not k.startswith('net.') or k.startswith('net.fc')]
138
+ dopt=torch.optim.Adam([{'params':pretrained,'lr':2e-5},{'params':heads,'lr':1e-4}],betas=(0.0,.99))
139
+ # Real backward + lazy R1 + checkpoint roundtrip before costly data preparation.
140
+ L=torch.zeros(2,1,256,256,device='cuda')
141
+ pred=model(L);critic.eval()
142
+ loss=loss_real(critic(pair(L,pred)));loss.backward()
143
+ assert all(torch.isfinite(v.grad).all() for v in model.parameters() if v.grad is not None)
144
+ x=pair(L,torch.zeros_like(pred)).detach().requires_grad_(True)
145
+ grad=torch.autograd.grad(score_mean(critic(x)).sum(),x,create_graph=True)[0]
146
+ (grad.square().flatten(1).sum(1).mean()*5).backward()
147
+ model.zero_grad(set_to_none=True);critic.zero_grad(set_to_none=True)
148
+ del L,pred,loss,x,grad
149
+ save_semantic(model,OUT/'initial')
150
+ reload=load_semantic(OUT/'initial','cuda')
151
+ for k,v in model.state_dict().items():assert torch.equal(v,reload.state_dict()[k])
152
+ del reload
153
+ r.write('smoke.json',{'backward_finite':True,'R1_second_derivative':True,'checkpoint_roundtrip':True,'generator_parameters':gc,'critic_parameters':dc})
154
+ durable.sync('preflight passed before data preparation')
155
+ print('PREFLIGHT_PASSED',gc,dc,flush=True)
156
+ tables=[]
157
+ 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'])]:
158
+ for f in files:tables.append(pq.read_table(hf_hub_download(repo,f,repo_type='dataset',revision=rev),columns=['image']))
159
+ table=pa.concat_tables(tables);manifest=json.loads((source/'previous_manifest.json').read_text())
160
+ ids=list(manifest['train']);provenance=[]
161
+ files=sorted(f for f in api.list_repo_files('detection-datasets/coco',repo_type='dataset',revision='26ddc382fe75dfc2a0655b5977e296ea10efebce') if f.startswith('default/train/') and f.endswith('.parquet') and not f.endswith(('/0000.parquet','/0001.parquet')))
162
+ for f in files:
163
+ if len(ids)>=60000:break
164
+ extra=pq.read_table(hf_hub_download('detection-datasets/coco',f,repo_type='dataset',revision='26ddc382fe75dfc2a0655b5977e296ea10efebce'),columns=['image'])
165
+ offset=len(table);take=min(len(extra),60000-len(ids));ids.extend(range(offset,offset+take))
166
+ table=pa.concat_tables([table,extra]);provenance.append({'file':f,'take':take})
167
+ fresh=pq.read_table(hf_hub_download('detection-datasets/coco','default/val/0001.parquet',repo_type='dataset',revision='26ddc382fe75dfc2a0655b5977e296ea10efebce'),columns=['image'])
168
+ prior=json.loads(Path(hf_hub_download(REPO,'experiments/round7-20260928/manifest.json',revision=BASE)).read_text())
169
+ excluded=set(prior['development']+prior['heldout'])
170
+ candidates=[int(i) for i in np.random.default_rng(SEED).permutation(len(fresh)) if int(i) not in excluded][:512]
171
+ dev=candidates[:128];held=candidates[128:]
172
+ evalhash={hashlib.sha256(fresh['image'][i].as_py()['bytes']).digest() for i in candidates}
173
+ ids=[i for i in ids if hashlib.sha256(table['image'][i].as_py()['bytes']).digest() not in evalhash]
174
+ # Real colour data only, but retain naturally neutral regions within these photos.
175
+ keep=[]
176
+ for i in ids:
177
+ im=decode_image(table,i).resize((32,32)).convert('RGB')
178
+ lab=rgb2lab(np.asarray(im,np.float32)/255)
179
+ if float(np.linalg.norm(lab[...,1:],axis=-1).mean())>=4:keep.append(i)
180
+ ids=keep
181
+ r.write('manifest.json',{'train':ids,'development':dev,'heldout':held,'additional_train_shards':provenance,'coco_revision':'26ddc382fe75dfc2a0655b5977e296ea10efebce','imagenette_revision':'771c1310a2487e8076ede6b7d6307244aa8400af','exact_duplicate_filter':True,'near_duplicates_not_excluded':True,'heldout_not_used_for_checkpoint_selection':True})
182
+ baseline=load_semantic(initial,'cuda')
183
+ olddir=WORK/'prior_27m';olddir.mkdir(exist_ok=True)
184
+ for fn in ['model.safetensors','config.json']:
185
+ shutil.copy2(hf_hub_download(REPO,'experiments/large-critic-20260928/candidate/'+fn,revision='main'),olddir/fn)
186
+ oldcandidate=load_semantic(olddir,'cuda')
187
+ r.write('baseline_development.json',r.ev.evaluate(baseline,fresh,dev))
188
+ durable.sync('data and fresh evaluation split persisted')
189
+ loader=DataLoader(Photos(table,ids),batch_size=BATCH,shuffle=True,num_workers=8,pin_memory=True,drop_last=True,persistent_workers=True,generator=torch.Generator().manual_seed(SEED))
190
+ it=iter(loader);history=[];augp=.2;real_signs=[];step=0;stopreason='completed'
191
+ for update in range(1,STEPS+501):
192
+ if time.monotonic()-START>MAX_SECONDS:stopreason='time_budget';break
193
+ model.train();critic.train()
194
+ for mod in model.encoder.modules():
195
+ if isinstance(mod,nn.BatchNorm2d):mod.eval()
196
+ try:L,ab=next(it)
197
+ except StopIteration:it=iter(loader);L,ab=next(it)
198
+ L,ab=L.cuda(non_blocking=True),ab.cuda(non_blocking=True)
199
+ # D update. Generator graph is constructed only for its own update.
200
+ with torch.no_grad(),torch.autocast('cuda',dtype=torch.bfloat16):pred=model(L)
201
+ for v in critic.parameters():v.requires_grad_(True)
202
+ dopt.zero_grad(set_to_none=True)
203
+ realx=pair(L,ab);fakex=pair(L,pred)
204
+ joined=augment(torch.cat([realx,fakex],0),augp)
205
+ with torch.autocast('cuda',dtype=torch.bfloat16):
206
+ scores=critic(joined)
207
+ real=[s[:BATCH].float() for s in scores];fake=[s[BATCH:].float() for s in scores]
208
+ dl=loss_real(real)+loss_fake(fake)
209
+ dl.backward()
210
+ r1value=0.
211
+ if update%16==0:
212
+ rx=realx[:8].detach().requires_grad_(True)
213
+ rs=score_mean(critic(rx))
214
+ rg=torch.autograd.grad(rs.sum(),rx,create_graph=True)[0]
215
+ r1=rg.square().flatten(1).sum(1).mean()
216
+ (r1*(10/2)*16).backward();r1value=float(r1.detach())
217
+ torch.nn.utils.clip_grad_norm_(critic.parameters(),10,error_if_nonfinite=True);dopt.step()
218
+ real_signs.append(float((score_mean(real).detach()>0).float().mean()))
219
+ if update%100==0:
220
+ augp=float(np.clip(augp+np.sign(np.mean(real_signs)-.6)*BATCH*100/500000,0,.8));real_signs=[]
221
+ glvalue=None
222
+ if update>500:
223
+ step=update-500
224
+ for v in critic.parameters():v.requires_grad_(False)
225
+ critic.eval();opt.zero_grad(set_to_none=True)
226
+ mult=min(step/500,1)*(.2+.8*(1+math.cos(math.pi*step/STEPS))/2)
227
+ for group in opt.param_groups:group['lr']=group['base_lr']*mult
228
+ with torch.autocast('cuda',dtype=torch.bfloat16):
229
+ pred=model(L)
230
+ gl=loss_real([s.float() for s in critic(augment(pair(L,pred),augp))])
231
+ gl.backward();torch.nn.utils.clip_grad_norm_(model.parameters(),1,error_if_nonfinite=True);opt.step()
232
+ glvalue=float(gl.detach())
233
+ with torch.no_grad():
234
+ for ep,p in zip(ema.parameters(),model.parameters()):ep.lerp_(p,.001)
235
+ for eb,b in zip(ema.buffers(),model.buffers()):eb.copy_(b)
236
+ if update%100==0:
237
+ item={'update':update,'generator_step':step,'D':float(dl.detach()),'G':glvalue,'R1':r1value,'augmentation_p':augp,'real_score':float(score_mean(real).detach().mean()),'fake_score':float(score_mean(fake).detach().mean()),'elapsed_seconds':time.monotonic()-START}
238
+ history.append(item);print('TRAIN',json.dumps(item),flush=True)
239
+ if update==500 or (step>0 and step%250==0):
240
+ save(step,model,ema,critic,opt,dopt,history,augp,'training')
241
+ if step==0 or step%1000==0:
242
+ r.write(f'development_{step:05d}.json',r.ev.evaluate(ema,fresh,dev))
243
+ r.grid({'v3':baseline,'prior27':oldcandidate,'ema':ema,'raw':model.eval()},fresh,dev[:24],f'comparison_{step:05d}')
244
+ r.ev.probes(ema,f'probes_{step:05d}')
245
+ durable.sync('development review '+str(step))
246
+ save(step,model,ema,critic,opt,dopt,history,augp,'training_finished')
247
+ # Held-out report only after training; no best-checkpoint claims based on critic scores.
248
+ for name,m in [('v3',baseline),('prior27',oldcandidate),('final_ema',ema),('final_raw',model.eval())]:
249
+ r.write(name+'_heldout.json',r.ev.evaluate(m,fresh,held))
250
+ r.review(m,fresh,held[:80],name)
251
+ r.grid({'v3':baseline,'prior27':oldcandidate,'ema':ema,'raw':model.eval()},fresh,held[:32],'heldout_comparison')
252
+ save_semantic(ema,OUT/'candidate')
253
+ r.write('status.json',{'status':'completed','generator_steps':step,'stop_reason':stopreason,'elapsed_seconds':time.monotonic()-START,'production_approved':False,'next':'Review developmental checkpoint grids and held-out comparisons; this job does not automatically replace v3.'})
254
+ durable.sync('completed candidate and independent reports')
255
+
256
+ try:
257
+ main()
258
+ except BaseException:
259
+ r.write('failure.json',{'traceback':traceback.format_exc().replace(os.environ.get('HF_TOKEN','__EMPTY__'),'[REDACTED]')})
260
+ durable.sync('failure report and saved checkpoints')
261
+ raise