"""Budget-bounded architecture/curriculum screening. Sources pinned by bootstrap.""" import os,sys,json,time,math,random,hashlib,shutil,traceback,copy,io from pathlib import Path import numpy as np import torch from torch import nn import torch.nn.functional as F from PIL import Image,ImageDraw from skimage.color import rgb2lab import cv2,pyarrow as pa,pyarrow.parquet as pq from huggingface_hub import HfApi,hf_hub_download from safetensors.torch import load_file from semantic_model import load_semantic from decision_model import DecisionColorizer,initialize_from_v3,save_decision,load_decision import persistence import r7 as r from data import decode_image REPO='User-2468/mini-unet-colorizer' BASE='1a9eb8af2754ad2329a24cfe50d388cb559441d0' ARTIFACTS='12dfd8111622f692a7f69d0a961ce9b65f044c01' CACHE_REV='7b10485ee9d205ce2f7b660dfc6a90012d78b1cb' COCO_REV='26ddc382fe75dfc2a0655b5977e296ea10efebce' PREFIX=os.environ.get('EXPERIMENT_PREFIX','experiments/final-20260929') SEED=int(os.environ.get('SEED','9292026')) STEPS=int(os.environ.get('STEPS','24000'));BATCH=24 START=time.monotonic();TRAIN_LIMIT=float(os.environ.get('TRAIN_LIMIT_SECONDS','3000')) OUT=Path('/work/results');OUT.mkdir(exist_ok=True,parents=True) persistence.PREFIX=PREFIX;r.OUT=OUT;r.ev.OUT=OUT;r.DEVICE='cuda';r.ev.DEVICE='cuda';r.ev.deadline=lambda:None api=HfApi();DURABLE=None def write(n,obj): p=OUT/n;p.parent.mkdir(parents=True,exist_ok=True);p.write_text(json.dumps(obj,indent=2));return obj def dl(f,rev=ARTIFACTS):return Path(hf_hub_download(REPO,f,revision=rev)) def seed(s):torch.manual_seed(s);np.random.seed(s);random.seed(s) def objective(details,target,step): allout=details['all'];b,k=allout.shape[:2] losses=r.losses(allout.flatten(0,1),target[:,None].expand(-1,k,-1,-1,-1).flatten(0,1)).reshape(b,k) if k==1:return losses.mean(),losses.argmin(1),losses # One winning complete image, never a per-pixel best-of-K patchwork. winner=losses.detach().argmin(1) selected=losses.gather(1,winner[:,None]).mean() # Small nonwinner term prevents completely untrained branches; no saturation reward. loss=.95*selected+.05*losses.mean()+.1*F.cross_entropy(details['scores'],winner) return loss,winner,losses class Batches: def __init__(self,rgb,teacher): self.rgb=torch.from_numpy(rgb).cuda().permute(0,3,1,2) self.teacher=torch.from_numpy(teacher).cuda().float() self.n=len(rgb) @staticmethod def light(g): y=torch.where(g<=.04045,g/12.92,((g+.055)/1.055).pow(2.4)) l=torch.where(y>.008856,116*y.clamp_min(1e-9).pow(1/3)-16,903.296296*y) return l/50-1 def batch(self,step,recipe,original,curriculum_step=None): phase=step if curriculum_step is None else curriculum_step gen=torch.Generator(device='cuda').manual_seed(SEED*100000+step) ids=torch.randint(self.n,(BATCH,),generator=gen,device='cuda') rgb=self.rgb[ids].float()/255;t=self.teacher[ids];o=original[ids] g=(rgb*rgb.new_tensor([.299,.587,.114])[None,:,None,None]).sum(1,keepdim=True) # Emulate standard 8-bit PIL grayscale conversion before Lab conversion. g=(g*255).round()/255 L=self.light(g) if recipe=='staged' and phase>STEPS*2/3: take=torch.rand(BATCH,1,1,1,generator=gen,device='cuda')<.35 weights=torch.rand(BATCH,3,1,1,generator=gen,device='cuda')+rgb.new_tensor([.8,1.2,.2])[None,:,None,None] weights/=weights.sum(1,keepdim=True) film=(rgb*weights).sum(1,keepdim=True) gamma=.85+.4*torch.rand(BATCH,1,1,1,generator=gen,device='cuda') film=(film.pow(gamma)-.5)*(.8+.3*torch.rand(BATCH,1,1,1,generator=gen,device='cuda'))+.5 film=film+.005*torch.randn(film.shape,generator=gen,device='cuda') L=torch.where(take,self.light(film.clamp(0,1)),L) else: gain=.9+.2*torch.rand(BATCH,1,1,1,generator=gen,device='cuda') L=(L*gain).clamp(-1,1) flips=torch.rand(BATCH,1,1,1,generator=gen,device='cuda')<.5 L=torch.where(flips,L.flip(-1),L);t=torch.where(flips,t.flip(-1),t);o=torch.where(flips,o.flip(-1),o) prob=0 if recipe=='teacher' else .5 if recipe=='staged':prob*=min(max((phase/STEPS-.15)/.35,0),1) pick=torch.rand(BATCH,1,1,1,generator=gen,device='cuda')=.95*b['color_coverage'] and s['ab_error']<=b['ab_error']*1.10 val+=s['patch_excess']/max(b['patch_excess'],1e-6)+.2*s['ab_error']/b['ab_error']+2*s['missed_color']+s['neutral_spill'] return val/2,bool(okay) @torch.inference_mode() def diagnostics(model,table,ids): model.eval();rows=[] for idx in ids: L,t,_=r.ev.gray_arrays(decode_image(table,idx));x=torch.from_numpy(L)[None].cuda() v=model(x,True);a=v['all'];sel=int(v['scores'].argmax(1)) resized=model(F.interpolate(x,(224,224),mode='bilinear',align_corners=False),True) p=F.interpolate(resized['output'],(256,256),mode='bilinear',align_corners=False) flip=model(x.flip(-1),True) colors=a[0].mean((-1,-2));k=model.modes diff=(a[:,:,None]-a[:,None]).abs().mean((0,3,4,5)) target=torch.from_numpy(t)[None].cuda();target64=F.avg_pool2d(target,4) losses=r.losses(a.flatten(0,1),target64.expand(k,-1,-1,-1)) mask=v['weights'].clamp_min(1e-8) rows.append({'index':int(idx),'selected_mode':sel,'resized_mode':int(resized['scores'].argmax(1)), 'flip_mode':int(flip['scores'].argmax(1)),'resize_ab_change':float((p-v['output']).abs().mean()), 'flip_ab_change':float((flip['output'].flip(-1)-v['output']).abs().mean()), 'mode_diversity':float(diff.sum()/max(k*(k-1),1)), 'default_target_loss':float(losses[sel]),'oracle_target_loss':float(losses.min()), 'mask_entropy':float(-(mask*mask.log()).sum(1).mean()), 'palette_chroma':float(v['palette'].norm(dim=-1).mean()),'output_chroma':float(v['output'].norm(dim=1).mean())}) keys=['resize_ab_change','flip_ab_change','mode_diversity','default_target_loss','oracle_target_loss','mask_entropy','palette_chroma','output_chroma'] return {'rows':rows,'summary':{k:float(np.mean([x[k] for x in rows])) for k in keys}, 'mode_frequency':np.bincount([x['selected_mode'] for x in rows],minlength=model.modes).tolist(), 'resize_switch_fraction':float(np.mean([x['selected_mode']!=x['resized_mode'] for x in rows])), 'note':'Oracle uses unknown original only for diagnosis. Default mode is selected by the learned image-only score, never by ground truth.'} def paired_bootstrap(result,base): rng=np.random.default_rng(937);out={} for style in ['gray','film']: br={z['index']:z for z in base[style]['per_image']};rows=[z for z in result[style]['per_image'] if z['target_chroma']>=5] out[style]={} for key in ['patch_excess','raw_patch_excess','missed_color','neutral_spill','color_coverage','ab_error']: d=np.array([z[key]-br[z['index']][key] for z in rows]);boot=d[rng.integers(len(d),size=(3000,len(d)))].mean(1) out[style][key]={'delta':float(d.mean()),'ci95':np.quantile(boot,[.025,.975]).tolist()} return out @torch.inference_mode() def pipeline_grid(models,table,ids,label): from inference import colorize from PIL import ImageOps names=list(models);order=np.random.default_rng(9184).permutation(len(names)) key={chr(65+j):names[int(i)] for j,i in enumerate(order)};write(label+'_key.json',key) for start in range(0,min(32,len(ids)),4): canvas=Image.new('RGB',(320*(len(names)+1),344*len(ids[start:start+4])),'white');d=ImageDraw.Draw(canvas) for row,idx in enumerate(ids[start:start+4]): original=decode_image(table,idx).convert('L').convert('RGB') original.thumbnail((800,800)) d.text((2,row*344),f'{idx}: grayscale | '+' | '.join(key),fill='black') canvas.paste(ImageOps.pad(original,(320,320),color='white'),(0,row*344+24)) for col,j in enumerate(order): output=colorize(models[names[int(j)]].eval(),original,size=256,radius=8) canvas.paste(ImageOps.pad(output,(320,320),color='white'),((col+1)*320,row*344+24)) canvas.save(OUT/f'{label}_{start//4}.jpg',quality=93) @torch.inference_mode() def resolution_grid(model,table,records): from inference import colorize from PIL import ImageOps for start in range(0,len(table),3): count=min(3,len(table)-start) canvas=Image.new('RGB',(1280,344*count),'white');d=ImageDraw.Draw(canvas) for row,idx in enumerate(range(start,start+count)): im=decode_image(table,idx).convert('L').convert('RGB');im.thumbnail((1000,1000)) d.text((2,row*344),records[idx]['name']+' | input | 256 | 384 | 512',fill='black') canvas.paste(ImageOps.pad(im,(320,320),color='white'),(0,row*344+24)) for col,size in enumerate([256,384,512],1): pred=colorize(model,im,size=size,radius=8) canvas.paste(ImageOps.pad(pred,(320,320),color='white'),(col*320,row*344+24)) canvas.save(OUT/f'v3_resolution_{start//3}.jpg',quality=94) def main(): global DURABLE assert api.whoami()['name']=='User-2468' DURABLE=persistence.DurableRun(OUT) shutil.copytree('/work/source',OUT/'source',dirs_exist_ok=True,ignore=shutil.ignore_patterns('__pycache__')) shutil.copy2('/work/PROTOCOL.md',OUT/'PROTOCOL.md') if Path('/work/REVIEW_HANDOFF.md').exists():shutil.copy2('/work/REVIEW_HANDOFF.md',OUT/'REVIEW_HANDOFF.md') if Path('/work/supplement').exists():shutil.copytree('/work/supplement',OUT/'supplement',dirs_exist_ok=True) write('status.json',{'status':'preflight','seed':SEED,'production_approved':False}) DURABLE.sync('source and write verification before training') torch.set_num_threads(6);seed(SEED) torch.backends.cuda.matmul.allow_tf32=False;torch.backends.cudnn.allow_tf32=False initial=Path('/work/initial');initial.mkdir(exist_ok=True) for f in ['model.safetensors','config.json']:shutil.copy2(dl(f,BASE),initial/f) baseline=load_semantic(initial,'cuda').requires_grad_(False) counts={};smokes={} encoder_rev=None write('encoder_provenance.json',{'backbone':'MobileNetV3 Large','initial_revision':ARTIFACTS,'initial_subfolder':'experiments/decision-20260928-v2/multi_mix/candidate'}) constructors=[('multi',4,'mobilenet')] for name,k,backbone in constructors: seed(SEED);m=DecisionColorizer(modes=k,backbone=backbone,pretrained=backbone=='mobilevit',encoder_revision=encoder_rev if backbone=='mobilevit' else None).cuda() copied=initialize_from_v3(m,baseline) seedpath=Path('/work/pilot');seedpath.mkdir(exist_ok=True) for f in ['config.json','model.safetensors']:shutil.copy2(dl('experiments/decision-20260928-v2/multi_mix/candidate/'+f),seedpath/f) m=load_decision(seedpath,'cuda');m.eval();count=sum(p.numel() for p in m.parameters());assert count<4_000_000 counts[name]=count x=torch.rand(2,1,97,113,device='cuda')*2-1 with torch.no_grad(): v=m(x,True);assert v['all'].shape==(2,k,2,97,113) and torch.isfinite(v['all']).all() parity=float((v['output']-baseline(x)).abs().max()) if name=='single' else None if parity is not None:assert parity<1e-3,parity loss,_,_=objective(m(x,True),torch.randn(2,2,64,64,device='cuda'),1) loss.backward();assert torch.isfinite(loss) and all(torch.isfinite(p.grad).all() for p in m.parameters() if p.grad is not None) m.zero_grad(set_to_none=True);save_decision(m,OUT/'initial'/name) re=load_decision(OUT/'initial'/name,'cuda') with torch.no_grad():assert torch.allclose(m(x),re(x),atol=1e-4) smokes[name]={'parameters':count,'copied_tensors':len(copied),'odd_shape':True,'finite_backward':True,'reload':True,'baseline_max_abs':parity} del m,re write('smoke.json',smokes);DURABLE.sync('all three architectures pass preflight') torch.backends.cuda.matmul.allow_tf32=True;torch.backends.cudnn.allow_tf32=True;torch.backends.cudnn.benchmark=True tables=[] for repo,rev,files in [('johnowhitaker/imagenette2-320','771c1310a2487e8076ede6b7d6307244aa8400af',['default/train/0000.parquet']),('detection-datasets/coco',COCO_REV,['default/train/0000.parquet','default/train/0001.parquet'])]: for f in files: print('PREP downloading',repo,f,flush=True) tables.append(pq.read_table(hf_hub_download(repo,f,repo_type='dataset',revision=rev),columns=['image'])) print('PREP loaded',repo,f,len(tables[-1]),flush=True) table=pa.concat_tables(tables) full_manifest=json.loads(dl('experiments/region-coherence-20260928/source/previous_manifest.json').read_text()) full_ids=full_manifest['train'] # Same fixed training subset for every arm, independent of new dev/test. ids=list(full_ids);pos={i:j for j,i in enumerate(ids)} teacher=np.empty((len(ids),2,64,64),np.float16) sums=json.loads(dl('experiments/round6-20260928/SHA256SUMS.json',CACHE_REV).read_text()) for start in range(0,len(full_ids),2048): f=f'teacher_cache/{start:05d}.npz';p=dl('experiments/round6-20260928/'+f,CACHE_REV) assert hashlib.sha256(p.read_bytes()).hexdigest()==sums[f] with np.load(p) as z: # NpzFile indexing decompresses a member on every access. Materialize once. cache_ids=z['indices'];cache_targets=z['target'] assert cache_ids.tolist()==full_ids[start:start+len(cache_ids)] src=[j for j,i in enumerate(cache_ids) if int(i) in pos] dst=[pos[int(cache_ids[j])] for j in src] teacher[dst]=cache_targets[src] print('PREP cache',start,'selected',len(src),flush=True) print('PREP decoding training images',len(ids),flush=True) rgb=np.stack([np.asarray(decode_image(table,i).resize((256,256),Image.Resampling.BILINEAR)) for i in ids]) original=np.stack([cv2.resize(rgb2lab(im.astype(np.float32)/255)[...,1:],(64,64),interpolation=cv2.INTER_AREA).transpose(2,0,1) for im in rgb]).astype(np.float32) batcher=Batches(rgb,teacher);original=torch.from_numpy(original).cuda();del rgb,teacher print('PREP training tensors ready',flush=True) fresh=pq.read_table(hf_hub_download('detection-datasets/coco','default/val/0001.parquet',repo_type='dataset',revision=COCO_REV),columns=['image']) exclude=set() for path,keys in [('decision-20260928-v2',['development','test']),('round7-20260928',['development','heldout']),('region-coherence-20260928',['development_regression','heldout_regression','new_holdout']),('recurrent-refinement-20260928',['fresh_holdout','development_regression'])]: manifest=json.loads(dl('experiments/'+path+'/manifest.json').read_text()) for key in keys:exclude.update(manifest[key]) hashes={hashlib.sha256(table['image'][i].as_py()['bytes']).digest() for i in full_ids} freshids=[int(i) for i in np.random.default_rng(SEED).permutation(len(fresh)) if int(i) not in exclude and hashlib.sha256(fresh['image'][int(i)].as_py()['bytes']).digest() not in hashes][:288] assert len(freshids)==288 dev,test=freshids[:128],freshids[128:] write('manifest.json',{'train':ids,'development':dev,'test':test,'prior_eval_excluded':sorted(exclude),'coco_revision':COCO_REV,'fresh_file':'default/val/0001.parquet','teacher_revision':CACHE_REV,'base_revision':BASE,'seed':SEED,'near_duplicates_and_upstream_overlap_not_excluded':True}) base_dev=r.ev.evaluate(baseline,fresh,dev);write('baseline_development.json',base_dev) DURABLE.sync('dataset checks and unseen development/test split frozen') arms=[('multi','mix')] if os.environ.get('ARMS'):arms=[tuple(x.split(':')) for x in os.environ['ARMS'].split(',')] selected={};histories={};dev_results={};begin_train=time.monotonic() for head,recipe in arms: name=head+'_'+recipe;seed(SEED) m=load_decision(OUT/'initial'/head,'cuda') # New encoder gets a separate alignment warmup, explicitly outside matched six arms. warmup=1000 if head=='mobilevit' else 0 groups=[{'params':[p for n,p in m.named_parameters() if n.startswith('encoder.')],'lr':2e-5 if head=='mobilevit' else 2e-6}, {'params':[p for n,p in m.named_parameters() if not n.startswith('encoder.')],'lr':2e-4 if head=='mobilevit' else 2e-5}] opt=torch.optim.AdamW(groups,weight_decay=.01) for g in opt.param_groups:g['base_lr']=g['lr'] history=[];beststep=0;last=0;arm_start=time.monotonic();mode_counts=np.zeros(m.modes,np.int64);stale=0 initial_res=r.ev.evaluate(m,fresh,dev);val,okay=score(initial_res,base_dev);best=val+(0 if okay else 100) save_decision(m,OUT/name/'candidate');dev_results[name]=initial_res write(name+'/development_00000.json',initial_res);DURABLE.sync('retain initial pilot candidate before continuation') for s in range(1,STEPS+warmup+1): if time.monotonic()-START>TRAIN_LIMIT:break step=max(1,s-warmup);active_recipe='teacher' if s<=warmup else recipe L,target,_,fraction=batcher.batch(s,active_recipe,original,curriculum_step=step) m.train() for layer in m.encoder.modules(): if isinstance(layer,nn.BatchNorm2d):layer.eval() opt.zero_grad(set_to_none=True) factor=min(s/150,1)*(.15+.85*(1+math.cos(math.pi*max(0,s-warmup)/STEPS))/2) for g in opt.param_groups:g['lr']=g['base_lr']*factor with torch.autocast('cuda',dtype=torch.bfloat16): details=m(L,True) loss,winners,terms=objective(details,target,step) consistency=loss.new_zeros(()) if active_recipe=='staged' and step>STEPS*2/3 and s%4==0: small=m(F.interpolate(L,(192,192),mode='bilinear',align_corners=False),True) a=F.adaptive_avg_pool2d(small['all'].flatten(0,1).float(),(64,64)) b=F.adaptive_avg_pool2d(details['all'].flatten(0,1).detach().float(),(64,64)) consistency=F.l1_loss(a/20,b/20) if m.modes>1:consistency+=.1*F.kl_div(small['scores'].log_softmax(1),details['scores'].detach().softmax(1),reduction='batchmean') loss=loss+.4*consistency # every fourth step -> mean weight .1 assert torch.isfinite(loss),name loss.backward();torch.nn.utils.clip_grad_norm_(m.parameters(),1,error_if_nonfinite=True);opt.step();last=s mode_counts+=np.bincount(winners.detach().cpu().numpy(),minlength=m.modes) if s%100==0: row={'step':s,'loss':float(loss.detach()),'reconstruction':float(terms.mean().detach()),'consistency':float(consistency.detach()),'photo_fraction_this_batch':fraction,'winner_counts':mode_counts.tolist(),'seconds':time.monotonic()-arm_start} history.append(row);print('STEP',name,s,round(row['loss'],4),flush=True) if s%2000==0 or s==STEPS+warmup: m.eval();res=r.ev.evaluate(m,fresh,dev);val,okay=score(res,base_dev) write(name+f'/development_{s:05d}.json',res) write(name+'/history.json',history) save_decision(m,OUT/name/f'step_{s:05d}') # Ineligible models still retained for diagnosis, ranked after eligible. ranked=val+(0 if okay else 100) if ranked=best-.002 and beststep!=s:stale+=1 write('status.json',{'status':'training','arm':name,'step':s,'elapsed_seconds':time.monotonic()-START,'production_approved':False}) DURABLE.sync(name+' checkpoint '+str(s)) if stale>=6: print('EARLY_STOP development score unchanged for six evaluations',flush=True);break if not last:raise RuntimeError('Budget exhausted before any final training update') write(name+'/history.json',history) torch.save({'optimizer':opt.state_dict(),'step':last,'seed':SEED,'torch_rng':torch.get_rng_state(),'cuda_rng':torch.cuda.get_rng_state_all(),'exact_resume':False},OUT/name/'optimizer.pt') selected[name]={'step':beststep,'rank_score':best,'parameters':counts[head],'steps_completed':last,'warmup_steps':warmup,'seconds':time.monotonic()-arm_start} write('candidates.json',selected);DURABLE.sync(name+' selected checkpoint on development only') del m,opt;torch.cuda.empty_cache() if last