Download experiments/recurrent-refinement-20260928/data_and_eval.py from User-2468/mini-unet-colorizer: direct link, hf CLI and curl.
- Browser
- Download file 14.3 kB
-
https://huggingface.co/User-2468/mini-unet-colorizer/resolve/7ec6b1210fd5af800b5f8a5628a856ca0098b0a5/experiments/recurrent-refinement-20260928/data_and_eval.py
- Command line
-
hf download hf://User-2468/mini-unet-colorizer@7ec6b1210fd5af800b5f8a5628a856ca0098b0a5/experiments/recurrent-refinement-20260928/data_and_eval.py
-
curl -L -o data_and_eval.py https://huggingface.co/User-2468/mini-unet-colorizer/resolve/7ec6b1210fd5af800b5f8a5628a856ca0098b0a5/experiments/recurrent-refinement-20260928/data_and_eval.py
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) | |
| 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)) | |
| 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 | |
| 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 | |