User-2468's picture
Colorizer round6: smoke and initial model verified
28a3d7d verified
Raw History Blame
14.3 kB
import os,sys,json,random,time,hashlib,shutil,math,traceback,io
from pathlib import Path
import numpy as np
import torch
import torch.nn.functional as F
from torch.utils.data import Dataset,DataLoader
from huggingface_hub import hf_hub_download
import pyarrow as pa
import pyarrow.parquet as pq
from PIL import Image,ImageDraw
from skimage.color import rgb2lab,lab2rgb
from skimage.segmentation import slic
import cv2
import persistence
from semantic_model import load_semantic,save_semantic
from objectives import color_loss
from data import decode_image
import legacy_eval as ev
REPO='User-2468/mini-unet-colorizer'
REV='7b10485ee9d205ce2f7b660dfc6a90012d78b1cb'
OLD='experiments/round6-20260928'
OUT=Path('/work/results');OUT.mkdir(parents=True,exist_ok=True)
persistence.PREFIX='experiments/round7-20260928'
ev.OUT=OUT
START=time.monotonic()
SEED=709
STEPS=2000
DEVICE='cpu' if os.environ.get('SMOKE')=='1' else 'cuda'
ev.DEVICE=DEVICE
DURABLE=None
def write(name,obj):
p=OUT/name;p.parent.mkdir(parents=True,exist_ok=True);p.write_text(json.dumps(obj,indent=2));return obj
def deadline():
if time.monotonic()-START>9000:raise TimeoutError('Graceful deadline')
ev.deadline=deadline
def losses(p,t):
p=F.adaptive_avg_pool2d(p.float(),(64,64));t=t.float()
pixel=(p-t).abs().mean((1,2,3))/20
a=F.avg_pool2d(p,2).flatten(2);b=F.avg_pool2d(t,2).flatten(2)
angles=torch.arange(8,device=p.device)*math.pi/8
directions=torch.stack([angles.cos(),angles.sin()],1)
ap=torch.einsum('kc,bcn->bkn',directions,a).sort(-1).values
bp=torch.einsum('kc,bcn->bkn',directions,b).sort(-1).values
distribution=(ap-bp).abs().mean((1,2))/20
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
return pixel+.25*distribution+.1*edge
def select_loss(p,teacher,original,mode):
a=losses(p,teacher)
if mode=='teacher':return a.mean(),torch.zeros_like(a,dtype=torch.bool)
b=losses(p,original);choose=b<a
return torch.where(choose,b,a).mean(),choose
def smoke():
torch.set_num_threads(4);torch.manual_seed(SEED)
p=torch.randn(2,2,64,64,requires_grad=True);t=torch.randn_like(p)
assert torch.allclose(losses(p,t).mean(),color_loss(p,t)[0],atol=1e-6)
# Whole-image choice: exact red or blue is better than an average or a patchwork.
red=torch.zeros(1,2,64,64);red[:,0]=30
blue=-red;mix=red.clone();mix[:,:,:,32:]=blue[:,:,:,32:]
assert select_loss(red,red,blue,'set')[0]==0
assert select_loss(torch.zeros_like(red),red,blue,'set')[0]>0
assert select_loss(mix,red,blue,'set')[0]>0
m=load_semantic('/work/initial',DEVICE)
assert sum(p.numel() for p in m.parameters())==3994676
x=torch.rand(2,1,96,128,device=DEVICE)*2-1
pred=m(x);loss,_=select_loss(pred,t.to(DEVICE),-t.to(DEVICE),'set')
loss.backward()
assert torch.isfinite(loss) and all(torch.isfinite(p.grad).all() for p in m.parameters() if p.grad is not None)
save_semantic(m,OUT/'initial')
re=load_semantic(OUT/'initial',DEVICE)
with torch.no_grad():assert torch.allclose(m.eval()(x),re(x),atol=1e-5)
write('smoke.json',{'passed':True,'parameters':3994676,'loss_equivalence':True,'whole_image_choice':True,'finite_backward':True,'reload':True})
del m,re
class Photos(Dataset):
def __init__(self,table,ids,cache,archive):
self.table,self.ids,self.cache,self.archive=table,ids,cache,archive;self.epoch=0
def __len__(self):return len(self.ids)
def __getitem__(self,k):
im=decode_image(self.table,self.ids[k]);L,ab,gray=ev.gray_arrays(im)
# Independent per-example RNG; augmentation differences never alter image ordering/flips.
rng=np.random.default_rng(np.random.SeedSequence([SEED,self.epoch,k]))
target=np.array(self.cache[k],np.float32)
if self.archive and rng.random()<.5:
rgb=np.asarray(im.resize((256,256),Image.Resampling.BILINEAR),np.float32)/255
weights=rng.dirichlet(np.array([3.,5.,2.]))
g=(rgb*weights).sum(-1)
g=np.clip((g**rng.uniform(.7,1.5)-.5)*rng.uniform(.6,1.1)+.5+rng.uniform(-.08,.08),0,1)
if rng.random()<.5:g=cv2.GaussianBlur(g,(5,5),float(rng.uniform(.3,1.2)))
if rng.random()<.6:g=np.clip(g+rng.normal(0,rng.uniform(.003,.025),g.shape),0,1)
L=(rgb2lab(np.repeat(g[...,None],3,-1).astype(np.float32))[...,0][None]/50-1).astype(np.float32)
else:
if rng.random()<.5:L=np.clip(L*rng.uniform(.7,1.15)+rng.uniform(-.12,.12),-1,1)
flip=np.random.default_rng(np.random.SeedSequence([SEED+1,self.epoch,k])).random()<.5
if flip:L=L[...,::-1].copy();ab=ab[...,::-1].copy();target=target[...,::-1].copy()
original=cv2.resize(ab.transpose(1,2,0),(64,64),interpolation=cv2.INTER_AREA).transpose(2,0,1).copy()
return torch.from_numpy(L),torch.from_numpy(target),torch.from_numpy(original)
@torch.inference_mode()
def predict(m,L):
x=torch.from_numpy(L)[None].to(DEVICE);raw=m(x)
guided=ev.guided_chroma(x,raw,8)
return raw,guided
def toimage(L,ab):
rgb=lab2rgb(np.concatenate([(L[0]*50+50)[...,None],ab],-1))
return Image.fromarray(np.uint8(np.clip(rgb,0,1)*255))
@torch.inference_mode()
def review(m,table,ids,name):
m.eval();rows=[]
# Reference-derived surface regions for diagnostics only, never inference.
for i in ids:
L,target,_=ev.gray_arrays(decode_image(table,i))
raw,pred=predict(m,L);x=torch.from_numpy(L)[None].to(DEVICE)
resize=F.interpolate(x,(224,224),mode='bilinear',align_corners=False)
resized=F.interpolate(m(resize),(256,256),mode='bilinear',align_corners=False)
contrast=m((x*.95+.02).clamp(-1,1))
pab=pred[0].cpu().numpy().transpose(1,2,0)
rgb=np.asarray(decode_image(table,i).resize((256,256)),np.float32)/255
seg=slic(rgb,n_segments=100,compactness=10,start_label=0)
tab=target.transpose(1,2,0);variations=[]
for label in np.unique(seg):
mask=seg==label
if mask.sum()<64:continue
ground=tab[mask];center=np.median(ground,axis=0)
if np.percentile(np.linalg.norm(ground-center,axis=1),90)>5:continue
col=pab[mask];variations.append(float(np.mean(np.linalg.norm(col-np.median(col,axis=0),axis=1))))
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())})
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.'})
return rows
@torch.inference_mode()
def grid(models,table,ids,label):
order=np.random.default_rng(SEED+2).permutation(len(models))
names=list(models);key={chr(65+j):names[int(i)] for j,i in enumerate(order)}
write(label+'_key.json',key)
for page in range(0,min(len(ids),32),8):
chosen=ids[page:page+8];canvas=Image.new('RGB',(256*(len(models)+1),274*len(chosen)),'white');draw=ImageDraw.Draw(canvas)
for row,index in enumerate(chosen):
L,_,gray=ev.gray_arrays(decode_image(table,index));draw.text((2,row*274),f'{index}: gray | '+' | '.join(key),fill='black')
canvas.paste(Image.fromarray(np.uint8(gray*255)),(0,row*274+18))
for j,i in enumerate(order):
_,p=predict(models[names[int(i)]],L)
canvas.paste(toimage(L,p[0].cpu().numpy().transpose(1,2,0)),((j+1)*256,row*274+18))
canvas.save(OUT/f'{label}_{page//8}.jpg',quality=92)
def main():
global DURABLE
torch.set_num_threads(6);torch.manual_seed(SEED);random.seed(SEED);np.random.seed(SEED)
DURABLE=persistence.DurableRun(OUT)
shutil.copy2('/work/runner.py',OUT/'runner.py');shutil.copy2('/work/PROTOCOL.md',OUT/'PROTOCOL.md')
shutil.copytree('/work/source',OUT/'source',ignore=shutil.ignore_patterns('__pycache__'),dirs_exist_ok=True)
smoke();DURABLE.sync('round7 initial checkpoint and smoke verified')
if DEVICE=='cpu':return
torch.backends.cudnn.benchmark=True
tables=[]
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'])]:
for f in files:tables.append(pq.read_table(hf_hub_download(repo,f,repo_type='dataset',revision=rev),columns=['image']))
table=pa.concat_tables(tables);manifest=json.loads(Path('/work/source/previous_manifest.json').read_text());ids=manifest['train']
cache=np.empty((len(ids),2,64,64),np.float16)
sums=json.loads(Path(hf_hub_download(REPO,OLD+'/SHA256SUMS.json',revision=REV)).read_text())
for start in range(0,len(ids),2048):
f=f'teacher_cache/{start:05d}.npz';path=Path(hf_hub_download(REPO,OLD+'/'+f,revision=REV))
assert hashlib.sha256(path.read_bytes()).hexdigest()==sums[f]
with np.load(path) as z:
assert z['indices'].tolist()==ids[start:start+len(z['indices'])]
cache[start:start+len(z['target'])]=z['target']
fresh=pq.read_table(hf_hub_download('detection-datasets/coco','default/val/0001.parquet',repo_type='dataset',revision='26ddc382fe75dfc2a0655b5977e296ea10efebce'),columns=['image'])
hashes={hashlib.sha256(table['image'][int(i)].as_py()['bytes']).digest() for i in ids}
freshids=[]
for i in np.random.default_rng(SEED).permutation(len(fresh)):
if hashlib.sha256(fresh['image'][int(i)].as_py()['bytes']).digest() not in hashes:freshids.append(int(i))
if len(freshids)==160:break
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})
DURABLE.sync('data cache verified and split fixed')
m=load_semantic('/work/initial',DEVICE)
write('initial_validation.json',ev.evaluate(m,table,manifest['validation']))
review(m,fresh,freshids[:80],'initial_development')
review(m,fresh,freshids[80:],'initial_heldout')
write('initial_heldout_reference_metrics.json',ev.evaluate(m,fresh,freshids[80:]))
del m
arms=[('teacher_clean','teacher',False),('set_clean','set',False),('teacher_archive','teacher',True),('set_archive','set',True)]
for name,mode,archive in arms:
torch.manual_seed(SEED);np.random.seed(SEED);random.seed(SEED)
m=load_semantic('/work/initial',DEVICE)
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)
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))
ds=Photos(table,ids,cache,archive)
loader=DataLoader(ds,batch_size=32,shuffle=True,num_workers=4,pin_memory=True,generator=torch.Generator().manual_seed(SEED))
it=iter(loader);history=[];choices=0;seen=0
for step in range(1,STEPS+1):
deadline();m.train()
for layer in m.encoder.modules():
if isinstance(layer,torch.nn.BatchNorm2d):layer.eval()
try:L,t,o=next(it)
except StopIteration:ds.epoch+=1;it=iter(loader);L,t,o=next(it)
L,t,o=L.to(DEVICE),t.to(DEVICE),o.to(DEVICE);opt.zero_grad(set_to_none=True)
with torch.autocast('cuda',dtype=torch.bfloat16):p=m(L)
loss,choice=select_loss(p,t,o,mode)
assert torch.isfinite(loss)
loss.backward();torch.nn.utils.clip_grad_norm_(m.parameters(),1,error_if_nonfinite=True);opt.step();sched.step()
choices+=int(choice.sum());seen+=len(L)
if step%100==0:
history.append({'step':step,'loss':float(loss),'original_choice_fraction':choices/seen})
print('STEP',name,step,round(float(loss),4),flush=True)
if step%500==0:
path=OUT/name/f'step_{step:05d}';save_semantic(m,path)
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')
write(name+'/history.json',history);write('status.json',{'status':'training','arm':name,'step':step,'elapsed_seconds':time.monotonic()-START})
DURABLE.sync(name+' checkpoint '+str(step))
write(name+'_validation.json',ev.evaluate(m,table,manifest['validation']))
review(m,fresh,freshids[:80],name+'_development')
review(m,fresh,freshids[80:],name+'_heldout')
write(name+'_heldout_reference_metrics.json',ev.evaluate(m,fresh,freshids[80:]))
ev.probes(m,name)
DURABLE.sync(name+' complete diagnostics')
del m,opt,sched,it,loader;torch.cuda.empty_cache()
models={'round6':load_semantic('/work/initial',DEVICE)}
for name,_,_ in arms:models[name]=load_semantic(OUT/name/f'step_{STEPS:05d}',DEVICE)
grid(models,fresh,freshids[:80],'blinded_development')
grid(models,fresh,freshids[80:],'blinded_heldout')
# Familiar historical development probes; never claim these are a held-out archival benchmark.
import requests
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'}
for name,url in sources.items():
try:
response=requests.get(url,timeout=45);response.raise_for_status()
im=Image.open(io.BytesIO(response.content)).convert('RGB')
tb=pa.table({'image':[{'bytes':response.content,'path':None}]})
grid(models,tb,[0],name)
write(name+'_source.json',{'url':url,'sha256':hashlib.sha256(response.content).hexdigest(),'use':'previously seen archival development probe; original colors unknown'})
except Exception as e:write(name+'_error.json',{'type':type(e).__name__,'message':str(e)})
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.'})
DURABLE.sync('round7 four-arm experiment completed')
if __name__=='__main__':
try:main()
except BaseException as e:
write('failure.json',{'type':type(e).__name__,'traceback':traceback.format_exc().replace(os.environ.get('HF_TOKEN','__EMPTY__'),'[REDACTED]')})
if DURABLE:DURABLE.sync('round7 failure report and latest saved checkpoints')
raise