Bellman-Resolvent Uncertainty Targets / experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math, os, random
  2import numpy as np
  3
  4SEED = 123
  5np.random.seed(SEED)
  6random.seed(SEED)
  7
  8
  9def resolvent_iter(H, L, gamma, K):
 10    x = np.zeros_like(H, dtype=float)
 11    for _ in range(K):
 12        x = H + gamma * L.dot(x)
 13    return x
 14
 15
 16def toy_verification():
 17    # Scalar eigenmode: xi = H/(1-gamma*lambda), and finite-K xi_K.
 18    gamma = 0.9
 19    H = np.array([1.0])
 20    rows = []
 21    # Prediction 1: amplification is exactly 1/(1-gamma lambda).
 22    for lam in [0.0, 0.2, 0.5, 0.8, 0.95, 1.0]:
 23        L = np.array([[lam]])
 24        x = resolvent_iter(H, L, gamma, 1000)
 25        pred = 1.0 / (1.0 - gamma * lam)
 26        rows.append({"lambda": lam, "observed_gain": float(x[0]), "predicted_gain": pred,
 27                     "relative_error": abs(float(x[0])-pred)/pred})
 28    # Prediction 2: truncation residual after K is (gamma lambda)^K.
 29    lam = 0.8
 30    q = gamma * lam
 31    trunc = []
 32    for K in [1, 2, 4, 8, 12, 20]:
 33        observed = abs(float(resolvent_iter(H, np.array([[lam]]), gamma, K)[0] - 1/(1-q)))
 34        predicted = q**K / (1-q)
 35        trunc.append({"K": K, "observed_residual": observed, "predicted_residual": predicted,
 36                      "relative_error": abs(observed-predicted)/max(predicted, 1e-15)})
 37    # Prediction 3: instability at gamma*lambda > 1, with finite-K growth.
 38    boundary = []
 39    for lam in [0.8, 1.0, 1.05, 1.2]:
 40        q = gamma * lam
 41        vals = [abs(float(resolvent_iter(H, np.array([[lam]]), gamma, K)[0])) for K in [5, 10, 20]]
 42        boundary.append({"lambda": lam, "gamma_lambda": q, "abs_x_K5_K10_K20": vals,
 43                         "predicted": "bounded" if q < 1 else "diverges"})
 44    # Covariance prediction using scalar Gaussian/bootstrap samples.
 45    rng = np.random.default_rng(SEED)
 46    lam = 0.7; q = gamma*lam; sigma_h = 0.35
 47    hs = rng.normal(0, sigma_h, size=200000)
 48    xis = hs/(1-q)
 49    observed_var = float(np.var(xis)); predicted_var = sigma_h**2/(1-q)**2
 50    covariance = {"lambda": lam, "observed_variance": observed_var,
 51                  "predicted_variance": predicted_var,
 52                  "relative_error": abs(observed_var-predicted_var)/predicted_var}
 53    return {"gain_sweep": rows, "truncation_sweep": trunc, "boundary_sweep": boundary,
 54            "variance_check": covariance}
 55
 56
 57def offline_critic_experiment():
 58    # Small deterministic chain with noisy one-step rewards. The target variance is
 59    # estimated from bootstrap reward replicas and propagated through P.
 60    rng = np.random.default_rng(SEED + 1)
 61    n = 24; gamma = 0.9; episodes = 180; horizon = 12
 62    # Fixed policy transitions: mostly move right, with a terminal reward at last state.
 63    P = np.zeros((n,n))
 64    for s in range(n):
 65        if s == n-1: P[s,s] = 1
 66        else:
 67            P[s,s+1] = 0.82; P[s,s] += 0.18
 68    r_true = np.zeros(n); r_true[-1] = 1.0
 69    V_true = np.linalg.solve(np.eye(n)-gamma*P, r_true)
 70    # Dataset has state-dependent reward noise and sparse visits to late states.
 71    data = []
 72    for _ in range(episodes):
 73        s = 0
 74        for t in range(horizon):
 75            ns = int(rng.choice(n, p=P[s]))
 76            noise = rng.normal(0, 0.05 + 0.35*(s > 7))
 77            r = r_true[s] + noise
 78            data.append((s, ns, r)); s = ns
 79    # Replay counts and bootstrap target variance (known next-state here).
 80    counts = np.bincount([x[0] for x in data], minlength=n)
 81    B = 32
 82    samples = [[] for _ in range(n)]
 83    for s, ns, r in data:
 84        samples[s].append((ns,r))
 85    Hvar = np.zeros(n)
 86    Hmean = np.zeros(n)
 87    for s in range(n):
 88        if samples[s]:
 89            vals = np.array([r for _,r in samples[s]])
 90            # bootstrap means approximate one-step empirical-process uncertainty
 91            boot = np.array([rng.choice(vals, size=len(vals), replace=True).mean() for _ in range(B)])
 92            Hvar[s] = np.var(boot, ddof=1)
 93            Hmean[s] = vals.mean() - r_true[s]
 94        else:
 95            Hvar[s] = 0.5
 96    # Resolvent of standard deviation (diagonal approximation requested by idea).
 97    # Use transition propagation for the point uncertainty magnitude.
 98    Hstd = np.sqrt(Hvar + 1e-9)
 99    u = np.zeros(n)
