jcandane commited on
Commit
94d2d5e
·
verified ·
1 Parent(s): 1f2bd6b

Update src/dima/gplm.py

Browse files
Files changed (1) hide show
  1. 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, overload
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
- def sqdist_ab(A: np.ndarray, B: np.ndarray) -> np.ndarray:
35
- """
36
- Pairwise squared Euclidean distances between rows:
37
- A: (a,d), B: (b,d) -> D2: (a,b)
38
  """
39
- A = np.asarray(A, dtype=np.float64)
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
- class GPLM:
48
- """
49
- GPLM decoder with optional ANN sparse inducing prediction, plus a geodesic "flow" integrator.
50
-
51
- Required stored arrays/params:
52
- - Z_mx_w : (m,d) inducing points in *whitened latent coords*
53
- - M_mX : (m,D) inducing outputs
54
- - mean_X : (D,) output mean added to decode
55
- - lat_mean_x : (d,) latent mean (for whitening)
56
- - lat_std_x : (d,) latent std (for whitening)
57
- - beta, eps : RBF hyperparams used exactly as in your _decode
58
- - pred_k : if not None and < m, use ANN sparse neighbors
59
- - ann_Z : object with method: search(query: (a,d), k:int) -> (idx: (a,k), D2: (a,k))
 
60
 
61
  Notes:
62
- - __call__ and _decode are kept behavior-identical to your pasted version.
63
- - flow() integrates geodesics of the pullback metric g = J^T J induced by this decoder.
 
64
  """
65
 
