dima / TGPA.py
sparsetrace's picture
Create TGPA.py
c530583 verified
Raw
History Blame Contribute Delete
12.5 kB
# TGPA.py
# ============================================================
# TGPA: Takens Gaussian Process Autoregressor.
#
# Purpose:
# A nonparametric baseline that removes NLSA entirely.
# Uses raw Takens delay vectors z_t in R^{L*D}.
#
# Two modes:
# (A) "direct_ic": ICM-only. Fit GP on prefix pairs and forecast.
# (B) "global+ic_resid": optional global GPLM head trained on F_train,
# plus in-context GP residual correction (LISA-style).
#
# Dependencies:
# numpy, scipy
# (optional) GPLM.py if you want the global head
# ============================================================
from __future__ import annotations
from typing import Optional, Tuple, Literal
import numpy as np
from numpy.lib.stride_tricks import sliding_window_view
from scipy.linalg import cho_factor, cho_solve, solve_triangular
try:
from GPLM import GPLM # type: ignore
except Exception:
try:
from gplm import GPLM # type: ignore
except Exception:
GPLM = None # allow "direct_ic" mode without GPLM
# -----------------------------
# utils
# -----------------------------
def _as_2d(X: np.ndarray) -> np.ndarray:
X = np.asarray(X, dtype=float)
if X.ndim == 1:
X = X[:, None]
return X
def _sliding_windows(F_tD: np.ndarray, L: int) -> np.ndarray:
"""
Return windows W (K,L,D) from F (N,D) with K=N-L+1.
"""
F = _as_2d(F_tD)
W = sliding_window_view(F, window_shape=int(L), axis=0)
a, b = W.shape[1], W.shape[2]
if (a, b) == (L, F.shape[1]):
return np.ascontiguousarray(W)
if (a, b) == (F.shape[1], L):
return np.ascontiguousarray(np.transpose(W, (0, 2, 1)))
raise ValueError(f"Unexpected window shape {W.shape} for L={L}, D={F.shape[1]}.")
def _flatten_windows(W_BLD: np.ndarray) -> np.ndarray:
"""
(B,L,D) -> (B, L*D)
"""
W = np.asarray(W_BLD, dtype=np.float64)
if W.ndim == 2:
W = W[None, :, :]
B, L, D = W.shape
return np.ascontiguousarray(W.reshape(B, L * D))
def _pairwise_sq_dists(X: np.ndarray) -> np.ndarray:
"""
Dense pairwise squared Euclidean distances (n,n).
"""
X = np.asarray(X, dtype=np.float64)
x2 = np.sum(X * X, axis=1, keepdims=True)
d2 = x2 + x2.T - 2.0 * (X @ X.T)
np.maximum(d2, 0.0, out=d2)
return d2
def _estimate_rbf_ell_from_d2(d2_mat: np.ndarray, q: float = 0.5, eps: float = 1e-12) -> float:
"""
Heuristic for RBF k = exp(-||x-y||^2/(2 ell^2)):
ell^2 ~ quantile(d^2)/2
"""
iu = np.triu_indices_from(d2_mat, k=1)
vals = d2_mat[iu]
if vals.size == 0:
return 1.0
qv = float(np.quantile(vals, q))
return float(np.sqrt(max(qv, 0.0) / 2.0 + eps))
# ============================================================
# TDGP
# ============================================================
ICMode = Literal["direct_ic", "global+ic_resid"]
class TGPA:
"""
TGPA: Takens GP Autoregressor.
- direct_ic:
Fit GP on prefix pairs (z_t -> x_{t+1}) and rollout.
This is "ICM alone" (no training beyond kernel hypers).
- global+ic_resid:
Train a global GPLM head on training data once,
then do in-context residual GP correction like LISA, but in raw delay space.
"""
def __init__(
self,
F_train: Optional[np.ndarray] = None,
*,
L: int,
ic_mode: ICMode = "direct_ic",
# ---------- global head (optional) ----------
gplm_kwargs: Optional[dict] = None,
center_outputs: bool = True,
# ---------- IC GP settings ----------
gp_noise2: float = 1e-3,
gp_rbf_ell: Optional[float] = None,
gp_rbf_q: float = 0.5,
# cap context points for dense IC solve
ctx_max_points: int = 1000,
ctx_k0: float = 10.0,
ctx_min_windows: Optional[int] = None,
# variance trust gating
use_var_gate: bool = True,
gate_tau2: float = 1.0,
gate_mode: Literal["rational", "exp"] = "rational",
seed: int = 0,
):
self.L = int(L)
self.ic_mode = str(ic_mode)
self._rng = np.random.default_rng(int(seed))
self.gp_noise2 = float(gp_noise2)
self.gp_rbf_ell = None if gp_rbf_ell is None else float(gp_rbf_ell)
self.gp_rbf_q = float(gp_rbf_q)
self.ctx_max_points = int(ctx_max_points)
self.ctx_k0 = float(ctx_k0)
self.ctx_min_windows = int(ctx_min_windows) if ctx_min_windows is not None else max(2, 8)
self.use_var_gate = bool(use_var_gate)
self.gate_tau2 = float(gate_tau2)
self.gate_mode = str(gate_mode).lower().strip()
# global head (optional)
self.has_global = False
self.center_outputs = bool(center_outputs)
self.mu_X: Optional[np.ndarray] = None
self.gplm: Optional[GPLM] = None
if self.ic_mode == "global+ic_resid":
if F_train is None:
raise ValueError("global+ic_resid requires F_train.")
if GPLM is None:
raise ImportError("GPLM not importable, but global+ic_resid requires GPLM.")
F_train = _as_2d(F_train)
self.D = int(F_train.shape[1])
# build training pairs (z_t -> x_{t+1})
W_all = _sliding_windows(F_train, self.L) # (K,L,D), K=N-L+1
K = W_all.shape[0]
N_pairs = K - 1
X_train = _flatten_windows(W_all[:N_pairs]) # (K-1, L*D)
Y_train = np.asarray(F_train[self.L:self.L + N_pairs, :], dtype=np.float64) # (K-1, D)
if self.center_outputs:
self.mu_X = Y_train.mean(axis=0)
Yc = Y_train - self.mu_X[None, :]
else:
self.mu_X = np.zeros((self.D,), dtype=np.float64)
Yc = Y_train
if gplm_kwargs is None:
gplm_kwargs = {}
gkw = dict(gplm_kwargs)
gkw.setdefault("center_X", False)
gkw.setdefault("sigma2", 1e-5)
gkw.setdefault("jitter", 1e-8)
gkw.setdefault("m", min(2048, X_train.shape[0]))
gkw.setdefault("inducing", "fps")
gkw.setdefault("seed", int(seed))
self.gplm = GPLM(X_train, Yc, **gkw)
self.has_global = True
else:
# direct_ic mode: we only need D at call-time
self.D = -1
# -----------------------------
# kernel bits
# -----------------------------
def _rbf_kernel_matrix(self, X: np.ndarray) -> Tuple[np.ndarray, float]:
d2 = _pairwise_sq_dists(X)
ell = _estimate_rbf_ell_from_d2(d2, q=self.gp_rbf_q) if self.gp_rbf_ell is None else float(self.gp_rbf_ell)
K = np.exp(-0.5 * d2 / (ell**2 + 1e-12))
return K, ell
def _rbf_kernel_eval(self, X: np.ndarray, xq: np.ndarray, ell: float) -> np.ndarray:
diff = X - xq[None, :]
d2 = np.einsum("nd,nd->n", diff, diff, optimize=True)
return np.exp(-0.5 * d2 / (ell**2 + 1e-12))
def _gate_from_var(self, var_f: float) -> float:
if not self.use_var_gate:
return 1.0
v = max(float(var_f), 0.0)
tau2 = max(float(self.gate_tau2), 1e-18)
if self.gate_mode == "exp":
return float(np.exp(-v / tau2))
return float(tau2 / (tau2 + v))
# -----------------------------
# global prediction
# -----------------------------
def _global_pred(self, Z_Bp: np.ndarray) -> np.ndarray:
assert self.gplm is not None
Yc = self.gplm(Z_Bp) # centered
if self.center_outputs:
return Yc + self.mu_X[None, :]
return Yc
# ============================================================
# main forecast
# ============================================================
def __call__(self, prefix: np.ndarray, *, steps: int = 1, return_var: bool = False):
"""
prefix: (ell,D), ell>=L
returns preds: (steps,D)
"""
prefix = _as_2d(prefix)
ell, D = prefix.shape
if ell < self.L:
raise ValueError(f"Need prefix length ell >= L={self.L}.")
if self.D < 0:
self.D = int(D)
if D != self.D:
raise ValueError(f"TDGP expects D={self.D}, got D={D}.")
H = int(steps)
if H <= 0:
out = np.zeros((0, D), dtype=np.float64)
return (out, np.zeros((0,), dtype=np.float64)) if return_var else out
# seed rolling window
cur = prefix[-self.L:, :].copy()
# build context pairs from prefix
K_ctx = ell - self.L
if K_ctx < self.ctx_min_windows:
# not enough context: fall back
if self.ic_mode == "global+ic_resid":
return self._rollout_global(cur, H, return_var=return_var)
raise ValueError("direct_ic needs ell > L with enough context pairs.")
W_all = _sliding_windows(prefix, self.L) # (ell-L+1, L, D)
W_ctx = np.ascontiguousarray(W_all[:K_ctx, :, :]) # windows with observed next step
X_ctx = _flatten_windows(W_ctx) # (K_ctx, L*D)
Y_ctx = np.asarray(prefix[self.L:self.L + K_ctx, :], dtype=np.float64) # (K_ctx, D)
# possibly subsample context to keep dense GP affordable
M = min(int(K_ctx), int(self.ctx_max_points))
if M < K_ctx:
# simple but effective: pick evenly spaced indices
idx = np.linspace(0, K_ctx - 1, M).round().astype(np.int64)
X_fit = X_ctx[idx]
Y_fit = Y_ctx[idx]
else:
X_fit = X_ctx
Y_fit = Y_ctx
# choose targets for IC solve
if self.ic_mode == "global+ic_resid":
# residual targets
Y_glob_fit = self._global_pred(X_fit) # (M,D)
T_fit = Y_fit - Y_glob_fit
else:
# direct map
T_fit = Y_fit
# fit dense GP on (X_fit -> T_fit)
K_mat, ell_used = self._rbf_kernel_matrix(X_fit)
K_reg = K_mat + self.gp_noise2 * np.eye(M, dtype=np.float64)
cf = cho_factor(K_reg, lower=True, check_finite=False)
alpha = cho_solve(cf, T_fit, check_finite=False) # (M,D)
Lfac, lower = cf
# base context mixing weight
w_ctx_base = float(M) / float(M + self.ctx_k0) if self.ctx_k0 > 0 else 1.0
preds = np.zeros((H, D), dtype=np.float64)
vars_out = np.zeros((H,), dtype=np.float64) if return_var else None
# rollout
for h in range(H):
zq = _flatten_windows(cur)[0] # (L*D,)
if self.ic_mode == "global+ic_resid":
y_base = self._global_pred(zq[None, :])[0]
else:
y_base = np.zeros((D,), dtype=np.float64)
k_eval = self._rbf_kernel_eval(X_fit, zq, ell_used) # (M,)
t_mean = k_eval @ alpha # (D,)
# function variance proxy
u = solve_triangular(Lfac, k_eval, lower=lower, check_finite=False)
quad = float(np.dot(u, u))
var_f = max(0.0, 1.0 - quad)
if return_var:
vars_out[h] = var_f
w_gate = self._gate_from_var(var_f)
w_eff = w_ctx_base * w_gate
if self.ic_mode == "global+ic_resid":
y = y_base + w_eff * t_mean
else:
# direct_ic: optionally apply mixing too (stabilizes rollout)
y = w_eff * t_mean + (1.0 - w_eff) * y_base # y_base is zeros here
preds[h] = y
# update rolling window
if self.L > 1:
cur[:-1] = cur[1:]
cur[-1] = y
if return_var:
return preds, vars_out
return preds
def _rollout_global(self, seed_LD: np.ndarray, H: int, return_var: bool = False):
"""
Pure global rollout (only for global+ic_resid mode).
"""
assert self.gplm is not None
cur = np.asarray(seed_LD, dtype=np.float64).copy()
out = np.zeros((H, self.D), dtype=np.float64)
for h in range(int(H)):
zq = _flatten_windows(cur)[0]
y = self._global_pred(zq[None, :])[0]
out[h] = y
if self.L > 1:
cur[:-1] = cur[1:]
cur[-1] = y
if return_var:
return out, np.zeros((H,), dtype=np.float64)
return out
__all__ = ["TGPA"]