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