Download experiments/recurrent-refinement-20260928/train.py from User-2468/mini-unet-colorizer: direct link, hf CLI and curl.
- Browser
- Download file 18.5 kB
-
https://huggingface.co/User-2468/mini-unet-colorizer/resolve/7ec6b1210fd5af800b5f8a5628a856ca0098b0a5/experiments/recurrent-refinement-20260928/train.py
- Command line
-
hf download hf://User-2468/mini-unet-colorizer@7ec6b1210fd5af800b5f8a5628a856ca0098b0a5/experiments/recurrent-refinement-20260928/train.py
-
curl -L -o train.py https://huggingface.co/User-2468/mini-unet-colorizer/resolve/7ec6b1210fd5af800b5f8a5628a856ca0098b0a5/experiments/recurrent-refinement-20260928/train.py
18.5 kB
| # /// script | |
| # requires-python = ">=3.11" | |
| # dependencies = ["torch==2.8.0", "torchvision==0.23.0", "huggingface-hub>=0.34,<2", "safetensors", "numpy<3", "scipy", "scikit-image", "opencv-python-headless", "pillow", "pyarrow", "requests", "datasets"] | |
| # /// | |
| import os,sys,time,json,hashlib,shutil,random,math,traceback,io | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| from torch import nn | |
| import torch.nn.functional as F | |
| from torch.utils.data import Dataset,DataLoader | |
| from PIL import Image | |
| from skimage.color import rgb2lab | |
| import cv2,pyarrow as pa,pyarrow.parquet as pq | |
| from huggingface_hub import HfApi,hf_hub_download | |
| REPO='User-2468/mini-unet-colorizer' | |
| BASE='1a9eb8af2754ad2329a24cfe50d388cb559441d0' | |
| OLD='experiments/region-coherence-20260928' | |
| PREFIX='experiments/recurrent-refinement-20260928' | |
| WORK=Path('/work');WORK.mkdir(exist_ok=True);os.chdir(WORK) | |
| START=time.monotonic();SEED=92872;STEPS=2000;BATCH=24 | |
| api=HfApi(token=os.environ['HF_TOKEN']) | |
| assert api.whoami()['name']=='User-2468' | |
| HEAD=api.model_info(REPO).sha | |
| api.upload_file(repo_id=REPO,path_in_repo=PREFIX+'/access_check.json',path_or_fileobj=json.dumps({'status':'write_access_verified','base_revision':BASE}).encode(),commit_message='Verify experiment artifact write access') | |
| print('WRITE_ACCESS_VERIFIED',flush=True) | |
| source=WORK/'source';source.mkdir(exist_ok=True) | |
| def download(f,rev=HEAD):return Path(hf_hub_download(REPO,f,revision=rev)) | |
| sums=json.loads(download(OLD+'/SHA256SUMS.json').read_text()) | |
| def verified(f): | |
| p=download(OLD+'/'+f) | |
| assert hashlib.sha256(p.read_bytes()).hexdigest()==sums[f],f | |
| return p | |
| for f in ['data.py','legacy_eval.py','metrics.py','model.py','objectives.py','persistence.py','semantic_model.py','spatial.py']: | |
| shutil.copy2(verified('source/'+f),source/f) | |
| shutil.copy2(verified('data_and_eval.py'),WORK/'r7.py') | |
| sys.path.insert(0,str(source));sys.path.insert(0,str(WORK)) | |
| RECURRENT_SOURCE="\nimport json\nfrom pathlib import Path\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nfrom safetensors.torch import load_file,save_file\nfrom semantic_model import SemanticColorizer\n\nclass RecurrentColorizer(SemanticColorizer):\n def __init__(self,steps=4,feedback=True,**kwargs):\n super().__init__(**kwargs)\n self.steps=steps;self.feedback=feedback\n self.config.update(architecture='RecurrentColorizer',steps=steps,feedback=feedback,format_version=2)\n def forward(self,L,return_all=False):\n h,w=L.shape[-2:]\n x=F.pad(self.neutral_rgb(L),(0,(-w)%32,0,(-h)%32),mode='replicate')\n x=(x-self.rgb_mean)/self.rgb_std\n features=[]\n for i,layer in enumerate(self.encoder):\n x=layer(x)\n if i in [3,6,12,16]:features.append(x)\n projected=[layer(f) for layer,f in zip(self.lateral,features)]\n x=projected[-1]\n for i in range(2,-1,-1):\n x=self.refine[2-i](F.interpolate(x,size=projected[i].shape[-2:],mode='bilinear',align_corners=False)+projected[i])\n memory=self.memory_norm(torch.cat([F.adaptive_avg_pool2d(f,(8,8)).flatten(2).transpose(1,2) for f in projected[1:]],1))\n q=self.queries[None].expand(L.shape[0],-1,-1)\n px=self.pixel(x);outputs=[];weights=None\n for step in range(self.steps):\n mem=memory\n if step and self.feedback:\n mass=weights.flatten(2).sum(-1,keepdim=True).clamp_min(1e-5)\n pooled=torch.einsum('bqn,bnd->bqd',weights.flatten(2),px.float().flatten(2).transpose(1,2))/mass\n mem=torch.cat([memory,self.memory_norm(pooled.to(memory.dtype))],1)\n old=q\n for block in self.query_blocks:q=block(q,mem)\n if step:q=old+.25*(q-old)\n nq=self.query_norm(q)\n palette=80*torch.tanh(self.palette(nq))\n masks=torch.einsum('bqd,bdhw->bqhw',nq,px)/(self.config['width']**.5)\n weights=F.softmax(masks.float(),dim=1)\n ab=torch.einsum('bqhw,bqc->bchw',weights,palette.float())+2*torch.tanh(self.residual(x).float())\n outputs.append(F.interpolate(ab,size=(x.shape[-2]*4,x.shape[-1]*4),mode='bilinear',align_corners=False)[...,:h,:w])\n return outputs if return_all else outputs[-1]\n\ndef load_recurrent(path,device='cpu'):\n path=Path(path);cfg=json.loads((path/'config.json').read_text())\n m=RecurrentColorizer(**{k:cfg[k] for k in ['head','width','queries','steps','feedback']})\n m.load_state_dict(load_file(str(path/'model.safetensors')))\n return m.to(device).eval()\n" | |
| (source/'recurrent_model.py').write_text(RECURRENT_SOURCE) | |
| import r7 as r | |
| from data import decode_image | |
| from semantic_model import load_semantic,save_semantic | |
| from recurrent_model import RecurrentColorizer,load_recurrent | |
| OUT=r.OUT;r.persistence.PREFIX=PREFIX | |
| durable=r.persistence.DurableRun(OUT);r.ev.deadline=lambda:None | |
| initial=WORK/'initial';initial.mkdir(exist_ok=True) | |
| for f in ['model.safetensors','config.json']:shutil.copy2(download(f,BASE),initial/f) | |
| torch.set_num_threads(6);torch.manual_seed(SEED);np.random.seed(SEED);random.seed(SEED) | |
| torch.backends.cuda.matmul.allow_tf32=True;torch.backends.cudnn.benchmark=True | |
| class TrainPhotos(Dataset): | |
| def __init__(self,rgb,targets,regions): | |
| self.rgb,self.targets,self.regions=rgb,targets,regions | |
| def __len__(self):return len(self.rgb) | |
| def __getitem__(self,k): | |
| rgb=self.rgb[k].astype(np.float32)/255 | |
| if random.random()<.5: | |
| g=rgb@np.random.dirichlet([3.,5.,2.]).astype(np.float32) | |
| g=np.clip((g**random.uniform(.8,1.25)-.5)*random.uniform(.75,1.15)+.5+random.uniform(-.05,.05),0,1) | |
| if random.random()<.4:g=cv2.GaussianBlur(g,(5,5),random.uniform(.3,1.0)) | |
| if random.random()<.4:g=np.clip(g+np.random.normal(0,random.uniform(.002,.012),g.shape),0,1) | |
| else:g=np.asarray(Image.fromarray(self.rgb[k]).convert('L'),np.float32)/255 | |
| L=rgb2lab(np.repeat(g[...,None],3,-1).astype(np.float32))[...,0][None]/50-1 | |
| target=self.targets[k].astype(np.float32);regions=self.regions[k].astype(np.int64) | |
| if random.random()<.5:L=L[...,::-1];target=target[...,::-1];regions=regions[...,::-1] | |
| return torch.from_numpy(L.copy()),torch.from_numpy(target.copy()),torch.from_numpy(regions.copy()) | |
| def save(m,path): | |
| save_semantic(m,path) | |
| if isinstance(m,RecurrentColorizer): | |
| (Path(path)/'README.md').write_text('# Experimental recurrent colorizer\n\nRequires source/recurrent_model.py and source/semantic_model.py. Load with load_recurrent(path). Parameters include encoder: 3,994,676. Shared decoder runs four passes by default; encoder runs once. Not production approved.\n') | |
| def main(): | |
| shutil.copytree(source,OUT/'source',dirs_exist_ok=True,ignore=shutil.ignore_patterns('__pycache__')) | |
| shutil.copy2(__file__,OUT/'train.py');shutil.copy2(WORK/'r7.py',OUT/'data_and_eval.py') | |
| protocol={'base_revision':BASE,'cache_revision':HEAD,'parameters_including_encoder':3994676,'seed':SEED,'steps_target':STEPS,'batch_size':BATCH, | |
| 'arms':['one_pass_control','recurrent_four_pass'], | |
| 'architecture':'Reuse pretrained MobileNetV3 and spatial decoder once. Repeat SAME two QueryBlocks; after first pass apply damped residual update .25. Pool spatial features using current color-region weights and append normalized pooled tokens to memory. No extra parameters.', | |
| 'loss':'Same fixed selected whole-image targets and losses from region pilot. Control single output; recurrent mean loss across all four passes. No region auxiliary or critic. Matched input batches, initial weights, optimizer and update count.', | |
| 'learning_rates':{'encoder':5e-6,'decoder':1e-4},'evaluation':'Fresh 96-image final holdout excluding earlier evaluation IDs and byte-identical training images. Development reused, explicitly nonindependent. 1/2/4 trained-depth and8 extrapolation. Equal-guided-chroma processing. Raw metrics also retained.', | |
| 'limitations':['Single seed short architecture pilot','8 passes beyond trained depth may worsen','Fixed targets inherit earlier model selection bias','Does not test diffusion or demonstrate architectural ceiling','No guarantees extra passes help','Near duplicates and upstream pretraining overlap unknown'], | |
| 'research':['https://arxiv.org/abs/1806.02919','https://arxiv.org/abs/1904.05290','https://arxiv.org/abs/2212.11613','https://arxiv.org/abs/2111.05826'], | |
| 'budget':{'hardware':'l4x1','hard_timeout_minutes':30,'maximum_hardware_usd_at_published_rate':.40,'training_graceful_stop_minutes':21}, | |
| 'release':'Experimental checkpoints to main subfolder, visual review required before production'} | |
| r.write('protocol.json',protocol) | |
| durable.sync('protocol and sources before GPU training') | |
| # Validate equivalent first-pass computation on CPU before enabling GPU kernels. | |
| baseline=load_semantic(initial,'cpu').requires_grad_(False) | |
| control=load_semantic(initial,'cpu') | |
| recurrent=RecurrentColorizer(steps=4).cpu() | |
| recurrent.load_state_dict(control.state_dict(),strict=True) | |
| assert sum(p.numel() for p in recurrent.parameters())==3994676 | |
| recurrent.eval() | |
| checks={} | |
| with torch.no_grad(): | |
| for h,w in [(96,128),(97,113)]: | |
| x=torch.rand(2,1,h,w)*2-1 | |
| ys=recurrent(x,return_all=True);bp=baseline(x) | |
| checks[f'cpu_{h}x{w}_max_abs']=float((ys[0]-bp).abs().max()) | |
| assert torch.allclose(ys[0],bp,atol=2e-5),checks | |
| assert all(torch.isfinite(y).all() for y in ys) | |
| baseline=baseline.cuda();control=control.cuda();recurrent=recurrent.cuda() | |
| x=torch.rand(2,1,96,128,device='cuda')*2-1 | |
| # Strict parity checks require full precision, not approximate TF32 kernels. | |
| torch.backends.cuda.matmul.allow_tf32=False | |
| torch.backends.cudnn.allow_tf32=False | |
| torch.backends.cudnn.benchmark=False | |
| torch.backends.cudnn.deterministic=True | |
| with torch.no_grad(): | |
| ys=recurrent(x,return_all=True);bp=baseline(x) | |
| checks['gpu_fp32_max_abs']=float((ys[0]-bp).abs().max()) | |
| print('PARITY_CHECKS',json.dumps(checks),flush=True) | |
| assert torch.allclose(ys[0],bp,atol=2e-5),checks | |
| recurrent(x).square().mean().backward() | |
| assert all(torch.isfinite(p.grad).all() for p in recurrent.parameters() if p.grad is not None) | |
| recurrent.zero_grad(set_to_none=True) | |
| save(recurrent,OUT/'initial_recurrent') | |
| test=load_recurrent(OUT/'initial_recurrent','cuda') | |
| with torch.no_grad():assert torch.allclose(test(x),recurrent(x),atol=1e-5) | |
| del test | |
| r.write('smoke.json',{'passed':True,'first_pass_matches_v3':True,'parity':checks,'finite_recurrent_backward':True,'reload_equal':True,'parameter_count':3994676}) | |
| durable.sync('CPU and FP32 GPU parity, backward and reload verified') | |
| print('PREFLIGHT_PASSED',flush=True) | |
| torch.backends.cuda.matmul.allow_tf32=True | |
| torch.backends.cudnn.allow_tf32=True | |
| torch.backends.cudnn.benchmark=True | |
| torch.backends.cudnn.deterministic=False | |
| 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(verified('manifest.json').read_text());ids=manifest['train'] | |
| targets=np.empty((len(ids),2,64,64),np.float16);regions=np.empty((len(ids),64,64),np.uint8) | |
| for start in range(0,len(ids),1024): | |
| with np.load(verified(f'region_cache/{start:05d}.npz')) as z: | |
| assert z['indices'].tolist()==ids[start:start+len(z['indices'])] | |
| targets[start:start+len(z['indices'])]=z['target'];regions[start:start+len(z['indices'])]=z['regions'] | |
| rgbcache=np.stack([np.asarray(decode_image(table,i).convert('RGB').resize((256,256),Image.Resampling.BILINEAR)) for i in ids]) | |
| fresh=pq.read_table(hf_hub_download('detection-datasets/coco','default/val/0001.parquet',repo_type='dataset',revision=manifest['coco_revision']),columns=['image']) | |
| excludes=set(manifest['development_regression']+manifest['heldout_regression']+manifest['new_holdout']) | |
| oldeval=json.loads(download('experiments/round7-20260928/manifest.json',BASE).read_text()) | |
| excludes.update(oldeval['development']+oldeval['heldout']) | |
| oldtrain=json.loads(download('experiments/region-coherence-20260928/source/previous_manifest.json').read_text())['train'] | |
| hashes={hashlib.sha256(table['image'][i].as_py()['bytes']).digest() for i in oldtrain} | |
| hold=[int(i) for i in np.random.default_rng(SEED).permutation(len(fresh)) if int(i) not in excludes and hashlib.sha256(fresh['image'][int(i)].as_py()['bytes']).digest() not in hashes][:96] | |
| assert len(hold)==96 | |
| dev=manifest['development_regression'][:32] | |
| r.write('manifest.json',{'train':ids,'development_regression':dev,'fresh_holdout':hold,'source_manifest_revision':HEAD,'seed':SEED}) | |
| models={'control':control,'recurrent':recurrent} | |
| opts={n:torch.optim.AdamW([{'params':m.encoder.parameters(),'lr':5e-6,'base_lr':5e-6},{'params':[p for k,p in m.named_parameters() if not k.startswith('encoder.')],'lr':1e-4,'base_lr':1e-4}],weight_decay=.01) for n,m in models.items()} | |
| loader=DataLoader(TrainPhotos(rgbcache,targets,regions),batch_size=BATCH,shuffle=True,num_workers=4,pin_memory=True,drop_last=True,persistent_workers=True,generator=torch.Generator().manual_seed(SEED)) | |
| it=iter(loader);history=[];last=0 | |
| def persist(step): | |
| for n,m in models.items(): | |
| save(m,OUT/n/f'step_{step:05d}') | |
| torch.save({'step':step,'optimizer':opts[n].state_dict(),'torch_rng':torch.get_rng_state(),'cuda_rng':torch.cuda.get_rng_state_all(),'numpy_rng':np.random.get_state(),'python_rng':random.getstate(),'exact_dataloader_resume':False},OUT/n/'optimizer.pt') | |
| r.write('history.json',history);r.write('status.json',{'status':'training','step':step,'elapsed_seconds':time.monotonic()-START,'production_approved':False}) | |
| durable.sync('matched recurrent trial step '+str(step)) | |
| durable.sync('cached targets verified and holdout frozen') | |
| for step in range(1,STEPS+1): | |
| if time.monotonic()-START>21*60:break | |
| try:L,t,_=next(it) | |
| except StopIteration:it=iter(loader);L,t,_=next(it) | |
| L,t=L.cuda(non_blocking=True),t.cuda(non_blocking=True);row={'step':step} | |
| for n,m in models.items(): | |
| m.train() | |
| for layer in m.encoder.modules(): | |
| if isinstance(layer,nn.BatchNorm2d):layer.eval() | |
| opts[n].zero_grad(set_to_none=True) | |
| factor=min(step/100,1)*(.2+.8*(1+math.cos(math.pi*step/STEPS))/2) | |
| for group in opts[n].param_groups:group['lr']=group['base_lr']*factor | |
| with torch.autocast('cuda',dtype=torch.bfloat16): | |
| ps=m(L,return_all=True) if n=='recurrent' else [m(L)] | |
| terms=torch.stack([r.losses(p,t).mean() for p in ps]);loss=terms.mean() | |
| assert torch.isfinite(loss) | |
| loss.backward();torch.nn.utils.clip_grad_norm_(m.parameters(),1,error_if_nonfinite=True);opts[n].step() | |
| row[n]=[float(v.detach()) for v in terms] | |
| last=step | |
| if step%100==0: | |
| row['elapsed_seconds']=time.monotonic()-START;history.append(row);print('STEP',json.dumps(row),flush=True) | |
| if step%500==0:persist(step) | |
| assert last>0,'No training budget remained' | |
| persist(last) | |
| for n,m in models.items():save(m,OUT/n/'candidate') | |
| class View(nn.Module): | |
| def __init__(self,m,steps):super().__init__();self.m=m;self.steps=steps | |
| def forward(self,x): | |
| old=self.m.steps;self.m.steps=self.steps | |
| try:return self.m(x) | |
| finally:self.m.steps=old | |
| def decode(self,x,temperature=.38):return x | |
| views={'v3':baseline,'control':control} | |
| views.update({f'recurrent_{s}':View(recurrent,s) for s in [1,2,4,8]}) | |
| for n,m in views.items(): | |
| m.eval() | |
| r.write(n+'_fresh.json',r.ev.evaluate(m,fresh,hold)) | |
| r.ev.probes(m,n+'_probes') | |
| x=torch.zeros(1,1,256,256,device='cuda') | |
| with torch.inference_mode(): | |
| for _ in range(3):m(x) | |
| torch.cuda.synchronize();a=time.monotonic() | |
| for _ in range(20):m(x) | |
| torch.cuda.synchronize() | |
| r.write(n+'_latency.json',{'batch':1,'size':256,'precision':'float32','ms_per_image':(time.monotonic()-a)*1000/20,'hardware':torch.cuda.get_device_name()}) | |
| durable.sync(n+' final metrics and probes') | |
| r.grid({'v3':baseline,'control':control,'recurrent_4':views['recurrent_4']},fresh,hold,'fresh_architecture') | |
| r.grid({n:views[n] for n in ['recurrent_1','recurrent_2','recurrent_4','recurrent_8']},fresh,hold,'fresh_passes') | |
| r.grid({'v3':baseline,'control':control,'recurrent_4':views['recurrent_4']},fresh,dev,'development') | |
| import requests | |
| for name,url in {'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'}.items(): | |
| try: | |
| response=requests.get(url,timeout=30);response.raise_for_status() | |
| tb=pa.table({'image':[{'bytes':response.content,'path':None}]}) | |
| r.grid({'v3':baseline,'control':control,'recurrent_4':views['recurrent_4']},tb,[0],name) | |
| r.write(name+'_source.json',{'url':url,'sha256':hashlib.sha256(response.content).hexdigest(),'previously_seen_probe':True}) | |
| except Exception as e:r.write(name+'_error.json',{'type':type(e).__name__}) | |
| r.write('status.json',{'status':'completed','steps':last,'elapsed_seconds':time.monotonic()-START,'production_approved':False,'next':'Inspect images and 1/2/4/8 pass progression. Eight passes are extrapolation, not promised improvement.'}) | |
| durable.sync('recurrent experiment completed with visual comparisons') | |
| if __name__=='__main__': | |
| try:main() | |
| except BaseException: | |
| msg=traceback.format_exc().replace(os.environ.get('HF_TOKEN','__NONE__'),'[REDACTED]') | |
| r.write('failure.json',{'traceback':msg}) | |
| try:durable.sync('failure report') | |
| except Exception:pass | |
| print(msg,flush=True) | |
| raise SystemExit(1) | |