TSRDA / main_method /code /evaluate.py
Dhruv1000's picture
Organize complete final models, all ablations, logs and checkpoints with visual guides (part 7)
71d64bb verified
Raw History Blame Contribute Delete
5.88 kB
"""Evaluate complete final bundles without changing the scored input protocol."""
import argparse,hashlib,json,time
from pathlib import Path
import numpy as np
import torch
from torch.utils.data import DataLoader
import config as C
from dataset import PASTISPatchDataset,val_patch_ids,test_patch_ids,pad_collate,crop_ij
from model import UTAE_ASPP,UTAEClassificationDual
from training_feature_clusters import DepthCollector,probability,metrics_cm
from temporal_cnn import TemporalCNN,AttentionCNN
MAIN=Path(__file__).resolve().parents[1]
def digest(path):
h=hashlib.sha256()
with path.open('rb') as f:
for part in iter(lambda:f.read(8*1024*1024),b''):h.update(part)
return h.hexdigest()
def manifest(name):
folder=MAIN/'configurations'/f'32X32_{name}_seed3407'
return json.loads((folder/'manifest.json').read_text()),folder
def load_state(path):
return torch.load(path,map_location='cpu',weights_only=False)
def load_models(names,device,verify=True):
specs={n:manifest(n)[0] for n in names};shared=MAIN/'checkpoints/shared'
for spec in specs.values():
if verify:
for key in ['sam_checkpoint','feature_cluster_checkpoint','original_cnn_checkpoint','addition_checkpoint']:
file=spec[key]
if digest(shared/file)!=spec['sha256'][file]:raise RuntimeError('Checkpoint hash mismatch: '+file)
sam=UTAEClassificationDual(UTAE_ASPP(n_channels=C.N_CHANNELS,d_model=C.D_MODEL,n_heads=C.N_HEADS),d_model=C.D_MODEL,num_classes=C.N_CLASSES)
first=next(iter(specs.values()));sam.load_state_dict(load_state(shared/first['sam_checkpoint'])['model_state_dict'],strict=True)
sam=sam.to(device).eval().requires_grad_(False)
clusters=load_state(shared/first['feature_cluster_checkpoint'])
old=TemporalCNN(18);old.load_state_dict(load_state(shared/first['original_cnn_checkpoint'])['state'],strict=True);old=old.to(device).eval().requires_grad_(False)
additions={}
for name,spec in specs.items():
model=TemporalCNN(18) if spec['addition_kind']=='cnn_control' else AttentionCNN(18)
ck=load_state(shared/spec['addition_checkpoint']);assert ck['seed']==3407 and ck['classes']==list(range(1,19))
model.load_state_dict(ck['state'],strict=True);additions[name]=model.to(device).eval().requires_grad_(False)
return specs,sam,clusters,old,additions
@torch.inference_mode()
def main():
ap=argparse.ArgumentParser()
ap.add_argument('--configuration',choices=['recommended_validation','highest_observed_test','both'],default='recommended_validation')
ap.add_argument('--split',choices=['val','test'],default='test')
ap.add_argument('--output',type=Path,default=MAIN/'evaluation_results.json')
ap.add_argument('--skip_hash_check',action='store_true');args=ap.parse_args()
if not torch.cuda.is_available():raise RuntimeError('The scored inference recipe requires a working CUDA runtime; CPU loading is checked separately.')
torch.set_num_threads(4);torch.manual_seed(3407);torch.backends.cuda.matmul.allow_tf32=True;torch.backends.cudnn.allow_tf32=True
names=['recommended_validation','highest_observed_test'] if args.configuration=='both' else [args.configuration]
specs,sam,clusters,old,additions=load_models(names,'cuda',not args.skip_hash_check)
collector=DepthCollector(sam.utae);mean,std=clusters['mean'].cuda(),clusters['std'].cuda()
state={k:v.cuda() for k,v in clusters['state'].items()} if clusters['kind']=='linear' else clusters['state'].cuda()
ids=val_patch_ids() if args.split=='val' else test_patch_ids();batch=2 if args.split=='val' else 1
dl=DataLoader(PASTISPatchDataset(ids),batch_size=batch,shuffle=False,num_workers=2,pin_memory=True,collate_fn=pad_collate)
matrices={name:np.zeros((20,20),np.int64) for name in names};started=time.monotonic()
try:
for bi,((x,pos,_),y) in enumerate(dl):
x,pos=x.cuda(),pos.cuda()
for i,j in C.FT_CROP_IJ if args.split=='val' else [(0,0)]:
xc,yc=crop_ij(x,y,i,j,4 if args.split=='val' else 1);collector.begin(xc,pos)
with torch.autocast('cuda',dtype=torch.float16):
output=sam(xc,pos);original_logits=old(xc,pos)
additional_logits={name:net(xc,pos) for name,net in additions.items()}
features=collector.end();spatial=output['refined_logits'].float().softmax(1)
pp=probability(features,clusters['depth'],mean,std,clusters['kind'],state)
b,_,h,w=spatial.shape;cluster_probs=torch.from_numpy(pp).cuda().reshape(b,h,w,20).permute(0,3,1,2)
feature_blend=.5*spatial+.5*cluster_probs
original_probs=torch.zeros_like(spatial);original_probs[:,1:19]=original_logits.float().softmax(1)
reference=.75*feature_blend+.25*original_probs
gt=yc.numpy().ravel();valid=(gt>0)&(gt<19)
for name,spec in specs.items():
probs=torch.zeros_like(spatial);probs[:,1:19]=additional_logits[name].float().softmax(1)
alpha=spec['addition_alpha'];prediction=((1-alpha)*reference+alpha*probs).argmax(1).cpu().numpy().ravel()
matrices[name]+=np.bincount(20*gt[valid]+prediction[valid],minlength=400).reshape(20,20)
if bi%50==0:print(args.split,bi,len(dl),flush=True)
finally:collector.close()
result={'split':args.split,'patch_entries':len(ids),'batch_size':batch,'seed':3407,'seconds':time.monotonic()-started,'configurations':{name:{'metrics':metrics_cm(matrices[name]),'manifest':specs[name]} for name in names}}
args.output.parent.mkdir(parents=True,exist_ok=True);args.output.write_text(json.dumps(result,indent=2))
for name,item in result['configurations'].items():print(name,'mIoU',item['metrics']['miou'],flush=True)
if __name__=='__main__':main()