100    for _ in range(40): u = Hstd + gamma * P.dot(u)
101    # Fitted value iteration with a shared low-dimensional critic. This avoids the
102    # degenerate tabular case where each state has its own intercept and weights
103    # cancel from the per-state mean.
104    z = np.arange(n, dtype=float) / (n - 1)
105    Phi = np.column_stack([np.ones(n), z, z*z, (z > 0.55).astype(float)])
106    def fit(weighted):
107        theta = np.zeros(Phi.shape[1])
108        for _ in range(80):
109            numer = np.zeros(Phi.shape[1]); denom = np.zeros((Phi.shape[1], Phi.shape[1]))
110            for s,ns,r in data:
111                y = r + gamma * float(Phi[ns].dot(theta))
112                w = 1.0/(u[s] + 0.05) if weighted else 1.0
113                numer += w * Phi[s] * y
114                denom += w * np.outer(Phi[s], Phi[s])
115            theta_new = np.linalg.solve(denom + 1e-5*np.eye(Phi.shape[1]), numer)
116            theta = 0.7*theta + 0.3*theta_new
117        return Phi.dot(theta)
118    vb = fit(False); vi = fit(True)
119    # Test error is against known population value, and Bellman residual uses true P.
120    def metrics(v):
121        bell = r_true + gamma*P.dot(v) - v
122        return {"value_rmse": float(np.sqrt(np.mean((v-V_true)**2))),
123                "bellman_rmse": float(np.sqrt(np.mean(bell**2))),
124                "late_state_rmse": float(np.sqrt(np.mean((v[8:]-V_true[8:])**2)))}
125    # Calibration: correlation of propagated uncertainty with absolute value error.
126    cal = float(np.corrcoef(u, np.abs(vb-V_true))[0,1])
127    return {"baseline": metrics(vb), "resolvent_weighted": metrics(vi),
128            "uncertainty_error_correlation": cal,
129            "mean_uncertainty_early": float(np.mean(u[:8])),
130            "mean_uncertainty_late": float(np.mean(u[8:])),
131            "dataset_count_early": int(np.sum(counts[:8])), "dataset_count_late": int(np.sum(counts[8:]))}
132
133
134def main():
135    out = {"seed": SEED, "toy": toy_verification(), "offline": offline_critic_experiment()}
136    with open("results.json", "w") as f: json.dump(out, f, indent=2)
137    print(json.dumps(out, indent=2))
138
139if __name__ == '__main__': main()