Strongly-Rayleigh Forest Dropout / forest_dropout_experiment.py
Beats tuned baseline
1import itertools, json, math, random
2import numpy as np
3import torch
4from torch import nn
5
6SEED = 3020
7
8def set_seed(seed):
9 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
10
11# Six edges of K4. A 3-edge subset is a spanning tree iff it is acyclic.
12EDGES = [(0,1),(0,2),(0,3),(1,2),(1,3),(2,3)]
13
14def is_forest(subset):
15 parent = list(range(4))
16 def find(x):
17 while parent[x] != x:
18 parent[x] = parent[parent[x]]
19 x = parent[x]
20 return x
21 for ei in subset:
22 a,b = EDGES[ei]; ra,rb = find(a),find(b)
23 if ra == rb: return False
24 parent[ra] = rb
25 return True
26
27def enumerate_trees(weights=None):
28 weights = np.ones(6) if weights is None else np.asarray(weights, float)
29 masks=[]; ws=[]
30 for comb in itertools.combinations(range(6), 3):
31 if is_forest(comb):
32 m=np.zeros(6,dtype=np.int64); m[list(comb)]=1
33 masks.append(m); ws.append(float(np.prod(weights[list(comb)])))
34 ws=np.asarray(ws); return np.asarray(masks), ws/ws.sum()
35
36def math_check():
37 masks, prob = enumerate_trees()
38 inc = prob @ masks
39 joint = (prob[:,None,None] * masks[:,:,None] * masks[:,None,:]).sum(0)
40 cov = joint - inc[:,None]*inc[None,:]
41 # For increasing coordinate events {i selected}, negative association is pairwise visible.
42 pairwise_ok = bool(np.max(np.triu(cov,1)) <= 1e-12)
43 # Direct multiaffinity: every enumerated monomial has exponents 0/1.
44 multiaffine = bool(np.all((masks*masks)==masks))
45 # Check the stated log-submodular inequality for Z(A)=sum weights of trees contained in A.
46 z={}
47 for bits in range(1<<6):
48 A={i for i in range(6) if bits>>i & 1}
49 z[bits]=float(sum(w for m,w in zip(masks, np.ones(len(masks))) if set(np.flatnonzero(m)).issubset(A)))
50 min_slack=1e9; violations=0
51 for a in range(1<<6):
52 for b in range(1<<6):
53 union=a|b; inter=a&b
54 slack=z[a]*z[b]-z[union]*z[inter]
55 min_slack=min(min_slack,slack)
56 violations += slack < -1e-10
57 return {
58 'num_trees': int(len(masks)), 'inclusion_probs': inc.tolist(),
59 'max_pair_covariance': float(np.max(np.triu(cov,1))),
60 'min_pair_covariance': float(np.min(np.triu(cov,1)[np.triu_indices(6,1)])),
61 'negative_pairwise_dependence': pairwise_ok,
62 'multiaffine': multiaffine,
63 'log_submodular_min_slack': float(min_slack),
64 'log_submodular_violations': int(violations)
65 }
66
67def make_data(n, seed):
68 rng=np.random.default_rng(seed)
69 # Three latent factors generate correlated candidate routes; target uses all factors.
70 latent=rng.normal(size=(n,3)); x=np.empty((n,6))
71 x[:,0]=latent[:,0]+.18*rng.normal(size=n); x[:,1]=latent[:,0]+.18*rng.normal(size=n)
72 x[:,2]=latent[:,1]+.18*rng.normal(size=n); x[:,3]=latent[:,1]+.18*rng.normal(size=n)
73 x[:,4]=latent[:,2]+.18*rng.normal(size=n); x[:,5]=latent[:,2]+.18*rng.normal(size=n)
74 y=latent @ np.array([1.0,-0.8,0.6]) + .15*rng.normal(size=n)
75 return torch.tensor(x,dtype=torch.float32), torch.tensor(y[:,None],dtype=torch.float32)
76
77class SmallNet(nn.Module):
78 def __init__(self):
79 super().__init__(); self.net=nn.Sequential(nn.Linear(6,12),nn.Tanh(),nn.Linear(12,1))
80 def forward(self,x): return self.net(x)
81
82def train_once(kind, seed, steps=500):
83 set_seed(seed); device='cuda' if torch.cuda.is_available() else 'cpu'
84 try:
85 tr_x,tr_y=make_data(1600,seed+100); va_x,va_y=make_data(800,seed+200)
86 tr_x,tr_y,va_x,va_y=[v.to(device) for v in (tr_x,tr_y,va_x,va_y)]
87 model=SmallNet().to(device); opt=torch.optim.Adam(model.parameters(),lr=0.012)
88 forest_masks, forest_p=enumerate_trees()
89 rng=np.random.default_rng(seed+500)
90 for _ in range(steps):
91 idx=torch.randint(0,len(tr_x),(64,),device=device); xb=tr_x[idx]; yb=tr_y[idx]
92 if kind=='forest':
93 mi=rng.choice(len(forest_masks),size=len(idx),p=forest_p)
94 mask=torch.tensor(forest_masks[mi],dtype=torch.float32,device=device)
95 else:
96 mask=(torch.rand((len(idx),6),device=device)<0.5).float()
97 xb=xb*mask/0.5
98 loss=((model(xb)-yb)**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
99 with torch.no_grad():
100 val=((model(va_x)-va_y)**2).mean().item()
101 return val
102 except Exception:
103 # GPU failures (including shared-memory allocation issues) are retried on CPU.
104 torch.cuda.empty_cache() if torch.cuda.is_available() else None
105 torch.set_default_device('cpu')
106 return train_once_cpu(kind,seed,steps)
107
108def train_once_cpu(kind, seed, steps=500):
109 set_seed(seed); tr_x,tr_y=make_data(1600,seed+100); va_x,va_y=make_data(800,seed+200)
110 model=SmallNet(); opt=torch.optim.Adam(model.parameters(),lr=0.012); masks,p=enumerate_trees(); rng=np.random.default_rng(seed+500)
111 for _ in range(steps):
112 idx=torch.randint(0,len(tr_x),(64,)); xb=tr_x[idx]; yb=tr_y[idx]
113 m=torch.tensor(masks[rng.choice(len(masks),len(idx),p=p)] if kind=='forest' else (rng.random((len(idx),6))<.5),dtype=torch.float32)
114 loss=((model(xb*m/.5)-yb)**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
115 return ((model(va_x)-va_y)**2).mean().item()
116
117def experiment():
118 vals={k:[train_once(k,s) for s in range(3)] for k in ('bernoulli','forest')}
119 return {k:{'runs':v,'mean':float(np.mean(v)),'std':float(np.std(v,ddof=1))} for k,v in vals.items()}
120
121if __name__=='__main__':
122 out={'math':math_check(),'experiment':experiment()}
123 print(json.dumps(out,indent=2))