Colorizer round6: preflight passed before data preparation
Browse files- experiments/medium-critic-20260928/SHA256SUMS.json +26 -0
- experiments/medium-critic-20260928/data_and_eval.py +234 -0
- experiments/medium-critic-20260928/initial/README.md +5 -0
- experiments/medium-critic-20260928/initial/config.json +7 -0
- experiments/medium-critic-20260928/initial/model.safetensors +3 -0
- experiments/medium-critic-20260928/protocol.json +25 -0
- experiments/medium-critic-20260928/smoke.json +7 -0
- experiments/medium-critic-20260928/source/__pycache__/data.cpython-311.pyc +0 -0
- experiments/medium-critic-20260928/source/__pycache__/legacy_eval.cpython-311.pyc +0 -0
- experiments/medium-critic-20260928/source/__pycache__/metrics.cpython-311.pyc +0 -0
- experiments/medium-critic-20260928/source/__pycache__/model.cpython-311.pyc +0 -0
- experiments/medium-critic-20260928/source/__pycache__/objectives.cpython-311.pyc +0 -0
- experiments/medium-critic-20260928/source/__pycache__/persistence.cpython-311.pyc +0 -0
- experiments/medium-critic-20260928/source/__pycache__/semantic_model.cpython-311.pyc +0 -0
- experiments/medium-critic-20260928/source/__pycache__/spatial.cpython-311.pyc +0 -0
- experiments/medium-critic-20260928/source/data.py +116 -0
- experiments/medium-critic-20260928/source/legacy_eval.py +150 -0
- experiments/medium-critic-20260928/source/metrics.py +35 -0
- experiments/medium-critic-20260928/source/model.py +122 -0
- experiments/medium-critic-20260928/source/objectives.py +17 -0
- experiments/medium-critic-20260928/source/persistence.py +42 -0
- experiments/medium-critic-20260928/source/previous_manifest.json +0 -0
- experiments/medium-critic-20260928/source/semantic_model.py +94 -0
- experiments/medium-critic-20260928/source/spatial.py +34 -0
- 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
|