"""Frozen-depth probes and patch-balanced prototype readouts on official fixed32. Caches only masked temporal mean/std summaries, not full time-series features. The readout sees 18 classes; probabilities are mapped to the original 20 channels. Validation is selection-only. Prototype banks and normalization use training only. """ import argparse import json import time from pathlib import Path import numpy as np import torch import torch.nn.functional as F from torch.utils.data import DataLoader from sklearn.cluster import KMeans import config as C from dataset import (PASTISPatchDataset, train_patch_ids, val_patch_ids, test_patch_ids, pad_collate, crop_ij) from model import UTAE_ASPP, UTAEClassificationDual from temporal_cnn import valid_from_index_positions ROOT = Path(__file__).resolve().parents[2] OUT = ROOT / 'main_method' / 'runs' / 'feature_cluster_readouts' SAM = ROOT / 'main_method/checkpoints/shared/sam_network_seed3407.tar' PRETRAIN = ROOT / 'checkpoints/pretrain/latest.utae.tar' def write_json(path, data): path.parent.mkdir(parents=True, exist_ok=True) temp = path.with_suffix('.json.tmp') temp.write_text(json.dumps(data, indent=2)) temp.replace(path) def metrics(y, pred): valid = (y > 0) & (y < 19) cm = np.bincount(20 * y[valid].astype(np.int64) + pred[valid], minlength=400).reshape(20, 20) return metrics_cm(cm) def metrics_cm(cm): tp = np.diag(cm)[1:19].astype(float) gt, pp = cm.sum(1)[1:19], cm.sum(0)[1:19] iou = tp / np.maximum(gt + pp - tp, 1) f1 = 2 * tp / np.maximum(gt + pp, 1) total = cm.sum() oa = cm.trace() / max(total, 1) chance = np.dot(cm.sum(0).astype(float), cm.sum(1)) / max(total ** 2, 1) return {'miou': float(iou.mean()), 'oa': float(oa), 'macro_f1': float(f1.mean()), 'kappa': float((oa-chance)/max(1-chance,1e-12)), 'per_class_iou': iou.tolist(), 'confusion_matrix': cm.tolist()} def frozen_model(source, device): enc = UTAE_ASPP(n_channels=C.N_CHANNELS, d_model=C.D_MODEL, n_heads=C.N_HEADS) net = UTAEClassificationDual(enc, d_model=C.D_MODEL, num_classes=C.N_CLASSES) ck = torch.load(SAM if source == 'sam' else PRETRAIN, map_location='cpu', weights_only=False) state = ck['model_state_dict'] if source == 'sam': net.load_state_dict(state, strict=True) else: if not any(k.startswith('spatial_encoder.') for k in state): state = {k[5:]: v for k, v in state.items() if k.startswith('utae.')} enc.load_state_dict(state, strict=True) return net.to(device).eval().requires_grad_(False) class DepthCollector: def __init__(self, encoder): self.handles = [block.register_forward_hook(self.hook(i)) for i, block in enumerate(encoder.temporal_encoder.blocks)] def begin(self, x, pos): b, _, _, h, w = x.shape self.shape = b, h, w self.valid = valid_from_index_positions(pos)[:, None, None] self.valid = self.valid.expand(b, h, w, -1).reshape(b*h*w, -1) self.parts = [[], [], []] self.offsets = [0, 0, 0] def hook(self, depth): def collect(module, inputs, output): start = self.offsets[depth] valid = self.valid[start:start + output.shape[0]] if valid.shape[1] < output.shape[1]: valid = F.pad(valid, (0, output.shape[1]-valid.shape[1]), value=False) weights = valid[..., None].float() values = output.float() count = weights.sum(1).clamp_min(1) mean = (values * weights).sum(1) / count var = ((values - mean[:, None]).square() * weights).sum(1) / count summary = torch.cat((mean, var.clamp_min(0).sqrt()), -1) self.parts[depth].append(summary.detach().half().cpu().numpy()) self.offsets[depth] += len(output) return collect def end(self): assert self.offsets == [np.prod(self.shape)] * 3 return np.stack([np.concatenate(parts) for parts in self.parts], axis=1) def close(self): for handle in self.handles: handle.remove() def extract_cache(): device = torch.device('cuda') started = time.monotonic() manifest = {'feature': 'valid-timestep mean/std at blocks 1,2,3; float16 summaries', 'crop_size': 32, 'batch': 2, 'sources': {}, 'normalizer': 'original', 'positions': 'original acquisition indices'} for source in ('sam', 'pretrain'): model = frozen_model(source, device) collector = DepthCollector(model.utae) for split, ids in [('train', train_patch_ids()), ('val', val_patch_ids())]: folder = OUT/'cache'/source/split folder.mkdir(parents=True, exist_ok=True) n = len(ids) * 2 * 32 * 32 files = { 'features': ((n, 3, 512), np.float16), 'labels': ((n,), np.uint8), 'patch_ids': ((n,), np.int32), 'crop_ids': ((n,), np.uint8), } if source == 'sam': files['base_probs'] = ((n, 20), np.float32) arrays = {k: np.lib.format.open_memmap(folder/f'{k}.npy',mode='w+',dtype=dtype,shape=shape) for k,(shape,dtype) in files.items()} loader = DataLoader(PASTISPatchDataset(ids), batch_size=2, shuffle=False, num_workers=2, pin_memory=True, collate_fn=pad_collate) offset = 0 with torch.inference_mode(): for bi, ((x, pos, _), y) in enumerate(loader): x, pos = x.to(device), pos.to(device) batch_ids = ids[2*bi:2*bi+x.shape[0]] for crop_index, (i, j) in enumerate(C.FT_CROP_IJ): crop, target = crop_ij(x, y, i, j, C.PRETRAIN_CROP_GRID) collector.begin(crop, pos) with torch.autocast('cuda', dtype=torch.float16): result = model(crop, pos) if source == 'sam' else model.utae(crop,pos) features = collector.end() end = offset + len(features) arrays['features'][offset:end] = features arrays['labels'][offset:end] = target.numpy().reshape(-1) arrays['patch_ids'][offset:end] = np.repeat(batch_ids,32*32) arrays['crop_ids'][offset:end] = crop_index if source == 'sam': pp = result['refined_logits'].float().softmax(1) arrays['base_probs'][offset:end] = pp.permute(0,2,3,1).reshape(-1,20).cpu().numpy() offset = end if bi % 10 == 0: print('cache',source,split,bi,len(loader),flush=True) assert offset == n for array in arrays.values(): array.flush() manifest['sources'][f'{source}_{split}'] = {'pixels':n,'entries':len(ids), 'unique_patches':len(set(ids))} del arrays collector.close() del model torch.cuda.empty_cache() y = np.load(OUT/'cache/sam/val/labels.npy') p = np.load(OUT/'cache/sam/val/base_probs.npy',mmap_mode='r') manifest['sam_baseline'] = metrics(y,p.argmax(1)) if abs(manifest['sam_baseline']['miou']-.5000673383500418)>1e-6: raise RuntimeError('Feature-cache SAM baseline does not reproduce its stored score') manifest['seconds'] = time.monotonic()-started write_json(OUT/'cache/manifest.json',manifest) def load_cache(source, split): folder = OUT/'cache'/source/split return {name:np.load(folder/f'{name}.npy',mmap_mode='r') for name in ('features','labels','patch_ids','crop_ids')} def rows(features, indices, depth): data = np.asarray(features[indices]) return (data[:,depth,:] if depth>=0 else data.reshape(len(data),-1)).astype(np.float32) def train_normalizer(features, indices, depth): width = 512 if depth>=0 else 1536 sums, squares = np.zeros(width,np.float64),np.zeros(width,np.float64) for start in range(0,len(indices),4096): x = rows(features,indices[start:start+4096],depth).astype(np.float64) sums += x.sum(0);squares += (x*x).sum(0) mean = sums/len(indices) std = np.sqrt(np.maximum(squares/len(indices)-mean*mean,0)).clip(.05) return torch.from_numpy(mean.astype(np.float32)).cuda(),torch.from_numpy(std.astype(np.float32)).cuda() def normalized(features, indices, depth, mean, std): return (torch.from_numpy(rows(features,indices,depth)).cuda()-mean)/std def make_prototypes(train, depth, mean, std, seed): rng = np.random.default_rng(seed) # Only the first occurrence of each (patch,crop) enters prototype memory. pid, cid, labels = train['patch_ids'],train['crop_ids'],train['labels'] unique = np.zeros(len(labels),bool);seen=set() for start in range(0,len(labels),1024): key=int(pid[start]),int(cid[start]) if key not in seen: unique[start:start+1024]=True;seen.add(key) protos=[] for cls in range(1,19): memory=[] for patch in np.unique(pid[unique & (labels==cls)]): eligible=np.flatnonzero(unique & (pid==patch) & (labels==cls)) memory.extend(rng.choice(eligible,32,replace=len(eligible)<32).tolist()) if not memory: raise RuntimeError(f'No training prototype support for class {cls}') embeddings=F.normalize(normalized(train['features'],np.array(memory),depth,mean,std),dim=1) km=KMeans(n_clusters=3,n_init=3,random_state=seed).fit(embeddings.cpu().numpy()) centers=F.normalize(torch.from_numpy(km.cluster_centers_).cuda(),dim=1) protos.append(centers) return torch.stack(protos) def fit_linear(train, depth, mean, std, seed): torch.manual_seed(seed) width=len(mean) head=torch.nn.Linear(width,18).cuda() optimizer=torch.optim.AdamW(head.parameters(),lr=1e-3,weight_decay=1e-3) rng=np.random.default_rng(seed) eligible=np.flatnonzero((train['labels']>0)&(train['labels']<19)) for epoch in range(30): shuffled=rng.permutation(eligible) for start in range(0,len(shuffled),4096): ix=shuffled[start:start+4096] x=normalized(train['features'],ix,depth,mean,std) y=torch.from_numpy(np.asarray(train['labels'][ix]).astype(np.int64)-1).cuda() optimizer.zero_grad(set_to_none=True) F.cross_entropy(head(x),y).backward();optimizer.step() return {k:v.detach() for k,v in head.state_dict().items()} def probability(features, depth, mean, std, kind, state): out=np.zeros((len(features),20),np.float32) with torch.inference_mode(): for start in range(0,len(features),4096): ix=np.arange(start,min(start+4096,len(features))) x=normalized(features,ix,depth,mean,std) if kind=='linear': logits=F.linear(x,state['weight'],state['bias']) else: embeddings=F.normalize(x,dim=-1) logits=torch.einsum('nd,ckd->nck',embeddings,state)/.1 logits=logits.logsumexp(-1)-np.log(3) out[start:start+len(ix),1:19]=logits.softmax(-1).cpu().numpy() return out def bootstrap(y, pred, baseline, patch_ids): unique=np.unique(patch_ids);a=[];b=[] for pid in unique: m=patch_ids==pid;valid=m&(y>0)&(y<19) a.append(np.bincount(20*y[valid].astype(int)+pred[valid],minlength=400).reshape(20,20)) b.append(np.bincount(20*y[valid].astype(int)+baseline[valid],minlength=400).reshape(20,20)) a,b=np.array(a),np.array(b);rng=np.random.default_rng(3407);d=[] for _ in range(1000): ix=rng.integers(0,len(unique),len(unique)) d.append(100*(metrics_cm(a[ix].sum(0))['miou']-metrics_cm(b[ix].sum(0))['miou'])) return np.percentile(d,[2.5,97.5]).tolist() def fit_readouts(): if not (OUT/'cache/manifest.json').exists(): raise RuntimeError('Incomplete feature cache') started=time.monotonic();records=[];models={};families={} base=np.load(OUT/'cache/sam/val/base_probs.npy',mmap_mode='r') target=np.load(OUT/'cache/sam/val/labels.npy') baseline=metrics(target,base.argmax(1))['miou'] for source in ('sam','pretrain'): train,val=load_cache(source,'train'),load_cache(source,'val') assert np.array_equal(target,val['labels']) eligible=np.flatnonzero((train['labels']>0)&(train['labels']<19)) for depth in (0,1,2,-1): mean,std=train_normalizer(train['features'],eligible,depth) for seed in (3407,42,1234): for kind in ('linear','prototype'): state=fit_linear(train,depth,mean,std,seed) if kind=='linear' else make_prototypes(train,depth,mean,std,seed) pp=probability(val['features'],depth,mean,std,kind,state) name=f'{source}_depth{depth}_{kind}_seed{seed}' path=OUT/'readouts'/f'{name}.tar';path.parent.mkdir(parents=True,exist_ok=True) cpu_state={k:v.cpu() for k,v in state.items()} if kind=='linear' else state.cpu() torch.save({'source':source,'depth':depth,'kind':kind,'seed':seed, 'mean':mean.cpu(),'std':std.cpu(),'state':cpu_state},path) for alpha in (1.,.1,.25,.5): pred=((1-alpha)*base+alpha*pp).argmax(1) m=metrics(target,pred) family=f'{source}_depth{depth}_{kind}_alpha{alpha}' entry={'source':source,'depth':depth,'kind':kind,'seed':seed, 'alpha':alpha,'metrics':m,'checkpoint':str(path),'family':family} records.append(entry);families.setdefault(family,[]).append(entry) print(name,'unblended',records[-4]['metrics']['miou'],flush=True) write_json(OUT/'probe_progress.json',{'completed_readouts':len(records)//4,'total_readouts':48}) del mean,std del train,val ranked=sorted(families.items(),key=lambda kv:np.mean([e['metrics']['miou'] for e in kv[1]]),reverse=True) family,entries=ranked[0] mean_score=float(np.mean([e['metrics']['miou'] for e in entries])) # Fixed seed deployment representative, never choose a seed by test score. chosen=next(e for e in entries if e['seed']==3407) ck=torch.load(chosen['checkpoint'],map_location='cpu',weights_only=False) ck['alpha']=chosen['alpha'];torch.save(ck,OUT/'probe_best.tar') val=load_cache(chosen['source'],'val') state={k:v.cuda() for k,v in ck['state'].items()} if ck['kind']=='linear' else ck['state'].cuda() pp=probability(val['features'],ck['depth'],ck['mean'].cuda(),ck['std'].cuda(),ck['kind'],state) pred=((1-ck['alpha'])*base+ck['alpha']*pp).argmax(1) result={'baseline_miou':baseline,'selected_family':family,'family_mean_val_miou':mean_score, 'family_std_val_miou':float(np.std([e['metrics']['miou'] for e in entries],ddof=1)), 'delta_mean_pp':100*(mean_score-baseline),'selected':chosen, 'paired_patch_delta_ci95_pp':bootstrap(target,pred,base.argmax(1),val['patch_ids']), 'promote':bool(mean_score>=baseline+.005),'records':records, 'seconds':time.monotonic()-started,'test_used_for_selection':False, 'scope':'Readout seeds vary on the same frozen encoder; not full encoder training seeds.'} write_json(OUT/'probe_results.json',result) print('Selected',family,'mean validation',mean_score,'promote',result['promote'],flush=True) def evaluate_selected(split): ck=torch.load(OUT/'probe_best.tar',map_location='cpu',weights_only=False) device=torch.device('cuda');base=frozen_model('sam',device) encoder_model=base if ck['source']=='sam' else frozen_model('pretrain',device) collector=DepthCollector(encoder_model.utae) ids=val_patch_ids() if split=='val' else test_patch_ids() batch=2 if split=='val' else 1 loader=DataLoader(PASTISPatchDataset(ids),batch_size=batch,shuffle=False,num_workers=2, pin_memory=True,collate_fn=pad_collate) state={k:v.cuda() for k,v in ck['state'].items()} if ck['kind']=='linear' else ck['state'].cuda() cm=np.zeros((20,20),np.int64);started=time.monotonic() with torch.inference_mode(): for bi,((x,pos,_),y) in enumerate(loader): x,pos=x.to(device),pos.to(device) grid=4 if split=='val' else 1 crops=C.FT_CROP_IJ if split=='val' else [(0,0)] for i,j in crops: crop,target=crop_ij(x,y,i,j,grid) collector.begin(crop,pos) with torch.autocast('cuda',dtype=torch.float16): output=base(crop,pos) if encoder_model is not base: encoder_model.utae(crop,pos) features=collector.end() bp=output['refined_logits'].float().softmax(1).permute(0,2,3,1).reshape(-1,20).cpu().numpy() pp=probability(features,ck['depth'],ck['mean'].cuda(),ck['std'].cuda(),ck['kind'],state) pred=((1-ck['alpha'])*bp+ck['alpha']*pp).argmax(1) gt=target.numpy().reshape(-1);valid=(gt>0)&(gt<19) cm+=np.bincount(20*gt[valid].astype(int)+pred[valid],minlength=400).reshape(20,20) if bi%25==0:print('probe evaluation',split,bi,len(loader),flush=True) collector.close() result={'metrics':metrics_cm(cm),'seconds':time.monotonic()-started, 'entries':len(ids),'batch':batch,'full_patch':split=='test', 'checkpoint':str(OUT/'probe_best.tar'),'source':ck['source'],'alpha':ck['alpha']} write_json(OUT/f'probe_{split}_results.json',result) if split=='val': expected=json.loads((OUT/'probe_results.json').read_text())['selected']['metrics']['miou'] if abs(result['metrics']['miou']-expected)>1e-6: raise RuntimeError('Online readout validation does not match cache selection') print(result,flush=True) def main(): parser=argparse.ArgumentParser();parser.add_argument('action',choices=['cache','fit','validate','test']) args=parser.parse_args();OUT.mkdir(parents=True,exist_ok=True) torch.set_num_threads(4);torch.manual_seed(3407) torch.backends.cuda.matmul.allow_tf32=True torch.backends.cudnn.allow_tf32=True if args.action=='cache':extract_cache() elif args.action=='fit':fit_readouts() else:evaluate_selected('val' if args.action=='validate' else 'test') if __name__=='__main__':main()