summaryrefslogtreecommitdiff
path: root/experiments/calib_tent.py
blob: a09fe931aff4e6b239ce44d7a7c5fd9499a6a233 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
import torch, collections, sys
from sdil.data import make_tentmap, onehot
from sdil.core import SDILNet, SDILConfig, sdil_step
from sdil.baselines import BPNet, dfa_config, evaluate
dev='cpu'
L=int(sys.argv[1]) if len(sys.argv)>1 else 7
tr,te,ni,no=make_tentmap(levels=L,n_in=3,n_train=50000,n_test=8000,seed=0,batch_size=128,device=dev)
cnt=collections.Counter(te.y.tolist()); print(f"tent levels={L} balance {dict(cnt)}",flush=True)
for W in [6,10,16]:
  print(f"--- width={W} PLAIN relu (residual=0) ---",flush=True)
  for d in [1,2,3,4,6,8]:
    netb=BPNet([ni]+[W]*d+[2],act='relu',device=dev,seed=1,residual=False)
    netd=SDILNet([ni]+[W]*d+[2],act='relu',device=dev,seed=1,residual=False)
    cfg=dfa_config(eta=0.02,momentum=0.9); s=0
    for e in range(20):
      for x,y in tr:
        netb.bp_step(x,y,0.02,0.9); sdil_step(netd,x,y,onehot(y,2),cfg,s); s+=1
    print(f"  depth={d:2d}: BP {evaluate(netb,te)[0]:.3f}  DFA {evaluate(netd,te)[0]:.3f}",flush=True)