Variance-aware gradient reduction trees / run_experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math
  2from dataclasses import dataclass
  3import numpy as np
  4
  5@dataclass
  6class Node:
  7    left: object = None
  8    right: object = None
  9    leaf: int = None
 10    def is_leaf(self): return self.leaf is not None
 11
 12def leaf(i): return Node(leaf=i)
 13
 14def depths(t, d=0, out=None):
 15    if out is None: out={}
 16    if t.is_leaf(): out[t.leaf]=d
 17    else:
 18        depths(t.left,d+1,out); depths(t.right,d+1,out)
 19    return out
 20
 21def internal_sizes(t):
 22    if t.is_leaf(): return 1, []
 23    na,sa=internal_sizes(t.left); nb,sb=internal_sizes(t.right)
 24    return na+nb, sa+sb+[na+nb]
 25
 26def lambdas(t,k):
 27    ds=depths(t)
 28    _, internal=internal_sizes(t)
 29    return float(sum(ds.values())), float(sum(x*x for x in internal))
 30
 31def Kmat(t,k):
 32    K=np.zeros((k,k))
 33    def walk(n, leaves):
 34        if n.is_leaf(): return [n.leaf]
 35        a=walk(n.left,leaves); b=walk(n.right,leaves); s=a+b
 36        for i in s:
 37            for j in s: K[i,j]+=1
 38        return s
 39    walk(t,[]); return K
 40
 41def balanced(ids):
 42    ids=list(ids)
 43    if len(ids)==1: return leaf(ids[0])
 44    m=len(ids)//2
 45    return Node(balanced(ids[:m]),balanced(ids[m:]))
 46
 47def chain(ids):
 48    ids=list(ids); t=leaf(ids[0])
 49    for i in ids[1:]: t=Node(t,leaf(i))
 50    return t
 51
 52def huffman(weights):
 53    # Min weighted external path length; children are immaterial to score.
 54    heap=[(float(w), i, leaf(i)) for i,w in enumerate(weights)]
 55    serial=len(heap)
 56    import heapq
 57    heapq.heapify(heap)
 58    while len(heap)>1:
 59        wa,_,a=heapq.heappop(heap); wb,_,b=heapq.heappop(heap)
 60        heapq.heappush(heap,(wa+wb,serial,Node(a,b))); serial+=1
 61    return heap[0][2]
 62
 63def proxy(t, variances, mean2=0.0, lam=0.0):
 64    d,l2=lambdas(t,len(variances))
 65    return float(np.dot(variances,[depths(t)[i] for i in range(len(variances))]) + lam*mean2*l2)
 66
 67def exact_identity_check(rng,k=8,n=120000):
 68    t=balanced(range(k)); K=Kmat(t,k); d,l2=lambdas(t,k)
 69    rows=[]
 70    for mu in [0.0,0.25,1.0,3.0]:
 71        for tau in [0.1,0.5,1.5]:
 72            p=rng.normal(mu,tau,size=(n,k))
 73            empirical=np.mean(np.einsum('bi,ij,bj->b',p,K,p))
 74            predicted=tau*tau*d+mu*mu*l2
 75            rows.append({'mu':mu,'tau':tau,'predicted':predicted,'observed':float(empirical),
 76                         'relative_error':float(abs(empirical-predicted)/predicted)})
 77    return {'tree':'balanced_8','Lambda1':d,'Lambda2':l2,'rows':rows,
 78            'max_relative_error':max(x['relative_error'] for x in rows)}
 79
 80def roundoff_scaling(rng,k=8,n=80000):
 81    # Conditional independent additive roundoff at each internal node:
 82    # Var(error | x)=u^2 x^2. This is precisely the paper's local model with nu=1.
 83    p=rng.normal(0.2,0.7,size=(n,k)); t=balanced(range(k)); K=Kmat(t,k)
 84    q=np.einsum('bi,ij,bj->b',p,K,p)
 85    rows=[]
 86    for u in [2**-5,2**-7,2**-9,2**-11]:
 87        # independently inject Gaussian error at each node, aggregate final error
 88        errs=np.zeros(n)
 89        def rec(node):
 90            if node.is_leaf(): return p[:,node.leaf]
 91            x=rec(node.left)+rec(node.right)
 92            e=rng.normal(size=n)*u*np.abs(x)
 93            nonlocal errs
 94            errs += e
 95            return x+e
 96        rec(t)
 97        observed=float(np.mean(errs**2)); predicted=float(u*u*np.mean(q))
 98        rows.append({'u':u,'observed_mse':observed,'predicted_mse':predicted,
 99                     'ratio_obs_over_u2':observed/(u*u),
100                     'relative_error':abs(observed-predicted)/predicted})
101    # slope in log-log should be 2
102    slope=float(np.polyfit(np.log([x['u'] for x in rows]),np.log([x['observed_mse'] for x in rows]),1)[0])
103    return {'rows':rows,'loglog_slope':slope,'predicted_slope':2.0}
104
105def heterogeneous_monte_carlo(rng, n=120000):
106    k=8; variances=np.array([16.,1.,1.,1.,1.,1.,1.,1.])
107    trees=[('balanced',balanced(range(k))),('huffman',huffman(variances))]
108    rows=[]
109    for name,t in trees:
110        K=Kmat(t,k)
111        p=rng.normal(size=(n,k))*np.sqrt(variances)[None,:]
112        observed=float(np.mean(np.einsum('bi,ij,bj->b',p,K,p)))
113        predicted=float(np.dot(variances,[depths(t)[i] for i in range(k)]))
114        rows.append({'tree':name,'predicted_cost':predicted,'observed_cost':observed,
115                     'relative_error':abs(observed-predicted)/predicted,
116                     'depths':depths(t)})
117    return rows
118
119def topology_sweep():
120    k=8; bal=balanced(range(k)); rows=[]
121    for ratio in [1,2,4,8,16,32,64]:
122        w=np.ones(k); w[0]=ratio
123        # Put the high-variance chunk at the shallowest available Huffman leaf.
124        ht=huffman(w); hd=depths(ht); order=sorted(hd,key=hd.get)
125        # huffman construction may assign leaf 0 shallow naturally; relabel if needed
126        # evaluate optimal assignment for the fixed shape by placing largest weights shallow.
127        ds=sorted(depths(ht).values()); weighted_h=float(np.dot(sorted(w,reverse=True),ds))
128        weighted_bal=float(np.dot(sorted(w,reverse=True),sorted(depths(bal).values())))
129        rows.append({'ratio':ratio,'balanced_cost':weighted_bal,'huffman_cost':weighted_h,
130                     'predicted_gain_fraction':1-weighted_h/weighted_bal,
131                     'huffman_depth_high_variance':min(ds)})
132    return rows
133
134def main():
135    rng=np.random.default_rng(1119)
136    identity=exact_identity_check(rng)
137    rounding=roundoff_scaling(rng)
138    topology=topology_sweep()
139    heterogeneous=heterogeneous_monte_carlo(rng)
140    out={'identity_check':identity,'roundoff_scaling':rounding,'variance_topology_sweep':topology,
141         'heterogeneous_monte_carlo':heterogeneous,
142         'notes':['All costs omit common nu*u^2 factors. Huffman minimizes sum sigma_i^2 depth_i among binary trees.',
143                  'The roundoff simulation uses the stated conditional Gaussian local-noise approximation.']}
144    with open('results.json','w') as f: json.dump(out,f,indent=2)
145    print(json.dumps(out,indent=2))
146if __name__=='__main__': main()