Spherical harmonic spectrum regularizer / spectrum_experiment.py
Mechanism failed
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()