Online Taylor Residual World Model / online_taylor_experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import json, math
  2import numpy as np
  3
  4SEED = 3100
  5rng = np.random.default_rng(SEED)
  6
  7
  8def monomials(z, degree):
  9    """[1, linear terms, all total-degree 2..degree] in deterministic order."""
 10    z = np.asarray(z, dtype=float)
 11    out = [1.0]
 12    n = len(z)
 13    # recursive exponent enumeration, grouped by total degree
 14    def comps(total, dims, prefix=()):
 15        if dims == 1:
 16            yield prefix + (total,)
 17        else:
 18            for a in range(total + 1):
 19                yield from comps(total-a, dims-1, prefix+(a,))
 20    for d in range(1, degree+1):
 21        for alpha in comps(d, n):
 22            v = 1.0
 23            for zi, ai in zip(z, alpha): v *= zi ** ai
 24            out.append(v)
 25    return np.asarray(out)
 26
 27
 28class OnlineTaylorRLS:
 29    def __init__(self, in_dim, out_dim, degree=1, lam=0.98, p0=100.0):
 30        self.in_dim, self.out_dim, self.degree = in_dim, out_dim, degree
 31        self.lam = lam
 32        self.W = np.zeros((out_dim, len(monomials(np.zeros(in_dim), degree))))
 33        self.P = np.eye(self.W.shape[1]) * p0
 34
 35    def update(self, z, target):
 36        phi = monomials(z, self.degree)
 37        Pphi = self.P @ phi
 38        K = Pphi / (self.lam + phi @ Pphi)
 39        err = np.asarray(target) - self.W @ phi
 40        self.W += np.outer(err, K)
 41        self.P = (self.P - np.outer(K, phi @ self.P)) / self.lam
 42        self.P = (self.P + self.P.T) / 2
 43        return err
 44
 45    def predict_residual(self, z):
 46        return self.W @ monomials(z, self.degree)
 47
 48
 49def rls_exactness_check():
 50    # RLS with lambda=1 must equal ordinary least squares after every prefix.
 51    local = np.random.default_rng(SEED + 1)
 52    Z = local.normal(size=(30, 3)); Phi = np.array([monomials(z, 2) for z in Z])
 53    true_w = local.normal(size=(2, Phi.shape[1]))
 54    Y = Phi @ true_w.T + .01 * local.normal(size=(30, 2))
 55    r = OnlineTaylorRLS(3, 2, degree=2, lam=1.0, p0=1e8)
 56    errors=[]
 57    for i in range(len(Z)):
 58        r.update(Z[i], Y[i])
 59        # large p0 is effectively unregularized; compare after enough samples
 60        batch = np.linalg.lstsq(Phi[:i+1], Y[:i+1], rcond=None)[0].T
 61        errors.append(float(np.max(np.abs(r.W-batch))))
 62    # Verify the covariance-weighted identity more robustly via predictions on heldout points.
 63    pred_err = np.max(np.abs((r.W @ Phi[-5:].T) - (batch @ Phi[-5:].T)))
 64    return {"max_prefix_coef_error": max(errors[10:]), "final_prediction_error": float(pred_err)}
 65
 66
 67def forgetting_check():
 68    # Constant residual step: adaptation should be faster for lower lambda,
 69    # while noisy stationary variance should be larger (the stated tradeoff).
 70    def run(lam, noise, n=500):
 71        r=OnlineTaylorRLS(1,1,degree=0,lam=lam,p0=1.0)
 72        vals=[]; rg=np.random.default_rng(55)
 73        for k in range(n):
 74            y = (1.0 if k >= 30 else 0.0) + noise*rg.normal()
 75            r.update([0.0], [y]); vals.append(float(r.W[0,0]))
 76        return np.asarray(vals)
 77    fast=run(.90,.08); slow=run(.99,.08)
 78    def recovery(a):
 79        return int(np.argmax(a[30:] >= .9)+30)
 80    return {"lambda_.90_recovery_step": recovery(fast), "lambda_.99_recovery_step": recovery(slow),
 81            "lambda_.90_post_std": float(np.std(fast[-150:])), "lambda_.99_post_std": float(np.std(slow[-150:]))}
 82
 83
 84def dynamics_experiment():
 85    """Frozen nominal MLP versus online Taylor residual after a dynamics shift."""
 86    from sklearn.neural_network import MLPRegressor
 87    local = np.random.default_rng(SEED + 2)
 88
 89    def F_nom(x, u):
 90        return .82*x + .18*u + .08*x*x
 91    def F_changed(x, u):
 92        return .82*x + .18*u + .08*x*x + .32*u + .10
 93
 94    # Train only on nominal transitions; this is the frozen global prior.
 95    tx = local.uniform(-1, 1, 4000)
 96    tu = local.uniform(-1, 1, 4000)
 97    prior = MLPRegressor(hidden_layer_sizes=(24, 24), activation='tanh',
 98                         solver='lbfgs', alpha=1e-5, max_iter=300,
 99                         random_state=SEED)
100    prior.fit(np.c_[tx, tu], F_nom(tx, tu))
101    def base(x, u):
102        return float(prior.predict(np.array([[x, u]]))[0])
103
104    # Same shifted trajectory and inputs for every method.
105    n = 260
106    us = local.uniform(-.8, .8, n)
107    xs = np.zeros(n + 1)
108    for k in range(n):
109        xs[k+1] = np.clip(F_changed(xs[k], us[k]), -2, 2)
110
111    results = {}
112    for degree in [0, 1, 2, 3]:
113        r = None if degree == 0 else OnlineTaylorRLS(2, 1, degree=degree,
114                                                     lam=.94, p0=20.)
115        one_step, early, horizon = [], [], []
116        for k in range(n):
117            b = base(xs[k], us[k])
118            if r is None:
119                pred = b
120                predict = lambda xx, uu: base(xx, uu)
121            else:
122                predict = lambda xx, uu: base(xx, uu) + float(
123                    r.predict_residual(np.array([xx, uu]))[0])
124                pred = predict(xs[k], us[k])
125            if k >= 30:
126                one_step.append(abs(pred - xs[k+1]))
127                if k < 80:
128                    early.append(abs(pred - xs[k+1]))
129            if k >= 30 and k + 10 <= n:
130                xx = xs[k]
131                for j in range(10):
132                    xx = predict(xx, us[k+j])
133                    horizon.append(abs(xx - xs[k+j+1]))
134            if r is not None:
135                r.update(np.array([xs[k], us[k]]), [xs[k+1] - b])
136        results[str(degree)] = {
137            "mean_abs_one_step_error": float(np.mean(one_step)),
138            "early_postshift_error": float(np.mean(early)),
139            "mean_abs_10_step_rollout_error": float(np.mean(horizon)),
140        }
141    return results
142
143def main():
144    out={"seed":SEED,"rls_exactness":rls_exactness_check(),
145         "forgetting_tradeoff":forgetting_check(),"dynamics":dynamics_experiment()}
146    print(json.dumps(out, indent=2))
147
148if __name__ == '__main__': main()