Fisher-Observable Latent State Training / run_experiment.py

Failed on benchmark

Raw ⬇ ZIP
 1import json, random
 2from pathlib import Path
 3import numpy as np
 4import torch
 5from fisher_observable import fisher_information
 6
 7SEED = 2387
 8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
 9torch.set_default_dtype(torch.float64)
10
11def rollout(z0, k, dt=0.5):
12    z = z0
13    for _ in range(k):
14        z = torch.stack((z[0] + dt*z[2], z[1] + dt*z[3], z[2], z[3]))
15    return z
16
17def obs(z, use_brightness=True, c=1.0):
18    x, y = z[0], z[1]
19    bearing = torch.atan2(y, x)
20    if use_brightness:
21        return torch.stack((bearing, c/(x*x+y*y)))
22    return bearing.reshape(1)
23
24def jacobians(z0, K, use_brightness, dt=0.5):
25    return [torch.autograd.functional.jacobian(lambda q: obs(rollout(q,k,dt),use_brightness), z0) for k in range(K+1)]
26
27def fisher(z0, K, use_brightness, sig_b=0.01, sig_l=0.01):
28    js = jacobians(z0,K,use_brightness)
29    vars = [sig_b**2] if not use_brightness else [sig_b**2,sig_l**2]
30    return fisher_information(js, vars).detach().numpy()
31
32def recover(z_init, observations, K, use_brightness, sig_b, sig_l, steps=500):
33    q=torch.tensor(z_init, requires_grad=True); opt=torch.optim.Adam([q],lr=0.04)
34    target=torch.tensor(observations)
35    for _ in range(steps):
36        opt.zero_grad(); pred=torch.cat([obs(rollout(q,k),use_brightness) for k in range(K+1)])
37        if use_brightness:
38            rb=torch.atan2(torch.sin(pred[0::2]-target[0::2]), torch.cos(pred[0::2]-target[0::2]))
39            rl=pred[1::2]-target[1::2]; loss=(rb/sig_b).square().mean()+(rl/sig_l).square().mean()
40        else:
41            rr=torch.atan2(torch.sin(pred-target),torch.cos(pred-target)); loss=(rr/sig_b).square().mean()
42        loss.backward(); opt.step()
43    return q.detach().numpy()
44
45def main():
46    z=torch.tensor([3.0,1.0,0.4,-0.2]); out={"seed":SEED}
47    rank_rows=[]
48    for K in [0,1,2,4,8]:
49        row={"K":K}
50        for name,use in [("bearing",False),("bearing_brightness",True)]:
51            e=np.linalg.eigvalsh(fisher(z,K,use)); row[name+"_rank"]=int(np.sum(e>1e-8)); row[name+"_eigenvalues"]=e.tolist()
52        rank_rows.append(row)
53    out["rank_sweep"]=rank_rows
54    # At K=4 both position and velocity directions can be observed.
55    scaling=[]
56    for s in [0.005,0.01,0.02,0.04,0.08]:
57        e=np.linalg.eigvalsh(fisher(z,4,True,sig_b=0.01,sig_l=s))
58        scaling.append({"sigma_l":s,"sigma_inv_sq":1/s**2,"lambda_min":float(e[0]),"lambda_max":float(e[-1])})
59    out["noise_scaling"]=scaling
60    # Explicit transition sweep around the predicted equal-whitened-sensitivity noise.
61    transition=[]
62    for s in [0.05, 0.10, 0.15, 0.158113883, 0.20, 0.30, 0.60]:
63        e=np.linalg.eigvalsh(fisher(z,4,True,sig_b=0.01,sig_l=s))
64        transition.append({"sigma_l":s, "lambda_min":float(e[0]), "condition":float(e[-1]/max(e[0],1e-15))})
65    out["noise_transition"]=transition
66    # Independent central-difference check of the autodiff chain-rule Jacobian.
67    kcheck=2; h=1e-6; q=z.detach().numpy(); J=torch.cat(jacobians(z,kcheck,True),dim=0).numpy()
68    def flat_obs(qn):
69        qq=torch.tensor(qn); return np.concatenate([obs(rollout(qq,k),True).detach().numpy() for k in range(kcheck+1)])
70    Jfd=np.column_stack([(flat_obs(q+np.eye(4)[j]*h)-flat_obs(q-np.eye(4)[j]*h))/(2*h) for j in range(4)])
71    out["jacobian_check_max_abs_error"]=float(np.max(np.abs(J-Jfd)))
72    out["first_full_rank_K"]={name:next((r["K"] for r in rank_rows if r[name+"_rank"]==4),None) for name in ["bearing","bearing_brightness"]}
73    jb=jacobians(z,0,True)[0].numpy(); bearing_norm=np.linalg.norm(jb[0]/0.01)
74    threshold=1.0/np.linalg.norm(jb[1]/0.01)
75    out["predicted_equal_whitened_sensitivity_sigma_brightness"]=float(threshold)
76    K=4; sb=0.01; sl=0.01
77    clean=np.concatenate([obs(rollout(z,k),True).detach().numpy() for k in range(K+1)])
78    rng=np.random.default_rng(SEED); noisy=clean.copy(); noisy[0::2]+=rng.normal(0,sb,K+1); noisy[1::2]+=rng.normal(0,sl,K+1)
79    init=np.array([2.7,1.3,0.25,-0.05]); base=recover(init,noisy[0::2],K,False,sb,sl); idea=recover(init,noisy,K,True,sb,sl)
80    true=z.numpy()
81    out["recovery"]={"true":true.tolist(),"bearing_only_estimate":base.tolist(),"bearing_plus_brightness_estimate":idea.tolist(),"bearing_only_rmse":float(np.sqrt(np.mean((base-true)**2))),"idea_rmse":float(np.sqrt(np.mean((idea-true)**2)))}
82    out["checks"]={"bearing_rank_at_K8":rank_rows[-1]["bearing_rank"],"combined_rank_at_K1":rank_rows[1]["bearing_brightness_rank"],"lambda_min_ratio_sigma_.01_to_.02":scaling[1]["lambda_min"]/scaling[2]["lambda_min"],"expected_inverse_variance_ratio":4.0,
83      "bearing_condition_at_K4":float(rank_rows[3]["bearing_eigenvalues"][-1]/max(rank_rows[3]["bearing_eigenvalues"][1],1e-15)),
84      "combined_condition_at_K4":float(rank_rows[3]["bearing_brightness_eigenvalues"][-1]/rank_rows[3]["bearing_brightness_eigenvalues"][0]),"bearing_whitened_norm":bearing_norm,"brightness_equal_norm_sigma":float(threshold)}
85    Path("results.json").write_text(json.dumps(out,indent=2)); print(json.dumps(out,indent=2))
86if __name__ == '__main__': main()