Fractional Memory State-Space Layer / fractional_memory_experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4from scipy.optimize import nnls
  5from scipy.special import gamma
  6
  7
  8def fit_bank(p, smin=1.0, smax=2048.0, J=24):
  9    lam = np.logspace(-math.log10(smax), -math.log10(smin), J)
 10    grid = np.geomspace(smin, smax, 500)
 11    target = grid ** (p - 1.0)
 12    A = np.exp(-grid[:, None] * lam[None, :])
 13    w, _ = nnls(A, target)
 14    pred = A @ w
 15    scale = np.exp(np.mean(np.log(target + 1e-30) - np.log(pred + 1e-30)))
 16    return lam, w * scale
 17
 18
 19def bank_impulse(lam, w, n=4096, dt=1.0):
 20    return np.exp(-np.outer(np.arange(n) * dt, lam)) @ w
 21
 22
 23def bank_transfer(lam, w, omega):
 24    return np.sum(w[None, :] / (lam[None, :] + 1j * omega[:, None]), axis=1)
 25
 26
 27def mechanism_checks():
 28    rows = []
 29    for p in [0.2, 0.5, 0.8]:
 30        lam, w = fit_bank(p)
 31        g = bank_impulse(lam, w)
 32        lags = np.arange(8, 512)
 33        slope = np.polyfit(np.log(lags), np.log(g[lags]), 1)[0]
 34        rows.append({"check": "power_law_slope", "p": p,
 35                     "predicted": p - 1.0, "observed": float(slope),
 36                     "abs_error": float(abs(slope - (p - 1.0))),
 37                     "positive_weights": bool(np.all(w >= 0))})
 38
 39    # Exact discretization predicts a spectral radius exp(-lambda_min*dt),
 40    # strictly below one for positive rates, for every tested time step.
 41    lam, w = fit_bank(0.5)
 42    for dt in [0.1, 1.0, 4.0]:
 43        rho = float(np.max(np.exp(-lam * dt)))
 44        predicted = float(np.exp(-np.min(lam) * dt))
 45        rows.append({"check": "discrete_stability", "dt": dt,
 46                     "predicted_rho": predicted, "observed_rho": rho,
 47                     "stable": rho < 1.0})
 48
 49    # For Q(w)=integral exp(-lambda s)x(t-s) ds, a fractional kernel has
 50    # Q phase approximately -pi*p/2 and magnitude slope -p. This is the
 51    # Fourier-sign counterpart of the paper's C_p phase +pi*p/2.
 52    for p in [0.2, 0.5, 0.8]:
 53        lam, w = fit_bank(p)
 54        omega = np.logspace(-2.3, -0.7, 300)
 55        H = bank_transfer(lam, w, omega)
 56        sel = (omega > 0.004) & (omega < 0.2)
 57        observed_phase = float(np.median(np.unwrap(np.angle(H))[sel]))
 58        observed_slope = float(np.polyfit(np.log(omega[sel]),
 59                                           np.log(np.abs(H[sel])), 1)[0])
 60        phase_pred = -math.pi * p / 2.0
 61        rows.append({"check": "fractional_transfer", "p": p,
 62                     "predicted_phase_rad": phase_pred,
 63                     "observed_phase_rad": observed_phase,
 64                     "phase_error_rad": float(abs(observed_phase-phase_pred)),
 65                     "predicted_log_slope": -p,
 66                     "observed_log_slope": observed_slope,
 67                     "slope_error": float(abs(observed_slope+p))})
 68    return rows
 69
 70def fractional_filter(x, p=0.5, smax=128, J=16):
 71    lam, w = fit_bank(p, 1.0, float(smax), J)
 72    q = np.zeros(J)
 73    out = np.zeros(len(x))
 74    decay = np.exp(-lam)
 75    gain = (1.0 - decay) / lam
 76    a = np.sum(w / lam)
 77    for t, xt in enumerate(x):
 78        q = decay * q + gain * xt
 79        out[t] = a * xt - np.dot(w, q)
 80    return out
 81
 82
 83def delayed_retrieval(seed=7, ntrain=6000, ntest=3000, delay=48):
 84    rng = np.random.default_rng(seed)
 85    # Isolated random pulses make the target unambiguously depend on a past
 86    # event. Evaluate the memory feature exactly at the requested delay.
 87    gap = delay + 1
 88    def make(n):
 89        x = np.zeros(n); y = np.zeros(n)
 90        for t in range(0, n-gap, gap):
 91            bit = rng.integers(0, 2)
 92            x[t] = bit
 93            y[t+delay] = bit
 94        return x, y
 95    train_x, train_y = make(ntrain)
 96    test_x, test_y = make(ntest)
 97    result = {}
 98    for kind in ["fractional", "exp", "raw"]:
 99        if kind == "fractional":
100            z = fractional_filter(train_x, 0.5, max(128, delay * 3), 16)
101            zt = fractional_filter(test_x, 0.5, max(128, delay * 3), 16)
102        elif kind == "exp":
103            lam = math.log(2) / delay
104            def filt(x):
105                q = 0.; out = np.zeros(len(x))
106                e = math.exp(-lam)
107                for t, xt in enumerate(x):
108                    q = e*q + (1-e)/lam*xt
109                    out[t] = q
110                return out
111            z, zt = filt(train_x), filt(test_x)
112        else:
113            z, zt = train_x, test_x
114        X = np.column_stack([np.ones(len(z)), z])
115        beta = np.linalg.lstsq(X, train_y, rcond=None)[0]
116        pred = np.column_stack([np.ones(len(zt)), zt]) @ beta
117        mask = test_y != 0
118        result[kind] = float(np.mean((pred[mask] - test_y[mask]) ** 2))
119    return result
120
121def main():
122    random.seed(0); np.random.seed(0)
123    checks = mechanism_checks()
124    retrieval = delayed_retrieval()
125    report = {"checks": checks, "retrieval_mse": retrieval}
126    Path("results.json").write_text(json.dumps(report, indent=2))
127    print(json.dumps(report, indent=2))
128
129if __name__ == "__main__":
130    main()