Fractional Memory State-Space Layer / fractional_memory_experiment.py
Failed on benchmark
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()