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