IMM Stale-Feedback Detector / nn_signature.py

Failed on benchmark

Raw ⬇ ZIP
 1import sys, math, json, random
 2import numpy as np
 3import torch
 4import torch.nn as nn
 5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 6from bench import get_dataset, make_model
 7from bench_imm import IMMMonitor, seed_all, vec_norm, EPOCHS, BATCH
 8
 9def run(seed, delay=0):
10    seed_all(seed); ds=get_dataset('tabular',seed,n_train=400,n_test=400)
11    net=make_model('mlp_tiny',ds['input_shape'],ds['out_dim']); opt=torch.optim.Adam(net.parameters(),lr=.01)
12    dev='cuda' if torch.cuda.is_available() else 'cpu'
13    try:
14        net.to(dev); x,y=ds['xtr'].to(dev),ds['ytr'].to(dev); mon=IMMMonitor(); queued=[]; p=[]; alarms=[]; inn=[]
15        for ep in range(8):
16            net.train(); perm=torch.randperm(len(x),device=dev)
17            for j in range(0,len(x),BATCH):
18                idx=perm[j:j+BATCH]; loss=nn.MSELoss()(net(x[idx]),y[idx]); opt.zero_grad(); loss.backward()
19                grads=[None if q.grad is None else q.grad.detach().clone() for q in net.parameters()]
20                queued.append(grads)
21                use=queued[max(0,len(queued)-1-delay)]
22                for q,g in zip(net.parameters(),use):
23                    if g is not None: q.grad=g
24                gn=math.sqrt(sum(float((g**2).sum()) for g in use if g is not None)); pn=vec_norm(net)
25                z=np.log1p(np.array([float(loss),float(loss),gn,pn,.01*gn]))
26                al=mon.step(z); alarms.append(al); p.append(mon.pi.copy()); inn.append(mon.last['innovation'])
27                opt.step()
28        p=np.asarray(p); return {'delay':delay,'final_posterior':p[-1].tolist(),'min_no_delay_posterior':float(p[:,0].min()),'alarm_count':int(sum(alarms)),'first_alarm':next((i+1 for i,a in enumerate(alarms) if a),None),'mean_innovation':float(np.mean(inn))}
29    except RuntimeError:
30        return {'delay':delay,'error':'runtime failure'}
31
32def main():
33    clean=[run(s,0) for s in range(8)]; stale=[run(s,2) for s in range(8)]
34    out={'clean':clean,'stale_delay_2':stale,'summary':{'clean_alarm_rate':float(np.mean([x['alarm_count']>0 for x in clean])),'stale_alarm_rate':float(np.mean([x['alarm_count']>0 for x in stale])),'stale_first_alarm_median':float(np.median([x['first_alarm'] for x in stale if x['first_alarm'] is not None])) if any(x['first_alarm'] is not None for x in stale) else None}}
35    print(json.dumps(out,indent=2))
36if __name__=='__main__': main()