"""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()