IMM Stale-Feedback Detector / nn_signature.py
Failed on benchmark
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()