Strongly-Rayleigh Forest Dropout / forest_dropout_experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  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))