Update src/dima/gplm.py
Browse files- src/dima/gplm.py +398 -349
src/dima/gplm.py
CHANGED
|
@@ -1,15 +1,15 @@
|
|
| 1 |
# src/dima/gplm.py
|
| 2 |
from __future__ import annotations
|
| 3 |
|
| 4 |
-
from typing import Any, Dict, Literal, Optional, Tuple, Union
|
| 5 |
|
| 6 |
import numpy as np
|
| 7 |
import scipy.linalg as la
|
| 8 |
|
| 9 |
-
|
| 10 |
from .ann import ANNBackend, make_ann
|
| 11 |
from .utils import fps_indices, median_eps_from_knn_d2, sqdist_ab
|
| 12 |
|
|
|
|
| 13 |
InducingMode = Literal["random_subset", "fps", "kmeans_medoids", "given"]
|
| 14 |
|
| 15 |
|
|
@@ -22,6 +22,7 @@ def _kmeans2_safe(Z: np.ndarray, m: int, seed: int = 0) -> np.ndarray:
|
|
| 22 |
m = int(min(max(1, m), Z.shape[0]))
|
| 23 |
try:
|
| 24 |
from scipy.cluster.vq import kmeans2 # type: ignore
|
|
|
|
| 25 |
C, _ = kmeans2(Z.astype(np.float64, copy=False), m, minit="points", seed=seed)
|
| 26 |
return C.astype(Z.dtype, copy=False)
|
| 27 |
except Exception:
|
|
@@ -29,105 +30,254 @@ def _kmeans2_safe(Z: np.ndarray, m: int, seed: int = 0) -> np.ndarray:
|
|
| 29 |
idx = rng.choice(Z.shape[0], size=m, replace=False)
|
| 30 |
return Z[idx]
|
| 31 |
|
| 32 |
-
from dataclasses import dataclass
|
| 33 |
|
| 34 |
-
|
| 35 |
-
"""
|
| 36 |
-
Pairwise squared Euclidean distances between rows:
|
| 37 |
-
A: (a,d), B: (b,d) -> D2: (a,b)
|
| 38 |
"""
|
| 39 |
-
|
| 40 |
-
B = np.asarray(B, dtype=np.float64)
|
| 41 |
-
AA = np.sum(A * A, axis=1, keepdims=True) # (a,1)
|
| 42 |
-
BB = np.sum(B * B, axis=1, keepdims=True).T # (1,b)
|
| 43 |
-
D2 = AA + BB - 2.0 * (A @ B.T)
|
| 44 |
-
return np.maximum(D2, 0.0)
|
| 45 |
|
|
|
|
|
|
|
|
|
|
| 46 |
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
|
|
|
| 60 |
|
| 61 |
Notes:
|
| 62 |
-
-
|
| 63 |
-
|
|
|
|
| 64 |
"""
|
| 65 |
|
| 66 |
def __init__(
|
| 67 |
self,
|
|
|
|
|
|
|
| 68 |
*,
|
| 69 |
-
|
| 70 |
-
M_mX: np.ndarray,
|
| 71 |
-
mean_X: np.ndarray,
|
| 72 |
-
lat_mean_x: Optional[np.ndarray] = None,
|
| 73 |
-
lat_std_x: Optional[np.ndarray] = None,
|
| 74 |
beta: float = 1.0,
|
| 75 |
-
eps: float = 1.0,
|
| 76 |
-
pred_k: Optional[int] = None,
|
| 77 |
-
ann_Z=None,
|
| 78 |
-
dtype=np.float32 ):
|
| 79 |
-
self.dtype = dtype
|
| 80 |
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 84 |
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
|
|
|
|
|
|
| 88 |
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
|
| 93 |
|
| 94 |
-
|
| 95 |
-
|
| 96 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 97 |
self.beta = float(beta)
|
| 98 |
-
self.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 99 |
|
| 100 |
-
|
| 101 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 102 |
|
| 103 |
-
|
| 104 |
-
|
|
|
|
|
|
|
|
|
|
| 105 |
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
raise ValueError("pred_k must be positive or None.")
|
| 109 |
-
if self.ann_Z is None:
|
| 110 |
-
raise ValueError("ann_Z must be provided when pred_k is not None.")
|
| 111 |
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
# -------------------------
|
| 115 |
-
def __call__(
|
| 116 |
-
self,
|
| 117 |
-
R_ax: Union[np.ndarray, list],
|
| 118 |
-
*,
|
| 119 |
-
batch_size: Optional[int] = None,
|
| 120 |
-
jacobian: bool = False,
|
| 121 |
-
metric: bool = False ):
|
| 122 |
-
"""
|
| 123 |
-
Decode latents to ambient.
|
| 124 |
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 131 |
R_ax = np.asarray(R_ax)
|
| 132 |
single = (R_ax.ndim == 1)
|
| 133 |
if single:
|
|
@@ -135,319 +285,218 @@ class GPLM:
|
|
| 135 |
R_ax = np.ascontiguousarray(R_ax.astype(self.dtype, copy=False))
|
| 136 |
|
| 137 |
if batch_size is None:
|
| 138 |
-
|
| 139 |
else:
|
| 140 |
bs = int(batch_size)
|
| 141 |
-
|
| 142 |
-
Js = [] if jacobian else None
|
| 143 |
-
Gs = [] if metric else None
|
| 144 |
for s in range(0, R_ax.shape[0], bs):
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
Ys.append(chunk)
|
| 148 |
-
elif jacobian and (not metric):
|
| 149 |
-
Yc, Jc = chunk
|
| 150 |
-
Ys.append(Yc); Js.append(Jc) # type: ignore[arg-type]
|
| 151 |
-
elif (not jacobian) and metric:
|
| 152 |
-
Yc, gc = chunk
|
| 153 |
-
Ys.append(Yc); Gs.append(gc) # type: ignore[arg-type]
|
| 154 |
-
else:
|
| 155 |
-
Yc, Jc, gc = chunk
|
| 156 |
-
Ys.append(Yc); Js.append(Jc); Gs.append(gc) # type: ignore[arg-type]
|
| 157 |
-
|
| 158 |
-
Y = np.vstack(Ys)
|
| 159 |
-
if (not jacobian) and (not metric):
|
| 160 |
-
out = Y
|
| 161 |
-
elif jacobian and (not metric):
|
| 162 |
-
out = (Y, np.vstack(Js)) # type: ignore[arg-type]
|
| 163 |
-
elif (not jacobian) and metric:
|
| 164 |
-
out = (Y, np.vstack(Gs)) # type: ignore[arg-type]
|
| 165 |
-
else:
|
| 166 |
-
out = (Y, np.vstack(Js), np.vstack(Gs)) # type: ignore[arg-type]
|
| 167 |
|
| 168 |
-
|
| 169 |
-
if single:
|
| 170 |
-
if (not jacobian) and (not metric):
|
| 171 |
-
return out[0] # type: ignore[index]
|
| 172 |
-
if jacobian and (not metric):
|
| 173 |
-
Y, J = out # type: ignore[misc]
|
| 174 |
-
return Y[0], J[0]
|
| 175 |
-
if (not jacobian) and metric:
|
| 176 |
-
Y, g = out # type: ignore[misc]
|
| 177 |
-
return Y[0], g[0]
|
| 178 |
-
Y, J, g = out # type: ignore[misc]
|
| 179 |
-
return Y[0], J[0], g[0]
|
| 180 |
-
|
| 181 |
-
return out
|
| 182 |
-
|
| 183 |
-
def _decode(
|
| 184 |
-
self,
|
| 185 |
-
R_ax: np.ndarray,
|
| 186 |
-
*,
|
| 187 |
-
jacobian: bool,
|
| 188 |
-
metric: bool ):
|
| 189 |
-
"""
|
| 190 |
-
Internal decode for a batch (a,d).
|
| 191 |
-
|
| 192 |
-
Jacobian is w.r.t. *unwhitened* input R_ax (chain rule applied if whiten_latent=True).
|
| 193 |
-
"""
|
| 194 |
-
Za = R_ax.astype(np.float64) # (a,d)
|
| 195 |
-
Za_w = (Za - self.lat_mean_x) / self.lat_std_x # (a,d)
|
| 196 |
|
| 197 |
-
|
| 198 |
-
|
|
|
|
| 199 |
|
| 200 |
if self.pred_k is None or self.pred_k == self.m:
|
| 201 |
-
|
| 202 |
-
|
| 203 |
-
|
| 204 |
-
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
|
| 208 |
-
|
| 209 |
-
# diff in whitened latent space: (a,m,d)
|
| 210 |
-
diff_amd = Za_w[:, None, :] - self.Z_mx_w[None, :, :]
|
| 211 |
-
|
| 212 |
-
# dW/dZa_w: (a,m,d)
|
| 213 |
-
dW_amd = (self._rbf_grad_scale * W_am[:, :, None]) * diff_amd
|
| 214 |
-
|
| 215 |
-
# chain to unwhitened: dW/dZa = dW/dZa_w * dZa_w/dZa = * (1/std)
|
| 216 |
-
dW_amd *= inv_std[None, None, :]
|
| 217 |
-
|
| 218 |
-
if jacobian:
|
| 219 |
-
# J_{a,D,d} = Σ_m dW_{a,m,d} * M_{m,D}
|
| 220 |
-
J_aDd = np.einsum("amd,mD->aDd", dW_amd, self.M_mX, optimize=True)
|
| 221 |
-
if metric:
|
| 222 |
-
g_add = np.einsum("aDd,aDe->ade", J_aDd, J_aDd, optimize=True)
|
| 223 |
-
return (Y + self.mean_X[None, :], J_aDd, g_add)
|
| 224 |
-
return (Y + self.mean_X[None, :], J_aDd)
|
| 225 |
-
|
| 226 |
-
# metric-only
|
| 227 |
-
# Build J implicitly then g; for simplicity we materialize J here.
|
| 228 |
-
J_aDd = np.einsum("amd,mD->aDd", dW_amd, self.M_mX, optimize=True)
|
| 229 |
-
g_add = np.einsum("aDd,aDe->ade", J_aDd, J_aDd, optimize=True)
|
| 230 |
-
return (Y + self.mean_X[None, :], g_add)
|
| 231 |
-
|
| 232 |
-
# sparse inducing neighbors for speed
|
| 233 |
-
j_aK, D2_aK = self.ann_Z.search(Za_w.astype(self.dtype, copy=False), self.pred_k) # (a,k), (a,k)
|
| 234 |
-
W_aK = np.exp(-self.beta * (D2_aK.astype(np.float64) / self.eps)) # (a,k)
|
| 235 |
-
M_aKD = self.M_mX[j_aK] # (a,k,D)
|
| 236 |
-
Y = np.sum(W_aK[:, :, None] * M_aKD, axis=1) # (a,D)
|
| 237 |
-
|
| 238 |
-
if not (jacobian or metric):
|
| 239 |
-
return Y + self.mean_X[None, :]
|
| 240 |
-
|
| 241 |
-
Zsel_aKd = self.Z_mx_w[j_aK] # (a,k,d) in whitened latent
|
| 242 |
-
diff_aKd = Za_w[:, None, :] - Zsel_aKd # (a,k,d)
|
| 243 |
-
|
| 244 |
-
dW_aKd = (self._rbf_grad_scale * W_aK[:, :, None]) * diff_aKd
|
| 245 |
-
dW_aKd *= inv_std[None, None, :]
|
| 246 |
-
|
| 247 |
-
if jacobian:
|
| 248 |
-
# J_{a,D,d} = Σ_k dW_{a,k,d} * M_{a,k,D}
|
| 249 |
-
J_aDd = np.einsum("akd,akD->aDd", dW_aKd, M_aKD, optimize=True)
|
| 250 |
-
if metric:
|
| 251 |
-
g_add = np.einsum("aDd,aDe->ade", J_aDd, J_aDd, optimize=True)
|
| 252 |
-
return (Y + self.mean_X[None, :], J_aDd, g_add)
|
| 253 |
-
return (Y + self.mean_X[None, :], J_aDd)
|
| 254 |
-
|
| 255 |
-
# metric-only
|
| 256 |
-
J_aDd = np.einsum("akd,akD->aDd", dW_aKd, M_aKD, optimize=True)
|
| 257 |
-
g_add = np.einsum("aDd,aDe->ade", J_aDd, J_aDd, optimize=True)
|
| 258 |
-
return (Y + self.mean_X[None, :], g_add)
|
| 259 |
-
|
| 260 |
-
# -------------------------
|
| 261 |
-
# NEW: Geodesic flow (analytic Christoffels) with generalized velocity-Verlet
|
| 262 |
-
# -------------------------
|
| 263 |
-
def flow(
|
| 264 |
-
self,
|
| 265 |
-
R_ax: Union[np.ndarray, list],
|
| 266 |
-
v0_ax: Union[np.ndarray, list],
|
| 267 |
-
*,
|
| 268 |
-
dt: float = 0.05,
|
| 269 |
-
reg: float = 1e-8,
|
| 270 |
-
keep_speed: bool = True,
|
| 271 |
-
) -> Tuple[np.ndarray, np.ndarray]:
|
| 272 |
-
"""
|
| 273 |
-
One generalized velocity-Verlet / leapfrog step for the geodesic ODE:
|
| 274 |
-
r_dot = v
|
| 275 |
-
v_dot^x = - Γ^x_{yz}(r) v^y v^z
|
| 276 |
-
|
| 277 |
-
Returns:
|
| 278 |
-
Q_ax : next latent positions (a,d)
|
| 279 |
-
v_ax : next latent velocities (a,d)
|
| 280 |
-
|
| 281 |
-
Use in a loop:
|
| 282 |
-
Q, v = R0, v0
|
| 283 |
-
for _ in range(T):
|
| 284 |
-
Q, v = gplm.flow(Q, v, dt=...)
|
| 285 |
-
"""
|
| 286 |
-
Z = np.asarray(R_ax, dtype=np.float64)
|
| 287 |
-
v = np.asarray(v0_ax, dtype=np.float64)
|
| 288 |
-
|
| 289 |
-
single = (Z.ndim == 1)
|
| 290 |
-
if single:
|
| 291 |
-
Z = Z[None, :]
|
| 292 |
-
if v.ndim == 1:
|
| 293 |
-
v = v[None, :]
|
| 294 |
|
| 295 |
-
|
| 296 |
-
raise ValueError(f"R_ax and v0_ax must have same shape; got {Z.shape} vs {v.shape}")
|
| 297 |
|
| 298 |
-
|
| 299 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 300 |
|
| 301 |
-
|
| 302 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 303 |
|
| 304 |
-
|
| 305 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 306 |
|
| 307 |
-
|
| 308 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 309 |
|
| 310 |
-
#
|
| 311 |
-
|
|
|
|
| 312 |
|
| 313 |
-
#
|
| 314 |
-
|
| 315 |
-
s0 = np.sqrt(np.maximum(np.einsum("ai,aij,aj->a", v, g0, v, optimize=True), 1e-16))
|
| 316 |
-
s1 = np.sqrt(np.maximum(np.einsum("ai,aij,aj->a", v_new, g1, v_new, optimize=True), 1e-16))
|
| 317 |
-
v_new = v_new * (s0 / s1)[:, None]
|
| 318 |
|
| 319 |
-
|
| 320 |
-
|
| 321 |
-
|
| 322 |
|
| 323 |
-
|
| 324 |
-
|
| 325 |
-
|
| 326 |
-
v_ax: np.ndarray,
|
| 327 |
-
*,
|
| 328 |
-
reg: float,
|
| 329 |
-
) -> Tuple[np.ndarray, np.ndarray]:
|
| 330 |
-
"""
|
| 331 |
-
Compute acceleration a_ax = -Γ^x_{yz}(r) v^y v^z using analytic Christoffels,
|
| 332 |
-
and return (a_ax, g_aij).
|
| 333 |
-
|
| 334 |
-
Shapes:
|
| 335 |
-
R_ax: (a,d)
|
| 336 |
-
v_ax: (a,d)
|
| 337 |
-
a_ax: (a,d)
|
| 338 |
-
g: (a,d,d)
|
| 339 |
-
"""
|
| 340 |
-
Z = np.asarray(R_ax, dtype=np.float64)
|
| 341 |
-
v = np.asarray(v_ax, dtype=np.float64)
|
| 342 |
-
a, d = Z.shape
|
| 343 |
-
if v.shape != (a, d):
|
| 344 |
-
raise ValueError(f"v_ax must have shape {(a, d)}, got {v.shape}")
|
| 345 |
|
| 346 |
-
#
|
| 347 |
-
|
| 348 |
-
|
| 349 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 350 |
|
| 351 |
-
|
| 352 |
|
| 353 |
-
|
| 354 |
-
|
| 355 |
-
diff_amd = Za_w[:, None, :] - self.Z_mx_w[None, :, :] # (a,m,d)
|
| 356 |
-
D2_am = np.sum(diff_amd * diff_amd, axis=2) # (a,m)
|
| 357 |
-
k_am = np.exp(-self.beta * (D2_am / self.eps)) # (a,m)
|
| 358 |
|
| 359 |
-
|
| 360 |
-
|
| 361 |
-
dk_amd *= inv_std[None, None, :] # chain to unwhitened
|
| 362 |
|
| 363 |
-
|
| 364 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 365 |
|
| 366 |
-
|
| 367 |
-
|
| 368 |
-
|
| 369 |
|
| 370 |
-
|
| 371 |
|
| 372 |
-
|
| 373 |
-
|
| 374 |
|
| 375 |
-
|
| 376 |
-
|
| 377 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 378 |
|
| 379 |
-
|
| 380 |
-
|
| 381 |
-
|
| 382 |
-
+ alpha * I[w][None, None, :]
|
| 383 |
-
) # (a,m,d) in whitened diffs
|
| 384 |
-
d2k_amy *= (inv_std[w] * inv_std[None])[None, None, :] # chain factors (1/σ_w)(1/σ_y)
|
| 385 |
|
| 386 |
-
|
| 387 |
-
|
| 388 |
|
| 389 |
-
|
| 390 |
-
|
| 391 |
-
|
| 392 |
-
|
| 393 |
-
|
| 394 |
-
|
| 395 |
|
| 396 |
-
|
| 397 |
-
|
| 398 |
-
|
| 399 |
-
|
|
|
|
|
|
|
| 400 |
|
| 401 |
-
|
| 402 |
-
diff_aKd = Za_w[:, None, :] - Zsel_aKd # (a,k,d)
|
| 403 |
-
M_aKX = self.M_mX[j_aK] # (a,k,D)
|
| 404 |
|
| 405 |
-
|
| 406 |
-
|
|
|
|
| 407 |
|
| 408 |
-
|
| 409 |
-
|
| 410 |
|
| 411 |
-
|
| 412 |
-
|
| 413 |
|
| 414 |
-
|
|
|
|
| 415 |
|
| 416 |
-
|
|
|
|
| 417 |
|
| 418 |
-
|
| 419 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 420 |
|
| 421 |
-
|
| 422 |
-
|
| 423 |
-
|
| 424 |
-
) # (a,k,d)
|
| 425 |
-
d2k_aKy *= (inv_std[w] * inv_std[None])[None, None, :]
|
| 426 |
|
| 427 |
-
|
|
|
|
|
|
|
|
|
|
| 428 |
|
| 429 |
-
|
| 430 |
-
np.einsum("aXd,aXe->ade", H_w_aXy, J_aXy, optimize=True)
|
| 431 |
-
+ np.einsum("aXd,aXe->ade", J_aXy, H_w_aXy, optimize=True)
|
| 432 |
-
)
|
| 433 |
-
dg_awyz[:, w, :, :] = dg_w
|
| 434 |
|
| 435 |
-
|
| 436 |
-
|
| 437 |
-
|
| 438 |
-
|
| 439 |
-
|
| 440 |
-
+ dg_awyz.transpose(0, 2, 1, 3)
|
| 441 |
-
- dg_awyz.transpose(0, 2, 3, 1)
|
| 442 |
-
) # (a,y,z,w)
|
| 443 |
|
| 444 |
-
|
| 445 |
-
|
| 446 |
|
| 447 |
-
|
| 448 |
-
|
| 449 |
|
| 450 |
-
|
|
|
|
|
|
|
| 451 |
|
| 452 |
|
| 453 |
-
__all__ = ["GPLM", "InducingMode"]
|
|
|
|
| 1 |
# src/dima/gplm.py
|
| 2 |
from __future__ import annotations
|
| 3 |
|
| 4 |
+
from typing import Any, Dict, Literal, Optional, Tuple, Union
|
| 5 |
|
| 6 |
import numpy as np
|
| 7 |
import scipy.linalg as la
|
| 8 |
|
|
|
|
| 9 |
from .ann import ANNBackend, make_ann
|
| 10 |
from .utils import fps_indices, median_eps_from_knn_d2, sqdist_ab
|
| 11 |
|
| 12 |
+
|
| 13 |
InducingMode = Literal["random_subset", "fps", "kmeans_medoids", "given"]
|
| 14 |
|
| 15 |
|
|
|
|
| 22 |
m = int(min(max(1, m), Z.shape[0]))
|
| 23 |
try:
|
| 24 |
from scipy.cluster.vq import kmeans2 # type: ignore
|
| 25 |
+
# minit="points" picks initial centers from data -> stable for medoids snapping
|
| 26 |
C, _ = kmeans2(Z.astype(np.float64, copy=False), m, minit="points", seed=seed)
|
| 27 |
return C.astype(Z.dtype, copy=False)
|
| 28 |
except Exception:
|
|
|
|
| 30 |
idx = rng.choice(Z.shape[0], size=m, replace=False)
|
| 31 |
return Z[idx]
|
| 32 |
|
|
|
|
| 33 |
|
| 34 |
+
class GPLM:
|
|
|
|
|
|
|
|
|
|
| 35 |
"""
|
| 36 |
+
Inducing-point / Nyström GP (kernel ridge) decoder on latents.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 37 |
|
| 38 |
+
Training:
|
| 39 |
+
R_ix: (N,d) latents
|
| 40 |
+
R_iX: (N,D) ambients
|
| 41 |
|
| 42 |
+
Choose inducing Z_mx (m << N), typically subset (medoids) of R_ix.
|
| 43 |
+
|
| 44 |
+
Latent kernel (unnormalized Gaussian affinity):
|
| 45 |
+
C_im = exp(-β * ||R_ix - Z_mx||^2 / ε) (N,m)
|
| 46 |
+
W_mn = exp(-β * ||Z_mx - Z_nx||^2 / ε) (m,m)
|
| 47 |
+
|
| 48 |
+
Nyström KRR/GP mean reduced solve:
|
| 49 |
+
M_mX = (C^T C + σ2 W + jitter I)^-1 (C^T (R_iX - mean_X))
|
| 50 |
+
|
| 51 |
+
Predict:
|
| 52 |
+
For novel R_ax:
|
| 53 |
+
find κ inducing neighbors (or all m if pred_κ=None)
|
| 54 |
+
C_am = exp(-β * ||R_ax - Z_mx||^2 / ε)
|
| 55 |
+
R_aX = C_am M_mX + mean_X
|
| 56 |
|
| 57 |
Notes:
|
| 58 |
+
- This implementation supports BOTH ascii kwargs and unicode kwargs
|
| 59 |
+
(β, ε, κ_eps, σ2, pred_κ, ε_use_kth, ε_mul).
|
| 60 |
+
- If whiten_latent=True, distances are computed in whitened latent space.
|
| 61 |
"""
|
| 62 |
|
| 63 |
def __init__(
|
| 64 |
self,
|
| 65 |
+
R_ix: np.ndarray,
|
| 66 |
+
R_iX: np.ndarray,
|
| 67 |
*,
|
| 68 |
+
# ASCII names (preferred for library APIs)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 69 |
beta: float = 1.0,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 70 |
|
| 71 |
+
# ε estimation
|
| 72 |
+
eps: Optional[float] = None,
|
| 73 |
+
k_eps: int = 256,
|
| 74 |
+
eps_use_kth: bool = True,
|
| 75 |
+
eps_mul: float = 1.0,
|
| 76 |
+
|
| 77 |
+
# regularization
|
| 78 |
+
sigma2: float = 1e-5,
|
| 79 |
+
jitter: float = 1e-8,
|
| 80 |
|
| 81 |
+
# inducing
|
| 82 |
+
m: int = 1024,
|
| 83 |
+
inducing: InducingMode = "kmeans_medoids",
|
| 84 |
+
Z_mx: Optional[np.ndarray] = None,
|
| 85 |
+
seed: int = 0,
|
| 86 |
|
| 87 |
+
# preprocess
|
| 88 |
+
center_X: bool = True,
|
| 89 |
+
whiten_latent: bool = False,
|
| 90 |
+
dtype: Any = np.float32,
|
| 91 |
|
| 92 |
+
# compute/memory
|
| 93 |
+
fit_block: int = 8192,
|
| 94 |
|
| 95 |
+
# inference
|
| 96 |
+
pred_k: Optional[int] = None,
|
| 97 |
+
ann_backend: ANNBackend = "auto",
|
| 98 |
+
ann_params: Optional[Dict[str, Any]] = None,
|
| 99 |
+
n_jobs: int = -1,
|
| 100 |
+
|
| 101 |
+
# accept unicode kwargs (β, ε, κ_eps, σ2, pred_κ, ...)
|
| 102 |
+
**kwargs: Any,
|
| 103 |
+
):
|
| 104 |
+
# ---- map unicode kwargs -> ascii ----
|
| 105 |
+
if "β" in kwargs:
|
| 106 |
+
beta = kwargs.pop("β")
|
| 107 |
+
if "ε" in kwargs:
|
| 108 |
+
eps = kwargs.pop("ε")
|
| 109 |
+
if "κ_eps" in kwargs:
|
| 110 |
+
k_eps = kwargs.pop("κ_eps")
|
| 111 |
+
if "ε_use_kth" in kwargs:
|
| 112 |
+
eps_use_kth = kwargs.pop("ε_use_kth")
|
| 113 |
+
if "ε_mul" in kwargs:
|
| 114 |
+
eps_mul = kwargs.pop("ε_mul")
|
| 115 |
+
if "σ2" in kwargs:
|
| 116 |
+
sigma2 = kwargs.pop("σ2")
|
| 117 |
+
if "pred_κ" in kwargs:
|
| 118 |
+
pred_k = kwargs.pop("pred_κ")
|
| 119 |
+
|
| 120 |
+
if kwargs:
|
| 121 |
+
raise TypeError(f"Unexpected kwargs: {sorted(kwargs.keys())}")
|
| 122 |
+
|
| 123 |
+
# ---- store params (provide both spellings) ----
|
| 124 |
self.beta = float(beta)
|
| 125 |
+
self.β = self.beta
|
| 126 |
+
|
| 127 |
+
self.sigma2 = float(sigma2)
|
| 128 |
+
self.σ2 = self.sigma2
|
| 129 |
+
|
| 130 |
+
self.jitter = float(jitter)
|
| 131 |
+
self.seed = int(seed)
|
| 132 |
+
self.dtype = dtype
|
| 133 |
+
self.fit_block = int(fit_block)
|
| 134 |
+
|
| 135 |
+
# ---- validate / cast ----
|
| 136 |
+
R_ix = np.ascontiguousarray(np.asarray(R_ix).astype(self.dtype, copy=False))
|
| 137 |
+
R_iX = np.ascontiguousarray(np.asarray(R_iX).astype(self.dtype, copy=False))
|
| 138 |
+
if R_ix.ndim != 2 or R_iX.ndim != 2 or R_ix.shape[0] != R_iX.shape[0]:
|
| 139 |
+
raise ValueError("R_ix must be (N,d) and R_iX must be (N,D) with same N.")
|
| 140 |
+
|
| 141 |
+
self.R_ix = R_ix
|
| 142 |
+
self.R_iX = R_iX
|
| 143 |
+
self.N, self.d_lat = R_ix.shape
|
| 144 |
+
_, self.D = R_iX.shape
|
| 145 |
+
|
| 146 |
+
# ---- center output ----
|
| 147 |
+
if center_X:
|
| 148 |
+
self.mean_X = R_iX.mean(axis=0).astype(np.float64)
|
| 149 |
+
Y = (R_iX.astype(np.float64) - self.mean_X[None, :])
|
| 150 |
+
else:
|
| 151 |
+
self.mean_X = np.zeros((self.D,), dtype=np.float64)
|
| 152 |
+
Y = R_iX.astype(np.float64)
|
| 153 |
+
|
| 154 |
+
# ---- latent whitening (optional) ----
|
| 155 |
+
Ztrain = R_ix.astype(np.float64)
|
| 156 |
+
if whiten_latent:
|
| 157 |
+
self.lat_mean_x = Ztrain.mean(axis=0)
|
| 158 |
+
self.lat_std_x = np.maximum(Ztrain.std(axis=0), 1e-12)
|
| 159 |
+
Ztrain_w = (Ztrain - self.lat_mean_x) / self.lat_std_x
|
| 160 |
+
else:
|
| 161 |
+
self.lat_mean_x = np.zeros((self.d_lat,), dtype=np.float64)
|
| 162 |
+
self.lat_std_x = np.ones((self.d_lat,), dtype=np.float64)
|
| 163 |
+
Ztrain_w = Ztrain
|
| 164 |
+
|
| 165 |
+
self.R_ix_w = Ztrain_w # (N,d) in float64
|
| 166 |
+
|
| 167 |
+
# ---- ANN on training latents (for eps + medoids snapping) ----
|
| 168 |
+
self.ann_train, self.ann_backend = make_ann(ann_backend, ann_params=ann_params, n_jobs=n_jobs)
|
| 169 |
+
self.ann_train.build(self.R_ix_w.astype(self.dtype, copy=False))
|
| 170 |
+
|
| 171 |
+
# ---- eps via kNN distances on latents ----
|
| 172 |
+
if eps is None:
|
| 173 |
+
k_eps = int(min(max(8, int(k_eps)), self.N - 1))
|
| 174 |
+
# ask for k_eps+1 to try to include self
|
| 175 |
+
j_iK1, D2_iK1 = self.ann_train.search(self.R_ix_w.astype(self.dtype, copy=False), k_eps + 1)
|
| 176 |
+
|
| 177 |
+
i = np.arange(self.N)[:, None]
|
| 178 |
+
is_self = (j_iK1 == i)
|
| 179 |
+
|
| 180 |
+
if np.any(is_self):
|
| 181 |
+
D2_iK = np.empty((self.N, k_eps), dtype=np.float64)
|
| 182 |
+
for ii in range(self.N):
|
| 183 |
+
keep = (j_iK1[ii] != ii)
|
| 184 |
+
D2_iK[ii] = D2_iK1[ii][keep][:k_eps]
|
| 185 |
+
else:
|
| 186 |
+
D2_iK = D2_iK1[:, :k_eps].astype(np.float64, copy=False)
|
| 187 |
|
| 188 |
+
eps_hat = median_eps_from_knn_d2(D2_iK, use_kth=bool(eps_use_kth))
|
| 189 |
+
else:
|
| 190 |
+
eps_hat = float(eps)
|
| 191 |
+
|
| 192 |
+
eps_hat *= float(eps_mul)
|
| 193 |
+
if eps_hat <= 0:
|
| 194 |
+
raise ValueError("eps must be > 0.")
|
| 195 |
+
self.eps = float(eps_hat)
|
| 196 |
+
self.ε = self.eps
|
| 197 |
+
|
| 198 |
+
# ---- choose inducing points (in whitened latent space) ----
|
| 199 |
+
rng = np.random.default_rng(self.seed)
|
| 200 |
+
m = int(min(max(1, int(m)), self.N))
|
| 201 |
+
|
| 202 |
+
if Z_mx is not None:
|
| 203 |
+
Zm = np.asarray(Z_mx, dtype=np.float64)
|
| 204 |
+
if Zm.ndim != 2 or Zm.shape[1] != self.d_lat:
|
| 205 |
+
raise ValueError("Z_mx must be (m, d_lat).")
|
| 206 |
+
Zm_w = (Zm - self.lat_mean_x) / self.lat_std_x
|
| 207 |
+
else:
|
| 208 |
+
if inducing == "random_subset":
|
| 209 |
+
idx = rng.choice(self.N, size=m, replace=False)
|
| 210 |
+
Zm_w = self.R_ix_w[idx]
|
| 211 |
+
elif inducing == "fps":
|
| 212 |
+
idx = fps_indices(self.R_ix_w, m=m, seed=self.seed)
|
| 213 |
+
Zm_w = self.R_ix_w[idx]
|
| 214 |
+
elif inducing == "kmeans_medoids":
|
| 215 |
+
C = _kmeans2_safe(self.R_ix_w, m, seed=self.seed).astype(np.float64, copy=False)
|
| 216 |
+
j_cm, _ = self.ann_train.search(C.astype(self.dtype, copy=False), 1)
|
| 217 |
+
idx = j_cm.reshape(-1).astype(np.int64)
|
| 218 |
+
|
| 219 |
+
# de-duplicate and refill if needed
|
| 220 |
+
idx_u = np.unique(idx)
|
| 221 |
+
if idx_u.size < m:
|
| 222 |
+
needed = m - idx_u.size
|
| 223 |
+
pool = np.setdiff1d(np.arange(self.N), idx_u, assume_unique=False)
|
| 224 |
+
if pool.size >= needed:
|
| 225 |
+
extra = rng.choice(pool, size=needed, replace=False)
|
| 226 |
+
else:
|
| 227 |
+
extra = rng.choice(self.N, size=needed, replace=True)
|
| 228 |
+
idx = np.concatenate([idx_u, extra])
|
| 229 |
+
else:
|
| 230 |
+
idx = idx_u[:m]
|
| 231 |
|
| 232 |
+
Zm_w = self.R_ix_w[idx]
|
| 233 |
+
elif inducing == "given":
|
| 234 |
+
raise ValueError("Provide Z_mx when inducing='given'.")
|
| 235 |
+
else:
|
| 236 |
+
raise ValueError(f"Unknown inducing mode: {inducing!r}")
|
| 237 |
|
| 238 |
+
self.Z_mx_w = np.ascontiguousarray(Zm_w.astype(np.float64, copy=False))
|
| 239 |
+
self.m = int(self.Z_mx_w.shape[0])
|
|
|
|
|
|
|
|
|
|
| 240 |
|
| 241 |
+
# also store raw inducing points (unwhitened) for serialization convenience
|
| 242 |
+
self.Z_mx = (self.Z_mx_w * self.lat_std_x[None, :]) + self.lat_mean_x[None, :]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 243 |
|
| 244 |
+
# ---- ANN on inducing points for fast prediction ----
|
| 245 |
+
self.ann_Z, _ = make_ann(ann_backend, ann_params=ann_params, n_jobs=n_jobs)
|
| 246 |
+
self.ann_Z.build(self.Z_mx_w.astype(self.dtype, copy=False))
|
| 247 |
+
|
| 248 |
+
# pred_k
|
| 249 |
+
if pred_k is None:
|
| 250 |
+
self.pred_k = None
|
| 251 |
+
else:
|
| 252 |
+
self.pred_k = int(min(max(1, int(pred_k)), self.m))
|
| 253 |
+
self.pred_κ = self.pred_k # unicode alias
|
| 254 |
+
|
| 255 |
+
# ---- W_mm ----
|
| 256 |
+
D2_mm = sqdist_ab(self.Z_mx_w, self.Z_mx_w)
|
| 257 |
+
W_mm = np.exp(-self.beta * (D2_mm.astype(np.float64) / self.eps))
|
| 258 |
+
W_mm.flat[:: self.m + 1] += self.jitter
|
| 259 |
+
self.W_mm = W_mm # (m,m)
|
| 260 |
+
|
| 261 |
+
# ---- accumulate G=C^T C and B=C^T Y ----
|
| 262 |
+
G_mm = np.zeros((self.m, self.m), dtype=np.float64)
|
| 263 |
+
B_mX = np.zeros((self.m, self.D), dtype=np.float64)
|
| 264 |
+
|
| 265 |
+
bs = int(self.fit_block)
|
| 266 |
+
for i0 in range(0, self.N, bs):
|
| 267 |
+
i1 = min(self.N, i0 + bs)
|
| 268 |
+
Zi = self.R_ix_w[i0:i1] # (b,d) float64
|
| 269 |
+
D2_im = sqdist_ab(Zi, self.Z_mx_w)
|
| 270 |
+
C_im = np.exp(-self.beta * (D2_im.astype(np.float64) / self.eps))
|
| 271 |
+
G_mm += C_im.T @ C_im
|
| 272 |
+
B_mX += C_im.T @ Y[i0:i1]
|
| 273 |
+
|
| 274 |
+
A_mm = G_mm + self.sigma2 * W_mm
|
| 275 |
+
A_mm.flat[:: self.m + 1] += self.jitter
|
| 276 |
+
|
| 277 |
+
cF = la.cho_factor(A_mm, lower=True, check_finite=False)
|
| 278 |
+
self.M_mX = la.cho_solve(cF, B_mX, check_finite=False) # (m,D)
|
| 279 |
+
|
| 280 |
+
def __call__(self, R_ax: Union[np.ndarray, list], *, batch_size: Optional[int] = None) -> np.ndarray:
|
| 281 |
R_ax = np.asarray(R_ax)
|
| 282 |
single = (R_ax.ndim == 1)
|
| 283 |
if single:
|
|
|
|
| 285 |
R_ax = np.ascontiguousarray(R_ax.astype(self.dtype, copy=False))
|
| 286 |
|
| 287 |
if batch_size is None:
|
| 288 |
+
Y = self._decode(R_ax)
|
| 289 |
else:
|
| 290 |
bs = int(batch_size)
|
| 291 |
+
out = []
|
|
|
|
|
|
|
| 292 |
for s in range(0, R_ax.shape[0], bs):
|
| 293 |
+
out.append(self._decode(R_ax[s:s + bs]))
|
| 294 |
+
Y = np.vstack(out)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 295 |
|
| 296 |
+
return Y[0] if single else Y
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 297 |
|
| 298 |
+
def _decode(self, R_ax: np.ndarray) -> np.ndarray:
|
| 299 |
+
Za = R_ax.astype(np.float64)
|
| 300 |
+
Za_w = (Za - self.lat_mean_x) / self.lat_std_x
|
| 301 |
|
| 302 |
if self.pred_k is None or self.pred_k == self.m:
|
| 303 |
+
D2_am = sqdist_ab(Za_w, self.Z_mx_w)
|
| 304 |
+
C_am = np.exp(-self.beta * (D2_am.astype(np.float64) / self.eps))
|
| 305 |
+
Y = C_am @ self.M_mX
|
| 306 |
+
else:
|
| 307 |
+
j_aK, D2_aK = self.ann_Z.search(Za_w.astype(self.dtype, copy=False), self.pred_k)
|
| 308 |
+
W = np.exp(-self.beta * (D2_aK.astype(np.float64) / self.eps)) # (a,k)
|
| 309 |
+
M = self.M_mX[j_aK] # (a,k,D)
|
| 310 |
+
Y = np.sum(W[:, :, None] * M, axis=1) # (a,D)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 311 |
|
| 312 |
+
return Y + self.mean_X[None, :]
|
|
|
|
| 313 |
|
| 314 |
+
# -----------------------------
|
| 315 |
+
# Geometry helpers (RBF kernel)
|
| 316 |
+
# -----------------------------
|
| 317 |
+
def _whiten(self, R_ax: np.ndarray) -> np.ndarray:
|
| 318 |
+
R = np.asarray(R_ax, dtype=np.float64)
|
| 319 |
+
return (R - self.lat_mean_x[None, :]) / self.lat_std_x[None, :]
|
| 320 |
|
| 321 |
+
def _inducing_idx(self, Rw_ax: np.ndarray) -> Optional[np.ndarray]:
|
| 322 |
+
"""
|
| 323 |
+
Return inducing indices per query if pred_k is set; else None means "use all m".
|
| 324 |
+
Shape: (A, k) int64 if pred_k; else None.
|
| 325 |
+
"""
|
| 326 |
+
if self.pred_k is None or self.pred_k == self.m:
|
| 327 |
+
return None
|
| 328 |
+
j_aK, _ = self.ann_Z.search(Rw_ax.astype(self.dtype, copy=False), self.pred_k)
|
| 329 |
+
return j_aK.astype(np.int64, copy=False)
|
| 330 |
|
| 331 |
+
def _rbf_terms_single(
|
| 332 |
+
self,
|
| 333 |
+
r_x: np.ndarray,
|
| 334 |
+
v_x: np.ndarray,
|
| 335 |
+
*,
|
| 336 |
+
idx_m: Optional[np.ndarray],
|
| 337 |
+
) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
|
| 338 |
+
"""
|
| 339 |
+
For a single point (no batch), compute:
|
| 340 |
+
g_dd: pullback metric in latent coords (d,d)
|
| 341 |
+
cF: Cholesky factorization handle for solves with g
|
| 342 |
+
f_d: geodesic momentum force term (d,)
|
| 343 |
+
p_d: momentum corresponding to input velocity (p = g v)
|
| 344 |
|
| 345 |
+
Uses efficient contractions with C_mm (no dependence on ambient dimension D).
|
| 346 |
+
"""
|
| 347 |
+
d = self.d_lat
|
| 348 |
+
c = self.beta / self.eps # = beta/eps
|
| 349 |
+
inv_std = 1.0 / self.lat_std_x # (d,)
|
| 350 |
+
inv_std2 = inv_std * inv_std # (d,)
|
| 351 |
+
|
| 352 |
+
# whitened coordinates (float64)
|
| 353 |
+
rw = (r_x.astype(np.float64) - self.lat_mean_x) / self.lat_std_x # (d,)
|
| 354 |
+
|
| 355 |
+
# select inducing points
|
| 356 |
+
if idx_m is None:
|
| 357 |
+
Zw = self.Z_mx_w # (m,d)
|
| 358 |
+
C = self.C_mm # (m,m)
|
| 359 |
+
else:
|
| 360 |
+
Zw = self.Z_mx_w[idx_m] # (k,d)
|
| 361 |
+
C = self.C_mm[np.ix_(idx_m, idx_m)] # (k,k)
|
| 362 |
|
| 363 |
+
Dw = rw[None, :] - Zw # (k,d) (k=m if idx_m is None)
|
| 364 |
+
D2 = np.sum(Dw * Dw, axis=1) # (k,)
|
| 365 |
+
k_m = np.exp(-c * D2) # (k,)
|
| 366 |
|
| 367 |
+
# Dw_over = Dw / std (this is the extra chain-rule factor for derivatives wrt unwhitened R)
|
| 368 |
+
Dw_over = Dw * inv_std[None, :] # (k,d)
|
|
|
|
|
|
|
|
|
|
| 369 |
|
| 370 |
+
# s_m = sum_y v_y * (Dw_y / std_y) (k,)
|
| 371 |
+
v = v_x.astype(np.float64, copy=False)
|
| 372 |
+
s_m = Dw_over @ v # (k,)
|
| 373 |
|
| 374 |
+
# grad_k[m,x] = d/dR_x k_m
|
| 375 |
+
# = -(2c) * k_m * (Dw_x / std_x)
|
| 376 |
+
grad_k = -(2.0 * c) * (k_m[:, None] * Dw_over) # (k,d)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 377 |
|
| 378 |
+
# pullback metric: g = grad^T C grad (d,d)
|
| 379 |
+
# (this equals J^T J with J_Xx = sum_m (dk_m/dR_x) M_mX)
|
| 380 |
+
Cg = C @ grad_k # (k,d)
|
| 381 |
+
g = grad_k.T @ Cg # (d,d)
|
| 382 |
+
g = 0.5 * (g + g.T) # symmetrize for numerical stability
|
| 383 |
+
# regularize to ensure PD for Cholesky
|
| 384 |
+
lam = 1e-10
|
| 385 |
+
g.flat[:: d + 1] += lam
|
| 386 |
|
| 387 |
+
cF = la.cho_factor(g, lower=True, check_finite=False)
|
| 388 |
|
| 389 |
+
# momentum from velocity: p = g v
|
| 390 |
+
p = g @ v
|
|
|
|
|
|
|
|
|
|
| 391 |
|
| 392 |
+
# S_m = v^T grad_k[m,:] = -(2c) k_m s_m
|
| 393 |
+
S_m = -(2.0 * c) * (k_m * s_m) # (k,)
|
|
|
|
| 394 |
|
| 395 |
+
# T_xm = sum_y v_y * d^2/dR_x dR_y k_m
|
| 396 |
+
# Hessian_{xy} = k_m * [ 4c^2 * (Dw_x/std_x)(Dw_y/std_y) - 2c * delta_xy / std_x^2 ]
|
| 397 |
+
# contraction over y gives:
|
| 398 |
+
# T_xm = k_m * [ 4c^2*(Dw_x/std_x)*s_m - 2c*(v_x/std_x^2) ]
|
| 399 |
+
term1 = (4.0 * c * c) * (k_m * s_m)[:, None] * Dw_over # (k,d)
|
| 400 |
+
term2 = (2.0 * c) * k_m[:, None] * (v[None, :] * inv_std2[None, :]) # (k,d)
|
| 401 |
+
T = (term1 - term2).T # (d,k)
|
| 402 |
|
| 403 |
+
# force: f = T @ (C @ S)
|
| 404 |
+
w = C @ S_m # (k,)
|
| 405 |
+
f = T @ w # (d,)
|
| 406 |
|
| 407 |
+
return g, cF, f, p
|
| 408 |
|
| 409 |
+
def _solve(self, cF: Tuple[np.ndarray, bool], b: np.ndarray) -> np.ndarray:
|
| 410 |
+
return la.cho_solve(cF, b, check_finite=False)
|
| 411 |
|
| 412 |
+
# -----------------------------
|
| 413 |
+
# Geodesic generalized leapfrog
|
| 414 |
+
# -----------------------------
|
| 415 |
+
def flow(
|
| 416 |
+
self,
|
| 417 |
+
R_ax: Union[np.ndarray, list],
|
| 418 |
+
v_ax: Union[np.ndarray, list],
|
| 419 |
+
*,
|
| 420 |
+
eps: float = 1e-2,
|
| 421 |
+
K_p: int = 5,
|
| 422 |
+
K_q: int = 5,
|
| 423 |
+
) -> Tuple[np.ndarray, np.ndarray]:
|
| 424 |
+
"""
|
| 425 |
+
One generalized-leapfrog step for geodesic flow on the pullback manifold.
|
| 426 |
|
| 427 |
+
Input:
|
| 428 |
+
R_ax: (A,d) or (d,)
|
| 429 |
+
v_ax: (A,d) or (d,) (velocity in *unwhitened* latent coordinates)
|
|
|
|
|
|
|
|
|
|
| 430 |
|
| 431 |
+
Output:
|
| 432 |
+
(R_next, v_next) with same shapes as input.
|
| 433 |
|
| 434 |
+
Notes:
|
| 435 |
+
- Uses fixed-point iterations with fixed K_p, K_q for determinism/time-reversibility.
|
| 436 |
+
- If pred_k is set, uses only the k nearest inducing points per query for geometry.
|
| 437 |
+
"""
|
| 438 |
+
R = np.asarray(R_ax, dtype=np.float64)
|
| 439 |
+
v = np.asarray(v_ax, dtype=np.float64)
|
| 440 |
|
| 441 |
+
single = (R.ndim == 1)
|
| 442 |
+
if single:
|
| 443 |
+
R = R[None, :]
|
| 444 |
+
v = v[None, :]
|
| 445 |
+
if R.shape != v.shape or R.ndim != 2 or R.shape[1] != self.d_lat:
|
| 446 |
+
raise ValueError(f"Expected R_ax and v_ax to have shape (A,{self.d_lat}) (or ({self.d_lat},)).")
|
| 447 |
|
| 448 |
+
A, d = R.shape
|
|
|
|
|
|
|
| 449 |
|
| 450 |
+
# inducing neighbor indices per point (optional)
|
| 451 |
+
Rw = self._whiten(R)
|
| 452 |
+
idx_aK = self._inducing_idx(Rw) # None or (A,k)
|
| 453 |
|
| 454 |
+
R_next = np.empty_like(R)
|
| 455 |
+
v_next = np.empty_like(v)
|
| 456 |
|
| 457 |
+
for a in range(A):
|
| 458 |
+
idx = None if idx_aK is None else idx_aK[a]
|
| 459 |
|
| 460 |
+
r_n = R[a]
|
| 461 |
+
v_n = v[a]
|
| 462 |
|
| 463 |
+
# metric/force at step start (for implicit p-half)
|
| 464 |
+
g_n, cF_n, f_n, p_n = self._rbf_terms_single(r_n, v_n, idx_m=idx)
|
| 465 |
|
| 466 |
+
# --- (1) implicit half-step in momentum via fixed-point on p ---
|
| 467 |
+
p = p_n.copy()
|
| 468 |
+
for _ in range(int(K_p)):
|
| 469 |
+
v_k = self._solve(cF_n, p) # v = g_n^{-1} p
|
| 470 |
+
_, _, f_k, _ = self._rbf_terms_single(r_n, v_k, idx_m=idx)
|
| 471 |
+
p = p_n + 0.5 * eps * f_k
|
| 472 |
+
p_half = p
|
| 473 |
|
| 474 |
+
# --- (2) implicit position update via fixed-point on r ---
|
| 475 |
+
v_half_n = self._solve(cF_n, p_half) # M(r_n) p_half
|
| 476 |
+
r = r_n + eps * v_half_n # init
|
|
|
|
|
|
|
| 477 |
|
| 478 |
+
for _ in range(int(K_q)):
|
| 479 |
+
g_r, cF_r, _, _ = self._rbf_terms_single(r, v_half_n, idx_m=idx) # metric at current r
|
| 480 |
+
v_half_r = self._solve(cF_r, p_half) # M(r) p_half
|
| 481 |
+
r = r_n + 0.5 * eps * (v_half_n + v_half_r)
|
| 482 |
|
| 483 |
+
r_np1 = r
|
|
|
|
|
|
|
|
|
|
|
|
|
| 484 |
|
| 485 |
+
# --- (3) explicit half-step in momentum at r_{n+1} ---
|
| 486 |
+
# recompute metric at r_{n+1} and compute force using v_mid = M(r_{n+1}) p_half
|
| 487 |
+
g_np1, cF_np1, _, _ = self._rbf_terms_single(r_np1, v_half_n, idx_m=idx)
|
| 488 |
+
v_mid = self._solve(cF_np1, p_half)
|
| 489 |
+
_, _, f_np1, _ = self._rbf_terms_single(r_np1, v_mid, idx_m=idx)
|
|
|
|
|
|
|
|
|
|
| 490 |
|
| 491 |
+
p_np1 = p_half + 0.5 * eps * f_np1
|
| 492 |
+
v_np1 = self._solve(cF_np1, p_np1)
|
| 493 |
|
| 494 |
+
R_next[a] = r_np1
|
| 495 |
+
v_next[a] = v_np1
|
| 496 |
|
| 497 |
+
if single:
|
| 498 |
+
return R_next[0], v_next[0]
|
| 499 |
+
return R_next, v_next
|
| 500 |
|
| 501 |
|
| 502 |
+
__all__ = ["GPLM", "InducingMode"]
|