Critical-Tail Multiscale Mixer / run_experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
 1import math, json, random, time
 2import numpy as np
 3import torch
 4import torch.nn as nn
 5import torch.nn.functional as F
 6
 7SEED=417
 8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
 9try:
10    device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
11    if device.type=='cuda': torch.cuda.empty_cache()
12except Exception:
13    device=torch.device('cpu')
14
15def kval(r, r0=2.0):
16    return 1.0/((r+r0)*math.log(r+r0)**2)
17
18def math_check():
19    rs=np.arange(1,300001,dtype=np.float64)
20    kv=1/((rs+2)*np.log(rs+2)**2)
21    tail=[]
22    for R in [8,16,32,64,128,256]:
23        s=float(kv[R:].sum()); q=1/math.log(R+2)
24        tail.append([R,s,q,s/q])
25    bands=[]
26    for k in range(7):
27        lo,hi=2**k,2**(k+1)
28        exact=float(kv[lo-1:hi-1].sum())
29        formula=1/((k+1)*(k+2)*math.log(2))
30        bands.append([k,exact,formula,exact/formula])
31    return {'tail_R_discrete_integral_ratio':tail,'bands_k_exact_formula_ratio':bands}
32
33def make_weights(L):
34    w=torch.zeros(L,device=device)
35    for r in range(1,L): w[r]=kval(r)
36    return w/(2*w.sum())
37
38def long_target(x,w):
39    # x [B,L,1], symmetric zero-padded convolution
40    B,L,D=x.shape; y=torch.zeros_like(x)
41    for r in range(1,L):
42        y[:,r:]+=w[r]*x[:,:L-r]; y[:,:L-r]+=w[r]*x[:,r:]
43    return y
44
45def dyadic_mix(x):
46    B,L,D=x.shape; out=torch.zeros_like(x); pref=torch.cat([torch.zeros(B,1,D,device=x.device),x.cumsum(1)],1)
47    K=int(math.log2(L))
48    for k in range(K):
49        lo,hi=2**k, min(2**(k+1),L)
50        # clipped ranges, with exact valid-count normalization per position
51        idx=torch.arange(L,device=x.device)
52        a=(idx-hi).clamp(min=0); b=(idx-lo).clamp(min=0)
53        left=(pref[:,b+1]-pref[:,a+1]) # positions [i-hi+1,i-lo], then remove invalid below zero
54        left_count=(b-a).clamp(min=1).float()[None,:,None]
55        # Correct range is j in [i-hi, i-lo], exclusive invalid j<0.
56        left=(pref[:,(idx-lo+1).clamp(min=0)]-pref[:,(idx-hi).clamp(min=0)])/left_count
57        a=(idx+lo).clamp(max=L); b=(idx+hi).clamp(max=L)
58        right=(pref[:,b]-pref[:,a])/(b-a).clamp(min=1).float()[None,:,None]
59        ak=1/((k+1)*(k+2)*math.log(2))
60        out += ak*(left+right)
61    return out
62
63def local_mix(x,r=8):
64    B,L,D=x.shape; out=torch.zeros_like(x)
65    for d in range(1,r+1):
66        out[:,d:]+=x[:,:L-d]; out[:,:L-d]+=x[:,d:]
67    return out/(2*r)
68
69class Dyadic(nn.Module):
70    def __init__(self): super().__init__(); self.scale=nn.Parameter(torch.tensor(.5))
71    def forward(self,x): return x+self.scale*dyadic_mix(x)
72class Local(nn.Module):
73    def __init__(self,r=8): super().__init__(); self.scale=nn.Parameter(torch.tensor(.5)); self.r=r
74    def forward(self,x): return x+self.scale*local_mix(x,self.r)
75
76def train(model, L, steps=180, batch=64):
77    model.to(device); opt=torch.optim.Adam(model.parameters(),lr=.03); w=make_weights(L)
78    losses=[]; t=time.time()
79    for s in range(steps):
80        x=torch.randn(batch,L,1,device=device); y=long_target(x,w)
81        pred=model(x); loss=F.mse_loss(pred,y); opt.zero_grad(); loss.backward(); opt.step()
82        losses.append(float(loss))
83    with torch.no_grad():
84        x=torch.randn(512,L,1,device=device); y=long_target(x,w); val=float(F.mse_loss(model(x),y))
85    return {'final_train':losses[-1],'val_mse':val,'seconds':time.time()-t,'learned_scale':float(model.scale.detach().cpu())}
86
87def main():
88    out={'device':str(device),'math':math_check(),'experiments':{}}
89    for L in [64,256]:
90        # same architecture/optimization; local control has same single scalar parameter
91        out['experiments'][str(L)]={'local_radius8':train(Local(8),L),'dyadic_tail':train(Dyadic(),L)}
92    print(json.dumps(out,indent=2))
93if __name__=='__main__': main()