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