| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| from __future__ import annotations |
|
|
| from dataclasses import dataclass |
| from typing import Any, Literal, Optional, Tuple, Union |
|
|
| import numpy as np |
| import scipy.linalg as la |
| import scipy.sparse.linalg as sla |
|
|
|
|
| Solver = Literal["chol", "eigh", "cg"] |
|
|
|
|
| def _as_2d(X: np.ndarray) -> np.ndarray: |
| X = np.asarray(X) |
| if X.ndim == 1: |
| return X[:, None] |
| return X |
|
|
|
|
| @dataclass |
| class _LatentPreproc: |
| mean: np.ndarray |
| std: np.ndarray |
|
|
| def apply(self, R: np.ndarray) -> np.ndarray: |
| return (R - self.mean[None, :]) / self.std[None, :] |
|
|
|
|
| class LGP: |
| """ |
| LGP = Linear-kernel GP mean (KRR) decoder. |
| |
| API is intentionally similar to GPLM: |
| - __init__(...) fits immediately |
| - __call__(R_ax, batch_size=None) predicts mean |
| - predict(R_ax, return_var=True) returns mean + scalar var proxy |
| |
| Notes: |
| - If center_X=True, we center outputs during training and add mean back at inference. |
| - If whiten_latent=True, we standardize latents dimension-wise (recommended if feature scales vary). |
| - sigma2 is the ridge/noise term (λ). Larger sigma2 -> smoother / more stable rollout. |
| """ |
|
|
| def __init__( |
| self, |
| R_ix: np.ndarray, |
| R_iX: np.ndarray, |
| *, |
| sigma2: float = 1e-5, |
| jitter: float = 1e-10, |
| center_X: bool = True, |
| whiten_latent: bool = False, |
| dtype: Any = np.float32, |
| solver: Solver = "chol", |
| cg_maxiter: int = 500, |
| cg_rtol: float = 1e-6, |
| cg_atol: float = 0.0, |
| |
| eig_clip: float = 1e-12, |
| ): |
| self.sigma2 = float(sigma2) |
| self.jitter = float(jitter) |
| self.center_X = bool(center_X) |
| self.whiten_latent = bool(whiten_latent) |
| self.dtype = dtype |
| self.solver: Solver = str(solver) |
|
|
| self.cg_maxiter = int(cg_maxiter) |
| self.cg_rtol = float(cg_rtol) |
| self.cg_atol = float(cg_atol) |
| self.eig_clip = float(eig_clip) |
|
|
| |
| R_ix = np.ascontiguousarray(np.asarray(R_ix).astype(self.dtype, copy=False)) |
| R_iX = np.ascontiguousarray(np.asarray(R_iX).astype(self.dtype, copy=False)) |
| if R_ix.ndim != 2 or R_iX.ndim != 2 or R_ix.shape[0] != R_iX.shape[0]: |
| raise ValueError("R_ix must be (N,d) and R_iX must be (N,D) with same N.") |
|
|
| self.R_ix = R_ix |
| self.R_iX = R_iX |
| self.N, self.d_lat = R_ix.shape |
| _, self.D = R_iX.shape |
|
|
| |
| Y = R_iX.astype(np.float64) |
| if self.center_X: |
| self.mean_X = Y.mean(axis=0) |
| Yc = Y - self.mean_X[None, :] |
| else: |
| self.mean_X = np.zeros((self.D,), dtype=np.float64) |
| Yc = Y |
|
|
| |
| R = R_ix.astype(np.float64) |
| if self.whiten_latent: |
| mu = R.mean(axis=0) |
| sd = np.maximum(R.std(axis=0), 1e-12) |
| self._pp = _LatentPreproc(mu, sd) |
| Rw = self._pp.apply(R) |
| else: |
| self._pp = _LatentPreproc(np.zeros((self.d_lat,), dtype=np.float64), |
| np.ones((self.d_lat,), dtype=np.float64)) |
| Rw = R |
|
|
| self.R_ix_w = np.ascontiguousarray(Rw, dtype=np.float64) |
|
|
| |
| lam = float(self.sigma2) |
| jit = float(self.jitter) |
|
|
| |
| |
| self._A = None |
| self._chol = None |
| self._eig = None |
| self.V_xX = None |
|
|
| if self.solver in ("chol", "eigh"): |
| self._fit_closed_form(Rw=self.R_ix_w, Yc=Yc, lam=lam, jit=jit) |
| elif self.solver == "cg": |
| self._fit_cg(Rw=self.R_ix_w, Yc=Yc, lam=lam, jit=jit) |
| else: |
| raise ValueError("solver must be one of {'chol','eigh','cg'}") |
|
|
| |
| |
| |
| def _fit_closed_form(self, *, Rw: np.ndarray, Yc: np.ndarray, lam: float, jit: float) -> None: |
| |
| G = Rw.T @ Rw |
| d = G.shape[0] |
| A = G + (lam + jit) * np.eye(d, dtype=np.float64) |
| B = Rw.T @ Yc |
|
|
| if self.solver == "eigh": |
| w, V = np.linalg.eigh(A) |
| w = np.maximum(w, float(self.eig_clip)) |
| self._eig = (w, V) |
| |
| self.V_xX = (V @ ((V.T @ B) / w[:, None])).astype(np.float64) |
| self._A = A |
| return |
|
|
| |
| cF = la.cho_factor(A, lower=True, check_finite=False) |
| self._chol = cF |
| self.V_xX = la.cho_solve(cF, B, check_finite=False).astype(np.float64) |
| self._A = A |
|
|
| def _fit_cg(self, *, Rw: np.ndarray, Yc: np.ndarray, lam: float, jit: float) -> None: |
| """ |
| Solve (R^T R + lam I) V = R^T Y with CG in d-space using matvecs. |
| Good when d is too large to factorize directly. |
| """ |
| N, d = Rw.shape |
| lam_eff = lam + jit |
|
|
| def matvec(v: np.ndarray) -> np.ndarray: |
| |
| v = np.asarray(v, dtype=np.float64) |
| return (Rw.T @ (Rw @ v)) + lam_eff * v |
|
|
| Aop = sla.LinearOperator((d, d), matvec=matvec, dtype=np.float64) |
|
|
| B = Rw.T @ Yc |
| V = np.zeros((d, self.D), dtype=np.float64) |
|
|
| |
| for j in range(self.D): |
| x0 = np.zeros((d,), dtype=np.float64) |
| sol, info = sla.cg( |
| Aop, |
| B[:, j], |
| x0=x0, |
| maxiter=self.cg_maxiter, |
| rtol=self.cg_rtol, |
| atol=self.cg_atol, |
| ) |
| if info != 0: |
| raise RuntimeError(f"LGP(CD) CG failed for output dim {j} with info={info}. " |
| f"Try larger cg_maxiter, looser rtol, or larger sigma2.") |
| V[:, j] = sol |
|
|
| self.V_xX = V |
| self._A = None |
| self._chol = None |
| self._eig = None |
|
|
| |
| |
| |
| def __call__(self, R_ax: Union[np.ndarray, list], *, batch_size: Optional[int] = None) -> np.ndarray: |
| """ |
| Mean prediction only. |
| R_ax: (A,d) or (d,) |
| returns: (A,D) or (D,) |
| """ |
| R_ax = np.asarray(R_ax) |
| single = (R_ax.ndim == 1) |
| if single: |
| R_ax = R_ax[None, :] |
|
|
| Ra = np.ascontiguousarray(R_ax.astype(np.float64, copy=False)) |
| if Ra.shape[1] != self.d_lat: |
| raise ValueError(f"Expected latent dim d={self.d_lat}, got {Ra.shape[1]}.") |
|
|
| if batch_size is None: |
| Y = self._decode_mean(Ra) |
| else: |
| bs = int(batch_size) |
| out = [] |
| for s in range(0, Ra.shape[0], bs): |
| out.append(self._decode_mean(Ra[s:s+bs])) |
| Y = np.vstack(out) |
|
|
| return Y[0] if single else Y |
|
|
| def predict( |
| self, |
| R_ax: Union[np.ndarray, list], |
| *, |
| return_var: bool = False, |
| batch_size: Optional[int] = None, |
| include_obs_noise: bool = False, |
| ) -> Union[np.ndarray, Tuple[np.ndarray, np.ndarray]]: |
| """ |
| Predict mean and optional scalar variance proxy per query. |
| |
| var_f(a) = sigma2 * r_a^T (R^T R + sigma2 I)^{-1} r_a |
| |
| If include_obs_noise=True, returns var_y = var_f + sigma2. |
| """ |
| R_ax = np.asarray(R_ax) |
| single = (R_ax.ndim == 1) |
| if single: |
| R_ax = R_ax[None, :] |
|
|
| Ra = np.ascontiguousarray(R_ax.astype(np.float64, copy=False)) |
| if Ra.shape[1] != self.d_lat: |
| raise ValueError(f"Expected latent dim d={self.d_lat}, got {Ra.shape[1]}.") |
|
|
| if not return_var: |
| mean = self.__call__(Ra, batch_size=batch_size) |
| if single and mean.ndim == 1: |
| mean = mean[None, :] |
| return mean[0] if single else mean |
|
|
| if batch_size is None: |
| mean, var = self._decode_mean_var(Ra, include_obs_noise=include_obs_noise) |
| else: |
| bs = int(batch_size) |
| ms, vs = [], [] |
| for s in range(0, Ra.shape[0], bs): |
| m, v = self._decode_mean_var(Ra[s:s+bs], include_obs_noise=include_obs_noise) |
| ms.append(m) |
| vs.append(v) |
| mean = np.vstack(ms) |
| var = np.concatenate(vs, axis=0) |
|
|
| if single: |
| return mean[0], var[0] |
| return mean, var |
|
|
| |
| |
| |
| def _decode_mean(self, Ra: np.ndarray) -> np.ndarray: |
| assert self.V_xX is not None |
|
|
| |
| Raw = self._pp.apply(Ra) |
|
|
| Yc = Raw @ self.V_xX |
| Y = Yc + self.mean_X[None, :] |
| return Y |
|
|
| def _solve_Ainv_vecs(self, X_dA: np.ndarray) -> np.ndarray: |
| """ |
| Solve (R^T R + sigma2 I)^{-1} X for X shape (d,A). |
| Only available for chol/eigh solvers. For CG solver we can do per-column CG. |
| """ |
| X = np.asarray(X_dA, dtype=np.float64) |
|
|
| if self.solver == "chol": |
| assert self._chol is not None |
| return la.cho_solve(self._chol, X, check_finite=False) |
|
|
| if self.solver == "eigh": |
| assert self._eig is not None |
| w, V = self._eig |
| return V @ ((V.T @ X) / w[:, None]) |
|
|
| |
| |
| Rw = self.R_ix_w |
| lam_eff = self.sigma2 + self.jitter |
| d = self.d_lat |
|
|
| def matvec(v: np.ndarray) -> np.ndarray: |
| return (Rw.T @ (Rw @ v)) + lam_eff * v |
|
|
| Aop = sla.LinearOperator((d, d), matvec=matvec, dtype=np.float64) |
|
|
| out = np.empty_like(X) |
| for j in range(X.shape[1]): |
| sol, info = sla.cg( |
| Aop, |
| X[:, j], |
| x0=np.zeros((d,), dtype=np.float64), |
| maxiter=self.cg_maxiter, |
| rtol=self.cg_rtol, |
| atol=self.cg_atol, |
| ) |
| if info != 0: |
| raise RuntimeError(f"LGP variance CG solve failed with info={info}.") |
| out[:, j] = sol |
| return out |
|
|
| def _decode_mean_var(self, Ra: np.ndarray, *, include_obs_noise: bool) -> Tuple[np.ndarray, np.ndarray]: |
| """ |
| Mean + scalar variance proxy: |
| var_f[a] = sigma2 * r_a^T A^{-1} r_a |
| """ |
| assert self.V_xX is not None |
|
|
| Raw = self._pp.apply(Ra) |
|
|
| |
| Yc = Raw @ self.V_xX |
| mean = Yc + self.mean_X[None, :] |
|
|
| |
| |
| Ainv_Rt = self._solve_Ainv_vecs(Raw.T) |
| quad = np.sum(Raw.T * Ainv_Rt, axis=0) |
| var_f = self.sigma2 * quad |
|
|
| if include_obs_noise: |
| var_f = var_f + self.sigma2 |
|
|
| return mean, var_f |
|
|
| |
| |
| |
| def kernel_mass(self, R_ax: Union[np.ndarray, list]) -> np.ndarray: |
| """ |
| For linear kernel, a simple "mass" proxy is ||r||^2. |
| """ |
| R_ax = np.asarray(R_ax, dtype=np.float64) |
| single = (R_ax.ndim == 1) |
| if single: |
| R_ax = R_ax[None, :] |
| Raw = self._pp.apply(R_ax) |
| mass = np.sum(Raw * Raw, axis=1) |
| return mass[0] if single else mass |
|
|
|
|
| __all__ = ["LGP"] |