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=b0 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