TSRDA / main_method /code /train_feature_clusters.py
Dhruv1000's picture
Finalize verified main-method and ablation layout; preserve all original artifacts
59a6762 verified
Raw History Blame Contribute Delete
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()