Download main_method/code/train_feature_clusters.py from Dhruv1000/TSRDA: direct link, hf CLI and curl.
- Browser
- Download file 18.9 kB
-
https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/train_feature_clusters.py
- Command line
-
hf download hf://Dhruv1000/TSRDA/main_method/code/train_feature_clusters.py
-
curl -L -o train_feature_clusters.py https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/train_feature_clusters.py
18.9 kB
| """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() | |