66
  def __init__(
67
  self,
 
 
68
  *,
69
- Z_mx_w: np.ndarray,
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
- self.Z_mx_w = np.asarray(Z_mx_w, dtype=np.float64) # (m,d) whitened inducing
82
- self.M_mX = np.asarray(M_mX, dtype=np.float64) # (m,D)
83
- self.mean_X = np.asarray(mean_X, dtype=np.float64) # (D,)
 
 
 
 
 
 
84
 
85
- self.m = int(self.M_mX.shape[0])
86
- self.d = int(self.Z_mx_w.shape[1])
87
- self.D = int(self.M_mX.shape[1])
 
 
88
 
89
- if lat_mean_x is None:
90
- lat_mean_x = np.zeros(self.d, dtype=np.float64)
91
- if lat_std_x is None:
92
- lat_std_x = np.ones(self.d, dtype=np.float64)
93
 
94
- self.lat_mean_x = np.asarray(lat_mean_x, dtype=np.float64) # (d,)
95
- self.lat_std_x = np.asarray(lat_std_x, dtype=np.float64) # (d,)
96
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
97
  self.beta = float(beta)
98
- self.eps = float(eps)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
99
 
100
- self.pred_k = pred_k
101
- self.ann_Z = ann_Z
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
102
 
103
- # matches your code: scale used in dW = scale * W * diff
104
- self._rbf_grad_scale = float(-2.0 * self.beta / self.eps)
 
 
 
105
 
106
- if self.pred_k is not None:
107
- if int(self.pred_k) <= 0:
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
- # Your original public API
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
- Returns:
126
- - if jacobian=False and metric=False: Y (a,D)
127
- - if jacobian=True, metric=False: (Y, J) where J is (a,D,d)
128
- - if jacobian=False, metric=True: (Y, g) where g is (a,d,d)
129
- - if jacobian=True, metric=True: (Y, J, g)
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
- out = self._decode(R_ax, jacobian=jacobian, metric=metric)
139
  else:
140
  bs = int(batch_size)
141
- Ys = []
142
- Js = [] if jacobian else None
143
- Gs = [] if metric else None
144
  for s in range(0, R_ax.shape[0], bs):
145
- chunk = self._decode(R_ax[s:s + bs], jacobian=jacobian, metric=metric)
146
- if (not jacobian) and (not metric):
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
- # unwrap single
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
- # We'll need this for chain rule in gradients:
198
- inv_std = (1.0 / self.lat_std_x).astype(np.float64) # (d,)
 
199
 
200
  if self.pred_k is None or self.pred_k == self.m:
201
- # full inducing
202
- D2_am = sqdist_ab(Za_w, self.Z_mx_w) # (a,m)
203
- W_am = np.exp(-self.beta * (D2_am.astype(np.float64) / self.eps)) # (a,m)
204
- Y = W_am @ self.M_mX # (a,D)
205
-
206
- if not (jacobian or metric):
207
- return Y + self.mean_X[None, :]
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
- if Z.shape != v.shape:
296
- raise ValueError(f"R_ax and v0_ax must have same shape; got {Z.shape} vs {v.shape}")
297
 
298
- # a_n = a(z_n, v_n)
299
- a0, g0 = self._geodesic_accel_analytic(Z, v, reg=reg)
 
 
 
 
300
 
301
- # v_{n+1/2}
302
- v_half = v + 0.5 * float(dt) * a0
 
 
 
 
 
 
 
303
 
304
- # z_{n+1}
305
- Q = Z + float(dt) * v_half
 
 
 
 
 
 
 
 
 
 
 
306
 
307
- # a_{n+1} = a(z_{n+1}, v_{n+1/2})
308
- a1, g1 = self._geodesic_accel_analytic(Q, v_half, reg=reg)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
309
 
310
- # v_{n+1}
311
- v_new = v_half + 0.5 * float(dt) * a1
 
312
 
313
- # Optional: keep Riemannian speed sqrt(v^T g v) constant per sample
314
- if keep_speed:
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
- if single:
320
- return Q[0], v_new[0]
321
- return Q, v_new
322
 
323
- def _geodesic_accel_analytic(
324
- self,
325
- R_ax: np.ndarray,
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
- # Precompute whitening + chain rule
347
- Za_w = (Z - self.lat_mean_x) / self.lat_std_x # (a,d)
348
- inv_std = (1.0 / self.lat_std_x).astype(np.float64) # (d,)
349
- alpha = float(self._rbf_grad_scale)
 
 
 
 
350
 
351
- I = np.eye(d, dtype=np.float64)
352
 
353
- if self.pred_k is None or self.pred_k == self.m:
354
- # ----- full inducing -----
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
- # ∂_y k_am = alpha * k_am * D_amy * (1/σ_y)
360
- dk_amd = (alpha * k_am[:, :, None]) * diff_amd # (a,m,d) (still whitened diffs)
361
- dk_amd *= inv_std[None, None, :] # chain to unwhitened
362
 
363
- # J_{aXy} = Σ_m (∂_y k_am) M_{mX}
364
- J_aXy = np.einsum("amd,mX->aXd", dk_amd, self.M_mX, optimize=True) # (a,D,d)
 
 
 
 
 
365
 
366
- # metric g_{ayz} = J_{aXy} J_{aXz}
367
- g_ayz = np.einsum("aXd,aXe->ade", J_aXy, J_aXy, optimize=True) # (a,d,d)
368
- g_ayz = 0.5 * (g_ayz + np.swapaxes(g_ayz, 1, 2)) # symmetrize
369
 
370
- g_inv = np.linalg.inv(g_ayz + reg * I[None, :, :]) # (a,d,d)
371
 
372
- # dg[a, w, y, z] = ∂_w g_{yz}
373
- dg_awyz = np.zeros((a, d, d, d), dtype=np.float64)
374
 
375
- # For each derivative index w, build Hessian slice H_{X w y} (as (a,D,y))
376
- for w in range(d):
377
- diff_w_am = diff_amd[:, :, w] # (a,m)
 
 
 
 
 
 
 
 
 
 
 
378
 
379
- # ∂_{w y} k = k * [ alpha^2 * D_w * D_y + alpha * δ_{wy} ] * (1/σ_w)(1/σ_y)
380
- d2k_amy = k_am[:, :, None] * (
381
- (alpha * alpha) * diff_w_am[:, :, None] * diff_amd
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
- # H_{aXy} for this w: H_{aXwy} = Σ_m (∂_{w y} k_am) M_{mX}
387
- H_w_aXy = np.einsum("amy,mX->aXd", d2k_amy, self.M_mX, optimize=True) # (a,D,d)
388
 
389
- # ∂_w g_{yz} = H_{X w y} J_{X z} + J_{X y} H_{X w z}
390
- dg_w = (
391
- np.einsum("aXd,aXe->ade", H_w_aXy, J_aXy, optimize=True)
392
- + np.einsum("aXd,aXe->ade", J_aXy, H_w_aXy, optimize=True)
393
- ) # (a,d,d)
394
- dg_awyz[:, w, :, :] = dg_w
395
 
396
- else:
397
- # ----- sparse ANN inducing neighbors -----
398
- j_aK, D2_aK = self.ann_Z.search(Za_w.astype(self.dtype, copy=False), int(self.pred_k)) # (a,k),(a,k)
399
- k_aK = np.exp(-self.beta * (D2_aK.astype(np.float64) / self.eps)) # (a,k)
 
 
400
 
401
- Zsel_aKd = self.Z_mx_w[j_aK] # (a,k,d)
402
- diff_aKd = Za_w[:, None, :] - Zsel_aKd # (a,k,d)
403
- M_aKX = self.M_mX[j_aK] # (a,k,D)
404
 
405
- dk_aKd = (alpha * k_aK[:, :, None]) * diff_aKd
406
- dk_aKd *= inv_std[None, None, :]
 
407
 
408
- # J_{aXy} = Σ_k (∂_y k_aK) M_{aKX}
409
- J_aXy = np.einsum("akd,akX->aXd", dk_aKd, M_aKX, optimize=True) # (a,D,d)
410
 
411
- g_ayz = np.einsum("aXd,aXe->ade", J_aXy, J_aXy, optimize=True) # (a,d,d)
412
- g_ayz = 0.5 * (g_ayz + np.swapaxes(g_ayz, 1, 2))
413
 
414
- g_inv = np.linalg.inv(g_ayz + reg * I[None, :, :]) # (a,d,d)
 
415
 
416
- dg_awyz = np.zeros((a, d, d, d), dtype=np.float64)
 
417
 
418
- for w in range(d):
419
- diff_w_aK = diff_aKd[:, :, w] # (a,k)
 
 
 
 
 
420
 
421
- d2k_aKy = k_aK[:, :, None] * (
422
- (alpha * alpha) * diff_w_aK[:, :, None] * diff_aKd
423
- + alpha * I[w][None, None, :]
424
- ) # (a,k,d)
425
- d2k_aKy *= (inv_std[w] * inv_std[None])[None, None, :]
426
 
427
- H_w_aXy = np.einsum("akd,akX->aXd", d2k_aKy, M_aKX, optimize=True) # (a,D,d)
 
 
 
428
 
429
- dg_w = (
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
- # Christoffels:
436
- # A[y,z,w] = ∂_y g_{zw} + ∂_z g_{yw} - ∂_w g_{yz}
437
- # with dg stored as dg[w,y,z] = ∂_w g_{yz}
438
- A_ayzw = (
439
- dg_awyz
440
- + dg_awyz.transpose(0, 2, 1, 3)
441
- - dg_awyz.transpose(0, 2, 3, 1)
442
- ) # (a,y,z,w)
443
 
444
- # Gamma[a,x,y,z] = 1/2 * g_inv[a,x,w] * A[a,y,z,w]
445
- Gamma_axyz = 0.5 * np.einsum("axw,ayzw->axyz", g_inv, A_ayzw, optimize=True)
446
 
447
- # acceleration: a^x = - Γ^x_{yz} v^y v^z
448
- a_ax = -np.einsum("axyz,ay,az->ax", Gamma_axyz, v, v, optimize=True)
449
 
450
- return a_ax, g_ayz
 
 
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"]