Spherical harmonic spectrum regularizer / spectrum_experiment.py

Mechanism failed

Raw ⬇ ZIP
 1import math, random, json
 2import numpy as np
 3import torch
 4
 5SEED = 1726
 6torch.set_default_dtype(torch.float64)
 7random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
 8
 9def normalize(z):
10    return z / (torch.linalg.vector_norm(z, dim=-1, keepdim=True) + 1e-12)
11
12def legendre_all(x, L):
13    vals = [torch.ones_like(x)]
14    if L >= 1: vals.append(x)
15    for l in range(2, L + 1):
16        vals.append(((2*l-1)*x*vals[-1] - (l-1)*vals[-2]) / l)
17    return vals
18
19def spectrum(u, L=4):
20    # Addition theorem: sum_m Y_lm(u_i)Y*_lm(u_j)=(2l+1)P_l(dot)/(4pi).
21    # Hence P_l=(1/(4pi))*sum_ij P_l(u_i dot u_j).
22    dots = torch.clamp(u @ u.T, -1.0, 1.0)
23    return torch.stack([v.sum()/(4*math.pi) for v in legendre_all(dots, L)])
24
25def spec_loss_z(z, L=4):
26    u = normalize(z); p = spectrum(u, L)
27    q = p[1:] / p[0].detach()
28    return (q*q).sum(), p, q
29
30def explicit_l1_power(u):
31    a = math.sqrt(3/(4*math.pi)) * u.sum(dim=0)
32    return (a*a).sum()/3.0
33
34def optimize(lam, steps=700, seed=SEED, N=64, L=4):
35    torch.manual_seed(seed)
36    z = torch.randn(N, 3)*0.35; z[:,2] += 1.0; z.requires_grad_()
37    opt = torch.optim.Adam([z], lr=0.035); trace = {}
38    for step in range(steps+1):
39        u = normalize(z); task = -u[:,2].mean(); sl,p,q = spec_loss_z(z,L)
40        if step in (0,10,50,200,700): trace[str(step)] = [float(task),float(sl),float(q[0])]
41        if step < steps:
42            opt.zero_grad(); (task + lam*sl).backward(); opt.step()
43    with torch.no_grad():
44        u=normalize(z); task=-u[:,2].mean(); sl,p,q=spec_loss_z(z,L)
45        dirs=normalize(torch.randn(100,3)); counts=torch.stack([(u@d >= math.cos(math.pi/4)).double().sum() for d in dirs])
46    return {'lambda':lam,'task':float(task),'spec_loss':float(sl),'q':q.tolist(),'cap_var':float(counts.var(unbiased=False)),'trace':trace}
47
48def numerical_checks():
49    torch.manual_seed(8); z=(torch.randn(11,3)*.4).requires_grad_()
50    u=normalize(z); p=spectrum(u,4); p1=explicit_l1_power(u)
51    sl,_,_=spec_loss_z(z,4); sl.backward(); g=z.grad.detach().clone()
52    d=torch.randn_like(z); eps=1e-5
53    with torch.no_grad():
54        plus=spec_loss_z(z.detach()+eps*d,4)[0]; minus=spec_loss_z(z.detach()-eps*d,4)[0]
55    fd=float((plus-minus)/(2*eps)); ad=float((g*d).sum())
56    return {'addition_theorem_l1_abs_error':float(abs(p1-p[1])),'gradient_fd':fd,'gradient_autodiff':ad,'gradient_relative_error':abs(fd-ad)/(abs(fd)+1e-12)}
57
58def gradient_scaling():
59    torch.manual_seed(91); z=(torch.randn(64,3)*.35); z[:,2]+=1; z.requires_grad_()
60    sl,_,_=spec_loss_z(z,4); sl.backward(); gs=z.grad.norm().item()
61    z2=z.detach().clone().requires_grad_(); (-normalize(z2)[:,2].mean()).backward(); gt=z2.grad.norm().item()
62    return {'spectrum_grad_norm':gs,'task_grad_norm':gt,'lambda_times_grad_ratio':{str(x):x*gs/gt for x in [0,.01,.1,1]}}
63
64def main():
65    sweep=[optimize(x) for x in [0,.001,.003,.01,.03,.1,.3,1]]
66    conv=[optimize(.1,steps=s)['spec_loss'] for s in [0,10,50,200,700]]
67    print(json.dumps({'checks':numerical_checks(),'gradient_scaling':gradient_scaling(),'sweep':sweep,'convergence_lambda_0.1':conv,'notes':'q_l=P_l/P_0, target q_1..q_4=0; task favors north-pole collapse.'},indent=2))
68
69if __name__=='__main__': main()