File size: 5,876 Bytes
71d64bb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 | """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()
|