User-2468's picture
Colorizer round6: protocol and sources before GPU training
25b9614 verified
Raw History Blame
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)