lrux 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
lrux/__init__.py ADDED
@@ -0,0 +1,3 @@
1
+ from .det_lru import det_lru, DetCarrier, init_det_carrier, det_lru_delayed
2
+ from .pfaffian import skew_eye, pf, slogpf
3
+ from .pf_lru import pf_lru, PfCarrier, init_pf_carrier, pf_lru_delayed
lrux/det_lru.py ADDED
@@ -0,0 +1,563 @@
1
+ from typing import Optional, Tuple, Union, Sequence, NamedTuple
2
+ from jax import Array
3
+ from jax.typing import ArrayLike
4
+ import jax
5
+ import jax.numpy as jnp
6
+ from jax._src.numpy import reductions
7
+
8
+
9
+ def _check_mat(mat: Array) -> None:
10
+ if mat.ndim != 2 or mat.shape[0] != mat.shape[1]:
11
+ raise ValueError(f"Expect input matrix shape (n, n), got {mat.shape}.")
12
+
13
+
14
+ def _standardize_uv(
15
+ u: Union[ArrayLike, Tuple[Array, ArrayLike]], n: int, dtype: jnp.dtype
16
+ ) -> Tuple[Array, Array]:
17
+ if isinstance(u, ArrayLike):
18
+ u = jnp.asarray(u)
19
+ if jnp.issubdtype(u.dtype, jnp.integer):
20
+ u = (jnp.empty((n, 0), dtype), u.flatten())
21
+ else:
22
+ u = (u.reshape(n, -1), jnp.array([], dtype=jnp.int32))
23
+ elif isinstance(u, Sequence):
24
+ u = (jnp.asarray(u[0]).reshape(n, -1), jnp.asarray(u[1]).flatten())
25
+ else:
26
+ raise ValueError(f"Got unsupported u or v data type {type(u)}.")
27
+ return u
28
+
29
+
30
+ def _check_uv(u: Tuple[Array, Array], v: Tuple[Array, Array]) -> None:
31
+ rank_u = u[0].shape[1] + u[1].size
32
+ rank_v = v[0].shape[1] + v[1].size
33
+ if rank_u != rank_v:
34
+ raise ValueError(
35
+ f"The input u and v should have matched rank, got {rank_u} and {rank_v}."
36
+ )
37
+
38
+
39
+ def _get_R(Ainv: Array, u: Tuple[Array, Array], v: Tuple[Array, Array]) -> Array:
40
+ xu_Ainv_xv = jnp.einsum("nk,nm,ml->kl", u[0], Ainv, v[0])
41
+ eu_Ainv_xv = Ainv[u[1]] @ v[0]
42
+ xu_Ainv_ev = u[0].T @ Ainv[:, v[1]]
43
+ eu_Ainv_ev = Ainv[u[1], :][:, v[1]]
44
+ uT_Ainv_v = jnp.block([[xu_Ainv_ev, xu_Ainv_xv], [eu_Ainv_ev, eu_Ainv_xv]])
45
+ return uT_Ainv_v.at[jnp.diag_indices_from(uT_Ainv_v)].add(1)
46
+
47
+
48
+ def _det_and_lufac(R: Array) -> Tuple[Array, Tuple[Array, Array]]:
49
+ lu, pivot = jax.scipy.linalg.lu_factor(R)
50
+ iota = jnp.arange(pivot.size, dtype=pivot.dtype)
51
+ parity = reductions.count_nonzero(pivot != iota, axis=-1)
52
+ sign = jnp.array(-2 * (parity % 2) + 1, dtype=lu.dtype)
53
+ det = sign * jnp.prod(jnp.diag(lu))
54
+ return det, (lu, pivot)
55
+
56
+
57
+ def _update_Ainv(
58
+ Ainv: Array,
59
+ u: Tuple[Array, Array],
60
+ v: Tuple[Array, Array],
61
+ lu_and_piv: Tuple[Array, Array],
62
+ ) -> Array:
63
+ uT_Ainv = jnp.concatenate((u[0].T @ Ainv, Ainv[u[1], :]), axis=0)
64
+ Rinv_uT_Ainv = jax.scipy.linalg.lu_solve(lu_and_piv, uT_Ainv)
65
+ Ainv_v = jnp.concatenate((Ainv[:, v[1]], Ainv @ v[0]), axis=1)
66
+ return Ainv - Ainv_v @ Rinv_uT_Ainv
67
+
68
+
69
+ def det_lru(
70
+ Ainv: Array,
71
+ u: Union[ArrayLike, Tuple[Array, ArrayLike]],
72
+ v: Union[ArrayLike, Tuple[Array, ArrayLike]],
73
+ return_update: bool = False,
74
+ ) -> Union[Array, Tuple[Array, Array]]:
75
+ r"""
76
+ Low-rank update of determinant :math:`\det(A_1) = \det(A_0 + vu^T)`
77
+
78
+ :param Ainv:
79
+ Inverse of the original matrix :math:`A_0^{-1}`, shape (n, n)
80
+
81
+ :param u:
82
+ Low-rank update vector(s) :math:`u`. There are several acceptable inputs
83
+ of ``u`` as listed below.
84
+
85
+ An array with shape (n,) or (n, k):
86
+ Direct expression of full low-rank vector(s).
87
+
88
+ An integer or array of integers with size k:
89
+ One-hot vectors ``u_full = jnp.zeros((n, k)).at[u, jnp.arange(k)].set(1)``.
90
+ For example, when you need a full matrix
91
+
92
+ .. code-block:: python
93
+
94
+ u = jnp.array([
95
+ [0, 0],
96
+ [1, 0],
97
+ [0, 0],
98
+ [0, 1],
99
+ ])
100
+
101
+ you can alternatively specify
102
+
103
+ .. code-block:: python
104
+
105
+ u = jnp.array([1, 3])
106
+
107
+ A tuple of two arrays, with respective shapes (n, k0) and (k1,):
108
+ A concatenation of the previous two. For example, when you need a full matrix
109
+
110
+ .. code-block:: python
111
+
112
+ u = jnp.array([
113
+ [u00, u01, 0, 0],
114
+ [u10, u11, 1, 0],
115
+ [u20, u21, 0, 0],
116
+ [u30, u31, 0, 1],
117
+ ])
118
+
119
+ you can alternatively specify
120
+
121
+ .. code-block:: python
122
+
123
+ x = jnp.array([
124
+ [u00, u01],
125
+ [u10, u11],
126
+ [u20, u21],
127
+ [u30, u31],
128
+ ])
129
+ e = jnp.array([1, 3])
130
+ u = (x, e)
131
+
132
+ The matrix product of one-hot vectors is internally performed by matrix slicing
133
+ for better performance, so an input of indices is preferred.
134
+
135
+ :param v:
136
+ Low-rank update vector(s) :math:`v`. The acceptable inputs are similar to ``u``.
137
+ When the input is a tuple of two arrays ``v = (x, e)``, for convenience
138
+ the concatenation order is reversely given by ``jnp.concatenate((ve, vx), axis=1)``.
139
+ Therefore, when you need
140
+
141
+ .. code-block:: python
142
+
143
+ v = jnp.array([
144
+ [0, 0, v00, v01],
145
+ [1, 0, v10, v11],
146
+ [0, 0, v20, u21],
147
+ [0, 1, v30, v31],
148
+ ])
149
+
150
+ you can alternatively specify
151
+
152
+ .. code-block:: python
153
+
154
+ x = jnp.array([
155
+ [v00, v01],
156
+ [v10, v11],
157
+ [v20, v21],
158
+ [v30, v31],
159
+ ])
160
+ e = jnp.array([1, 3])
161
+ u = (x, e)
162
+
163
+ :param return_update:
164
+ Whether the new matrix inverse :math:`A_1^{-1}` should be returned,
165
+ defaul to False.
166
+
167
+ :return:
168
+ ratio:
169
+ The ratio between two determinants
170
+
171
+ .. math::
172
+
173
+ r = \frac{\det(A_1)}{\det(A_0)} = \det(R)
174
+
175
+ where
176
+
177
+ .. math::
178
+
179
+ R = I + u^T A_0^{-1} v
180
+
181
+ new_Ainv:
182
+ The new matrix inverse
183
+
184
+ .. math::
185
+
186
+ A_1^{-1} = (A_0 + vu^T)^{-1} = A_0^{-1} - A_0^{-1} v R^{-1} u^T A_0^{-1}
187
+
188
+ Only returned when ``return_update`` is True.
189
+
190
+ .. tip::
191
+
192
+ This function is compatible with ``jax.jit`` and ``jax.vmap``, while
193
+ ``return_update`` is a static argument which shouldn't be jitted or vmapped.
194
+
195
+ Furthermore, we recommend setting ``donate_argnums=0`` in ``jax.jit`` to reuse
196
+ the memory of ``Ainv`` if it's no longer needed. This helps to greatly reduce
197
+ the time and memory cost. For instance,
198
+
199
+ .. code-block:: python
200
+
201
+ lru_vmap = jax.vmap(det_lru, in_axes=(0, 0, 0, None))
202
+ lru_jit = jax.jit(lru_vmap, static_argnums=3, donate_argnums=0)
203
+
204
+ .. note::
205
+
206
+ Here are examples of how to define ``u`` and ``v`` before calling ``det_lru(Ainv, u, v)``.
207
+ Keep in mind that the low-rank update we need takes the form
208
+
209
+ .. math::
210
+
211
+ A_1 - A_0 = vu^T
212
+
213
+ **Rank-1 row update**
214
+
215
+ .. math::
216
+
217
+ A_1 - A_0 = \begin{pmatrix}
218
+ 0 & 0 & 0 & 0 \\
219
+ u_0 & u_1 & u_2 & u_3 \\
220
+ 0 & 0 & 0 & 0 \\
221
+ 0 & 0 & 0 & 0 \\
222
+ \end{pmatrix}
223
+ = \begin{pmatrix}
224
+ 0 \\ 1 \\ 0 \\ 0
225
+ \end{pmatrix}
226
+ (u_0, u_1, u_2, u_3)
227
+
228
+ .. code-block:: python
229
+
230
+ u = jnp.array([u0, u1, u2, u3])
231
+ v = 1
232
+
233
+ **Rank-1 column update**
234
+
235
+ .. math::
236
+
237
+ A_1 - A_0 = \begin{pmatrix}
238
+ 0 & 0 & v_0 & 0 \\
239
+ 0 & 0 & v_1 & 0 \\
240
+ 0 & 0 & v_2 & 0 \\
241
+ 0 & 0 & v_3 & 0 \\
242
+ \end{pmatrix}
243
+ = \begin{pmatrix}
244
+ v_0 \\ v_1 \\ v_2 \\ v_3
245
+ \end{pmatrix}
246
+ (0, 0, 1, 0)
247
+
248
+ .. code-block:: python
249
+
250
+ u = 2
251
+ v = jnp.array([v0, v1, v2, v3])
252
+
253
+ **Rank-2 row update**
254
+
255
+ .. math::
256
+
257
+ A_1 - A_0 = \begin{pmatrix}
258
+ 0 & 0 & 0 & 0 \\
259
+ u_{00} & u_{01} & u_{02} & u_{03} \\
260
+ 0 & 0 & 0 & 0 \\
261
+ u_{10} & u_{11} & u_{12} & u_{13} \\
262
+ \end{pmatrix}
263
+ = \begin{pmatrix}
264
+ 0 & 0 \\ 1 & 0 \\ 0 & 0 \\ 0 & 1
265
+ \end{pmatrix}
266
+ \begin{pmatrix}
267
+ u_{00} & u_{01} & u_{02} & u_{03} \\
268
+ u_{10} & u_{11} & u_{12} & u_{13} \\
269
+ \end{pmatrix}
270
+
271
+ .. code-block:: python
272
+
273
+ u = jnp.array([[u00, u10], [u01, u11], [u02, u12], [u03, u13]])
274
+ v = jnp.array([1, 3])
275
+
276
+ **Simultaneous update of row and column**
277
+
278
+ .. math::
279
+
280
+ A_1 - A_0 = \begin{pmatrix}
281
+ 0 & 0 & v_0 & 0 \\
282
+ u_0 & u_1 & u_2 + v_1 & u_3 \\
283
+ 0 & 0 & v_2 & 0 \\
284
+ 0 & 0 & v_3 & 0 \\
285
+ \end{pmatrix}
286
+ = \begin{pmatrix}
287
+ 0 & v_0 \\ 1 & v_1 \\ 0 & v_2 \\ 0 & v_3
288
+ \end{pmatrix}
289
+ \begin{pmatrix}
290
+ u_0 & u_1 & u_2 & u_3 \\
291
+ 0 & 0 & 1 & 0 \\
292
+ \end{pmatrix}
293
+
294
+ .. code-block:: python
295
+
296
+ u = (jnp.array([u0, u1, u2, u3]), 2)
297
+ v = (jnp.array([v0, v1, v2, v3]), 1)
298
+ """
299
+ _check_mat(Ainv)
300
+ u = _standardize_uv(u, Ainv.shape[0], Ainv.dtype)
301
+ v = _standardize_uv(v, Ainv.shape[0], Ainv.dtype)
302
+ _check_uv(u, v)
303
+
304
+ R = _get_R(Ainv, u, v)
305
+ ratio, lufac = _det_and_lufac(R)
306
+ if return_update:
307
+ Ainv = _update_Ainv(Ainv, u, v, lufac)
308
+ return ratio, Ainv
309
+ else:
310
+ return ratio
311
+
312
+
313
+ class DetCarrier(NamedTuple):
314
+ Ainv: Array
315
+ a: Array
316
+ b: Array
317
+
318
+
319
+ def init_det_carrier(A: Array, max_delay: int, max_rank: int = 1) -> DetCarrier:
320
+ r"""
321
+ Prepare the data and space for `~lrux.det_lru_delayed`
322
+
323
+ :param A:
324
+ The initial matrix :math:`A_0` with shape (n, n).
325
+
326
+ :param max_delay:
327
+ The maximum iterations T of delayed updates, usually chosen to be ~n/10.
328
+
329
+ :param max_rank:
330
+ The maximum rank K in delayed updates, default to 1.
331
+
332
+ :return:
333
+ A ``NamedTuple`` with the following attributes.
334
+
335
+ Ainv:
336
+ The initial matrix inverse :math:`A_0^{-1}` of shape (n, n).
337
+ a:
338
+ The delayed update vectors of shape (T, n, K), initialized to 0
339
+ b:
340
+ The delayed update vectors of shape (T, n, K), initialized to 0
341
+ """
342
+
343
+ if max_delay <= 0:
344
+ raise ValueError(
345
+ "`max_delay` should be a positive integer. "
346
+ "Otherwise, please use `det_lru` for non-delayed updates."
347
+ )
348
+ _check_mat(A)
349
+ Ainv = jnp.linalg.inv(A)
350
+ a = jnp.zeros((max_delay, A.shape[0], max_rank), A.dtype)
351
+ b = jnp.zeros_like(a)
352
+ return DetCarrier(Ainv, a, b)
353
+
354
+
355
+ def _update_ab(a: Array, new_a: Array, current_delay: int) -> Array:
356
+ k = new_a.shape[-1]
357
+ if k > a.shape[-1]:
358
+ raise ValueError(
359
+ "The rank of update exceeds max_rank specified in `init_det_carrier`."
360
+ )
361
+ return a.at[current_delay, :, :k].set(new_a)
362
+
363
+
364
+ def _get_delayed_output(
365
+ carrier: DetCarrier,
366
+ u: Tuple[Array, Array],
367
+ v: Tuple[Array, Array],
368
+ return_update: bool,
369
+ current_delay: int,
370
+ ) -> Union[Array, Tuple[Array, Array]]:
371
+ Ainv = carrier.Ainv
372
+ a = carrier.a[:current_delay]
373
+ b = carrier.b[:current_delay]
374
+ R0 = _get_R(Ainv, u, v)
375
+
376
+ xuT_a = jnp.einsum("nk,tnl->tkl", u[0], a)
377
+ euT_a = a[:, u[1], :]
378
+ uT_a = jnp.concatenate((xuT_a, euT_a), axis=1)
379
+
380
+ xvT_b = jnp.einsum("nk,tnl->tkl", v[0], b)
381
+ evT_b = b[:, v[1], :]
382
+ vT_b = jnp.concatenate((evT_b, xvT_b), axis=1)
383
+
384
+ R = R0 - jnp.einsum("tkl,tml->km", uT_a, vT_b)
385
+ ratio, lufac = _det_and_lufac(R)
386
+
387
+ if return_update:
388
+ a0 = jnp.concatenate((Ainv[:, v[1]], Ainv @ v[0]), axis=1)
389
+ new_a = a0 - jnp.einsum("tnk,tlk->nl", a, vT_b)
390
+ bT0 = jnp.concatenate((u[0].T @ Ainv, Ainv[u[1], :]), axis=0)
391
+ new_bT = bT0 - jnp.einsum("tkl,tnl->kn", uT_a, b)
392
+ new_bT = jax.scipy.linalg.lu_solve(lufac, new_bT)
393
+
394
+ a = _update_ab(carrier.a, new_a, current_delay)
395
+ b = _update_ab(carrier.b, new_bT.T, current_delay)
396
+
397
+ if current_delay == a.shape[0] - 1:
398
+ Ainv -= jnp.einsum("tnk,tmk->nm", a, b)
399
+ carrier = DetCarrier(Ainv, jnp.zeros_like(a), jnp.zeros_like(b))
400
+ else:
401
+ carrier = DetCarrier(Ainv, a, b)
402
+ return ratio, carrier
403
+ else:
404
+ return ratio
405
+
406
+
407
+ def det_lru_delayed(
408
+ carrier: DetCarrier,
409
+ u: Union[ArrayLike, Tuple[Array, ArrayLike]],
410
+ v: Union[ArrayLike, Tuple[Array, ArrayLike]],
411
+ return_update: bool = False,
412
+ current_delay: Optional[int] = None,
413
+ ) -> Union[Array, Tuple[Array, DetCarrier]]:
414
+ r"""
415
+ Delayed low-rank update of determinant
416
+
417
+ :param carrier:
418
+ The existing delayed update quantities, including :math:`A_0^{-1}`, and
419
+
420
+ .. math::
421
+
422
+ a_t = A_{t-1}^{-1} v_t
423
+
424
+ .. math::
425
+
426
+ b_t = (A_{t-1}^{-1})^T u_t
427
+
428
+ with :math:`t` from 1 to :math:`\tau-1`.
429
+ Initially provided by `~lrux.init_det_carrier`.
430
+
431
+ :param u:
432
+ Low-rank update vector(s) :math:`u_\tau`, the same as :math:`u` in `lrux.det_lru`.
433
+ The rank of u shouldn't exceed the maximum allowed rank specified
434
+ in `~lrux.init_det_carrier`.
435
+
436
+ :param v:
437
+ Low-rank update vector(s) :math:`v_\tau`, the same as :math:`v` in `lrux.det_lru`.
438
+ The rank of v shouldn't exceed the maximum allowed rank specified
439
+ in `~lrux.init_det_carrier`.
440
+
441
+ :param return_update:
442
+ Whether the new carrier with updated quantities should be returned,
443
+ defaul to False.
444
+
445
+ :param current_delay:
446
+ The current iterations :math:`\tau` of delayed updates. As python starts counting
447
+ from 0, the actual :math:`\tau` should be ``current_delay + 1``.
448
+ It must be specified when ``return_update`` is True.
449
+
450
+ :return:
451
+ ratio:
452
+ The ratio between two determinants
453
+
454
+ .. math::
455
+
456
+ r_\tau = \frac{\det(A_\tau)}{\det(A_{\tau-1})} = \det(R_\tau)
457
+
458
+ where
459
+
460
+ .. math::
461
+
462
+ R_\tau = I + u_\tau^T A_0^{-1} v_\tau - \sum_{t=1}^{\tau-1} (u_\tau^T a_t) (b_t^T v_\tau)
463
+
464
+ new_carrier:
465
+ Only returned when ``return_update`` is True. The new carrier contains
466
+ the quantities from the input carrier, and in addition
467
+
468
+ .. math::
469
+
470
+ a_\tau = A_{\tau-1}^{-1} v_\tau = A_0^{-1} v_\tau - \sum_{t=1}^{\tau-1} a_t (b_t^T v_\tau)
471
+
472
+ .. math::
473
+
474
+ b_\tau = (A_{\tau-1}^{-1})^T u_\tau = (A_0^{-1})^T u_\tau - \sum_{t=1}^{\tau-1} b_t (a_t^T u_\tau)
475
+
476
+ When :math:`\tau` reaches the maximum delayed iterations :math:`T`
477
+ specified in `~lrux.init_det_carrier`, i.e. ``current_delay == max_delay - 1``,
478
+ the current :math:`A_\tau` will be set as the new :math:`A_0`,
479
+ whose inverse is given by
480
+
481
+ .. math::
482
+
483
+ A_\tau^{-1} = A_0^{-1} - \sum_{t=1}^\tau a_t b_t^T
484
+
485
+ The ``Ainv`` in ``new_carrier`` will be replaced by :math:`A_\tau^{-1}`, and
486
+ ``a`` and ``b`` will be set to 0.
487
+
488
+ .. warning::
489
+
490
+ This function is only recommended for heavy users who understand why and when
491
+ to use delayed updates. Otherwise, please choose `~\lrux.det_lru`.
492
+
493
+ .. tip::
494
+
495
+ Similar to `~lrux.det_lru`, this function is compatible with ``jax.jit`` and
496
+ ``jax.vmap``, while ``return_update`` and ``current_delay`` are static arguments
497
+ which shouldn't be jitted or vmapped.
498
+
499
+ We still recommend setting ``donate_argnums=0`` in ``jax.jit`` to reuse
500
+ the memory of ``carrier`` if it's no longer needed. For instance,
501
+
502
+ .. code-block:: python
503
+
504
+ lru_vmap = jax.vmap(det_lru_delayed, in_axes=(0, 0, 0, None, None))
505
+ lru_jit = jax.jit(lru_vmap, static_argnums=(3, 4), donate_argnums=0)
506
+
507
+ Here is a complete example of delayed updates.
508
+
509
+ .. code-block:: python
510
+
511
+ import os
512
+ os.environ["JAX_ENABLE_X64"] = "1"
513
+
514
+ import random
515
+ import jax
516
+ import jax.numpy as jnp
517
+ import jax.random as jr
518
+ from lrux import det_lru_delayed, init_det_carrier
519
+
520
+ def _get_key():
521
+ seed = random.randint(0, 2**31 - 1)
522
+ return jr.key(seed)
523
+
524
+ dtype = jnp.float64
525
+ n = 10
526
+ max_delay = n // 2
527
+ max_rank = 2
528
+ A = jr.normal(_get_key(), (n, n), dtype)
529
+ carrier = init_det_carrier(A, max_delay, max_rank)
530
+ detA0 = jnp.linalg.det(A)
531
+
532
+ lru_fn = jax.jit(det_lru_delayed, static_argnums=(3, 4), donate_argnums=0)
533
+
534
+ for i in range(20):
535
+ current_delay = i % max_delay
536
+ k = random.randint(0, max_rank)
537
+ u = jr.normal(_get_key(), (n, k), dtype)
538
+ v = jr.normal(_get_key(), (n, k), dtype)
539
+ ratio, carrier = lru_fn(carrier, u, v, True, current_delay)
540
+
541
+ # verify the low-rank update result
542
+ A += v @ u.T
543
+ detA1 = jnp.linalg.det(A)
544
+ assert jnp.allclose(ratio, detA1 / detA0)
545
+ detA0 = detA1
546
+ """
547
+ max_delay = carrier.a.shape[0]
548
+ if current_delay is None:
549
+ if return_update:
550
+ raise ValueError("`current_delay` must be specified to return updates.")
551
+ current_delay = max_delay - 1
552
+
553
+ elif current_delay < 0 or current_delay >= max_delay:
554
+ raise ValueError(
555
+ f"`current_delay` should be in range [0, {max_delay}), got {current_delay}."
556
+ )
557
+
558
+ Ainv = carrier.Ainv
559
+ u = _standardize_uv(u, Ainv.shape[0], Ainv.dtype)
560
+ v = _standardize_uv(v, Ainv.shape[0], Ainv.dtype)
561
+ _check_uv(u, v)
562
+
563
+ return _get_delayed_output(carrier, u, v, return_update, current_delay)
lrux/pf_lru.py ADDED
@@ -0,0 +1,429 @@
1
+ from typing import Optional, Tuple, Union, NamedTuple
2
+ from jax.typing import ArrayLike
3
+ from jax import Array
4
+ import jax
5
+ import jax.numpy as jnp
6
+ from .det_lru import _standardize_uv, _update_ab
7
+ from .pfaffian import skew_eye, pf
8
+
9
+
10
+ def _check_mat(mat: Array) -> Array:
11
+ if mat.ndim != 2 or mat.shape[0] != mat.shape[1] or mat.shape[0] % 2 == 1:
12
+ raise ValueError(f"Expect input matrix shape (2n, 2n), got {mat.shape}.")
13
+ return (mat - mat.T) / 2
14
+
15
+
16
+ def _get_R(Ainv: Array, u: Tuple[Array, Array]) -> Array:
17
+ xu_Ainv_xu = jnp.einsum("nk,nm,ml->kl", u[0], Ainv, u[0])
18
+ xu_Ainv_eu = u[0].T @ Ainv[:, u[1]]
19
+ eu_Ainv_eu = Ainv[u[1], :][:, u[1]]
20
+ uT_Ainv_u = jnp.block([[xu_Ainv_xu, xu_Ainv_eu], [-xu_Ainv_eu.T, eu_Ainv_eu]])
21
+ J = skew_eye(uT_Ainv_u.shape[0] // 2, Ainv.dtype)
22
+ R = uT_Ainv_u + J
23
+ return (R - R.T) / 2 # ensure skew-symmetric
24
+
25
+
26
+ def _update_Ainv(Ainv: Array, u: Tuple[Array, Array], R: Array) -> Array:
27
+ Ainv_u = jnp.concatenate((Ainv @ u[0], Ainv[:, u[1]]), axis=1)
28
+ if R.shape[0] == 2:
29
+ Ainv_u1, Ainv_u2 = Ainv_u.T
30
+ outer = jnp.outer(Ainv_u1, Ainv_u2)
31
+ Ainv -= (outer - outer.T) / R[0, 1]
32
+ else:
33
+ Rinv_Ainv_u = jax.scipy.linalg.solve(R, Ainv_u.T)
34
+ Ainv += Ainv_u @ Rinv_Ainv_u
35
+ return (Ainv - Ainv.T) / 2
36
+
37
+
38
+ def pf_lru(
39
+ Ainv: Array,
40
+ u: Union[ArrayLike, Tuple[Array, ArrayLike]],
41
+ return_update: bool = False,
42
+ ) -> Union[Array, Tuple[Array, Array]]:
43
+ r"""
44
+ Low-rank update of pfaffian :math:`\mathrm{pf}(A_1) = \mathrm{pf}(A_0 - u J u^T)`.
45
+ Here :math:`J` is the skew-symmetric identity matrix
46
+
47
+ .. math::
48
+
49
+ J = \begin{pmatrix}
50
+ 0 & I \\ -I & 0
51
+ \end{pmatrix}
52
+
53
+ as given in `~lrux.skew_eye`.
54
+
55
+ :param Ainv:
56
+ Inverse of the original skew-symmetric matrix :math:`A_0^{-1}`, shape (n, n)
57
+
58
+ :param u:
59
+ Low-rank update vector(s) :math:`u`, the same as :math:`u` in `lrux.det_lru`.
60
+
61
+ :param return_update:
62
+ Whether the new matrix inverse :math:`A_1^{-1}` should be returned,
63
+ defaul to False.
64
+
65
+ :return:
66
+ ratio:
67
+ The ratio between two pfaffians
68
+
69
+ .. math::
70
+
71
+ r = \frac{\mathrm{pf}(A_1)}{\mathrm{pf}(A_0)} = \frac{\mathrm{pf}(R)}{\mathrm{pf}(J)}
72
+
73
+ where
74
+
75
+ .. math::
76
+
77
+ R = J + u^T A_0^{-1} u
78
+
79
+ new_Ainv:
80
+ The new matrix inverse
81
+
82
+ .. math::
83
+
84
+ A_1^{-1} = (A_0 + u J u^T)^{-1} = A_0^{-1} + (A_0^{-1} u) R^{-1} (A_0^{-1} u)^T
85
+
86
+ Only returned when ``return_update`` is True.
87
+
88
+ .. tip::
89
+
90
+ This function is compatible with ``jax.jit`` and ``jax.vmap``, while
91
+ ``return_update`` is a static argument which shouldn't be jitted or vmapped.
92
+
93
+ Furthermore, we recommend setting ``donate_argnums=0`` in ``jax.jit`` to reuse
94
+ the memory of ``Ainv`` if it's no longer needed. This helps to greatly reduce
95
+ the time and memory cost. For instance,
96
+
97
+ .. code-block:: python
98
+
99
+ lru_vmap = jax.vmap(pf_lru, in_axes=(0, 0, None))
100
+ lru_jit = jax.jit(lru_vmap, static_argnums=2, donate_argnums=0)
101
+
102
+ .. note::
103
+
104
+ Here are examples of how to define ``u`` before calling ``pf_lru(Ainv, u)``.
105
+ Keep in mind that the low-rank update we need should be skew-symmetric and takes
106
+ the form
107
+
108
+ .. math::
109
+
110
+ A_1 - A_0 = -u J u^T
111
+
112
+ **Update of 1 row and 1 column**
113
+
114
+ .. math::
115
+
116
+ A_1 - A_0 = \begin{pmatrix}
117
+ 0 & -u_0 & 0 & 0 \\
118
+ u_0 & 0 & u_2 & u_3 \\
119
+ 0 & -u_2 & 0 & 0 \\
120
+ 0 & -u_3 & 0 & 0 \\
121
+ \end{pmatrix}
122
+ = -\begin{pmatrix}
123
+ u_0 & 0 \\ u_1 & 1 \\ u_2 & 0 \\ u_3 & 0
124
+ \end{pmatrix}
125
+ \begin{pmatrix}
126
+ 0 & 1 \\ -1 & 0
127
+ \end{pmatrix}
128
+ \begin{pmatrix}
129
+ u_0 & u_1 & u_2 & u_3 \\
130
+ 0 & 1 & 0 & 0\\
131
+ \end{pmatrix}
132
+
133
+ .. code-block:: python
134
+
135
+ u = (jnp.array([u0, u1, u2, u3]), 1)
136
+
137
+ **Update of 2 rows and 2 columns**
138
+
139
+ .. math::
140
+
141
+ \begin{split}
142
+ A_1 - A_0 &= \begin{pmatrix}
143
+ 0 & -u_{00} & 0 & -u_{10} \\
144
+ u_{00} & 0 & u_{02} & u_{03} - u_{11} \\
145
+ 0 & -u_{02} & 0 & -u_{12} \\
146
+ u_{10} & u_{11} - u_{03} & u_{12} & 0 \\
147
+ \end{pmatrix} \\
148
+ &= -\begin{pmatrix}
149
+ u_{00} & u_{10} & 0 & 0 \\
150
+ u_{01} & u_{11} & 1 & 0 \\
151
+ u_{02} & u_{12} & 0 & 0 \\
152
+ u_{03} & u_{13} & 0 & 1 \\
153
+ \end{pmatrix}
154
+ \begin{pmatrix}
155
+ 0 & 0 & 1 & 0 \\
156
+ 0 & 0 & 0 & 1 \\
157
+ -1 & 0 & 0 & 0 \\
158
+ 0 & -1 & 0 & 0 \\
159
+ \end{pmatrix}
160
+ \begin{pmatrix}
161
+ u_{00} & u_{01} & u_{02} & u_{03} \\
162
+ u_{10} & u_{11} & u_{12} & u_{13} \\
163
+ 0 & 1 & 0 & 0 \\
164
+ 0 & 0 & 0 & 1 \\
165
+ \end{pmatrix}
166
+ \end{split}
167
+
168
+ .. code-block:: python
169
+
170
+ x = jnp.array([[u00, u10], [u01, u11], [u02, u12], [u03, u13]])
171
+ e = jnp.array([1, 3])
172
+ u = (x, e)
173
+ """
174
+ Ainv = _check_mat(Ainv)
175
+ u = _standardize_uv(u, Ainv.shape[0], Ainv.dtype)
176
+ k = u[0].shape[1] + u[1].size
177
+ if k % 2 == 1:
178
+ raise ValueError(f"The input u should have even rank, got rank {k}.")
179
+
180
+ R = _get_R(Ainv, u)
181
+ pfR = pf(R)
182
+ k_half = k // 2
183
+ ratio = jnp.where((k_half * (k_half - 1) // 2) % 2 == 0, pfR, -pfR)
184
+ if return_update:
185
+ Ainv = _update_Ainv(Ainv, u, R)
186
+ return ratio, Ainv
187
+ else:
188
+ return ratio
189
+
190
+
191
+ class PfCarrier(NamedTuple):
192
+ Ainv: Array
193
+ a: Array
194
+ Rinv: Array
195
+
196
+
197
+ def init_pf_carrier(A: Array, max_delay: int, max_rank: int = 2) -> PfCarrier:
198
+ r"""
199
+ Prepare the data and space for `~lrux.pf_lru_delayed`
200
+
201
+ :param A:
202
+ The initial skew-symmetric matrix :math:`A_0` with shape (n, n).
203
+
204
+ :param max_delay:
205
+ The maximum iterations T of delayed updates, usually chosen to be ~n/10.
206
+
207
+ :param max_rank:
208
+ The maximum rank K in delayed updates, default to 2.
209
+
210
+ :return:
211
+ A ``NamedTuple`` with the following attributes.
212
+
213
+ Ainv:
214
+ The initial matrix inverse :math:`A_0^{-1}` of shape (n, n).
215
+ a:
216
+ The delayed update vectors of shape (T, n, K), initialized to 0
217
+ Rinv:
218
+ The delayed update matrices :math:`R_t^{-1}` of shape (T, K, K), initialized to 0
219
+ """
220
+ if max_delay <= 0:
221
+ raise ValueError(
222
+ "`max_delay` should be a positive integer. "
223
+ "Otherwise, please use `pf_lru` for non-delayed updates."
224
+ )
225
+ A = _check_mat(A)
226
+ Ainv = jnp.linalg.inv(A)
227
+ Ainv = (Ainv - Ainv.T) / 2 # ensure skew-symmetric
228
+ a = jnp.zeros((max_delay, A.shape[0], max_rank), A.dtype)
229
+ Rinv = jnp.zeros((max_delay, max_rank, max_rank), A.dtype)
230
+ return PfCarrier(Ainv, a, Rinv)
231
+
232
+
233
+ def _get_delayed_updates(Ainv: Array, a: Array, Rinv: Array) -> Array:
234
+ if Rinv.shape[-1] == 2:
235
+ a1 = a[:, :, 0]
236
+ a2 = a[:, :, 1]
237
+ outer = jnp.einsum("tn,t,tm->nm", a1, Rinv[:, 0, 1], a2)
238
+ update = outer - outer.T
239
+ else:
240
+ update = jnp.einsum("tnj,tjk,tmk->nm", a, Rinv, a)
241
+
242
+ Ainv += update
243
+ return (Ainv - Ainv.T) / 2 # ensure skew-symmetric
244
+
245
+
246
+ def _get_delayed_output(
247
+ carrier: PfCarrier, u: Tuple[Array, Array], return_update: bool, current_delay: int
248
+ ) -> Union[Array, Tuple[Array, Array]]:
249
+ Ainv = carrier.Ainv
250
+ a = carrier.a[:current_delay]
251
+ Rinv = carrier.Rinv[:current_delay]
252
+ R0 = _get_R(Ainv, u)
253
+
254
+ xT_a = jnp.einsum("nk,tnl->tkl", u[0], a)
255
+ eT_a = a[:, u[1], :]
256
+ uT_a = jnp.concatenate((xT_a, eT_a), axis=1)
257
+
258
+ R = R0 + jnp.einsum("tjk,tkl,tml->jm", uT_a, Rinv, uT_a)
259
+ pfR = pf(R)
260
+ k = u[0].shape[1] + u[1].size
261
+ if k % 2 == 1:
262
+ raise ValueError(f"The input u should have even rank, got rank {k}.")
263
+ k_half = k // 2
264
+ ratio = jnp.where((k_half * (k_half - 1) // 2) % 2 == 0, pfR, -pfR)
265
+
266
+ if return_update:
267
+ a0 = jnp.concatenate((Ainv @ u[0], Ainv[:, u[1]]), axis=1)
268
+ new_a = a0 + jnp.einsum("tnj,tjk,tlk->nl", a, Rinv, uT_a)
269
+ a = _update_ab(carrier.a, new_a, current_delay)
270
+
271
+ if k == 2:
272
+ rinv = -1 / ratio
273
+ new_Rinv = jnp.array([[0, rinv], [-rinv, 0]], dtype=Rinv.dtype)
274
+ else:
275
+ new_Rinv = jnp.linalg.inv(R)
276
+ new_Rinv = (new_Rinv - new_Rinv.T) / 2 # ensure skew-symmetric
277
+ Rinv = carrier.Rinv.at[current_delay, :k, :k].set(new_Rinv)
278
+
279
+ if current_delay == a.shape[0] - 1:
280
+ Ainv = _get_delayed_updates(Ainv, a, Rinv)
281
+ carrier = PfCarrier(Ainv, jnp.zeros_like(a), jnp.zeros_like(Rinv))
282
+ else:
283
+ carrier = PfCarrier(Ainv, a, Rinv)
284
+ return ratio, carrier
285
+ else:
286
+ return ratio
287
+
288
+
289
+ def pf_lru_delayed(
290
+ carrier: PfCarrier,
291
+ u: Union[ArrayLike, Tuple[Array, ArrayLike]],
292
+ return_update: bool = False,
293
+ current_delay: Optional[int] = None,
294
+ ) -> Union[Array, Tuple[Array, PfCarrier]]:
295
+ r"""
296
+ Delayed low-rank update of pfaffian
297
+
298
+ :param carrier:
299
+ The existing delayed update quantities, including :math:`A_0^{-1}`, :math:`R_t^{-1}`, and
300
+
301
+ .. math::
302
+
303
+ a_t = A_{t-1}^{-1} u_t
304
+
305
+ with :math:`t` from 1 to :math:`\tau-1`.
306
+ Initially provided by `~lrux.init_pf_carrier`.
307
+
308
+ :param u:
309
+ Low-rank update vector(s) :math:`u_\tau`, the same as :math:`u` in `lrux.det_lru`.
310
+ The rank of u shouldn't exceed the maximum allowed rank specified
311
+ in `~lrux.init_pf_carrier`.
312
+
313
+ :param return_update:
314
+ Whether the new carrier with updated quantities should be returned,
315
+ defaul to False.
316
+
317
+ :param current_delay:
318
+ The current iterations :math:`\tau` of delayed updates. As python starts counting
319
+ from 0, the actual :math:`\tau` should be ``current_delay + 1``.
320
+ It must be specified when ``return_update`` is True.
321
+
322
+ :return:
323
+ ratio:
324
+ The ratio between two pfaffians
325
+
326
+ .. math::
327
+
328
+ r_\tau = \frac{\mathrm{pf}(A_\tau)}{\mathrm{pf}(A_{\tau-1})} = \frac{\mathrm{pf}(R_\tau)}{\mathrm{pf}(J)}
329
+
330
+ where
331
+
332
+ .. math::
333
+
334
+ R_\tau = J + u_\tau^T A_0^{-1} u_\tau + \sum_{t=1}^{\tau-1} (u_\tau^T a_t) (a_t^T u_\tau)
335
+
336
+ new_carrier:
337
+ Only returned when ``return_update`` is True. The new carrier contains
338
+ the quantities from the input carrier, and in addition :math:`R_\tau` and
339
+
340
+ .. math::
341
+
342
+ a_\tau = A_{\tau-1}^{-1} u_\tau = A_0^{-1} u_\tau + \sum_{t=1}^{\tau-1} a_t R_t^{-1} (a_t^T u_\tau)
343
+
344
+ When :math:`\tau` reaches the maximum delayed iterations :math:`T`
345
+ specified in `~lrux.init_pf_carrier`, i.e. ``current_delay == max_delay - 1``,
346
+ the current :math:`A_\tau` will be set as the new :math:`A_0`,
347
+ whose inverse is given by
348
+
349
+ .. math::
350
+
351
+ A_\tau^{-1} = A_0^{-1} + \sum_{t=1}^\tau a_t R_t^{-1} a_t^T
352
+
353
+ The ``Ainv`` in ``new_carrier`` will be replaced by :math:`A_\tau^{-1}`, and
354
+ ``a`` and ``Rinv`` will be set to 0.
355
+
356
+ .. warning::
357
+
358
+ This function is only recommended for heavy users who understand why and when
359
+ to use delayed updates. Otherwise, please choose `~\lrux.pf_lru`.
360
+
361
+ .. tip::
362
+
363
+ Similar to `~lrux.det_lru_delayed` and `~lrux.pf_lru`, this function is compatible
364
+ with ``jax.jit`` and ``jax.vmap``, while ``return_update`` and ``current_delay``
365
+ are static arguments which shouldn't be jitted or vmapped.
366
+
367
+ We still recommend setting ``donate_argnums=0`` in ``jax.jit`` to reuse
368
+ the memory of ``carrier`` if it's no longer needed. For instance,
369
+
370
+ .. code-block:: python
371
+
372
+ lru_vmap = jax.vmap(pf_lru_delayed, in_axes=(0, 0, None, None))
373
+ lru_jit = jax.jit(lru_vmap, static_argnums=(2, 3), donate_argnums=0)
374
+
375
+ Here is a complete example of delayed updates.
376
+
377
+ .. code-block:: python
378
+
379
+ import os
380
+ os.environ["JAX_ENABLE_X64"] = "1"
381
+
382
+ import random
383
+ import jax
384
+ import jax.numpy as jnp
385
+ import jax.random as jr
386
+ from lrux import skew_eye, pf, init_pf_carrier, pf_lru_delayed
387
+
388
+ def _get_key():
389
+ seed = random.randint(0, 2**31 - 1)
390
+ return jr.key(seed)
391
+
392
+ dtype = jnp.float64
393
+ n = 10
394
+ k = 2
395
+ max_delay = n // 2
396
+ A = jr.normal(_get_key(), (n, n), dtype)
397
+ A = (A - A.T) / 2
398
+ carrier = init_pf_carrier(A, max_delay, k)
399
+ pfA0 = pf(A)
400
+
401
+ lru_fn = jax.jit(pf_lru_delayed, static_argnums=(2, 3), donate_argnums=0)
402
+
403
+ for i in range(20):
404
+ current_delay = i % max_delay
405
+ ki = random.randint(0, k // 2) * 2 # ensure ki is even
406
+ u = jr.normal(_get_key(), (n, ki), dtype)
407
+
408
+ ratio, carrier = lru_fn(carrier, u, True, current_delay)
409
+ J = skew_eye(ki // 2, dtype)
410
+ A -= u @ J @ u.T
411
+ pfA1 = pf(A)
412
+ assert jnp.allclose(ratio, pfA1 / pfA0)
413
+ pfA0 = pfA1
414
+ """
415
+ max_delay = carrier.a.shape[0]
416
+ if current_delay is None:
417
+ if return_update:
418
+ raise ValueError("`current_delay` must be specified to return updates.")
419
+ current_delay = max_delay - 1
420
+
421
+ elif current_delay < 0 or current_delay >= max_delay:
422
+ raise ValueError(
423
+ f"`current_delay` should be in range [0, {max_delay}), got {current_delay}."
424
+ )
425
+
426
+ Ainv = carrier.Ainv
427
+ u = _standardize_uv(u, Ainv.shape[0], Ainv.dtype)
428
+
429
+ return _get_delayed_output(carrier, u, return_update, current_delay)
lrux/pfaffian.py ADDED
@@ -0,0 +1,300 @@
1
+ from typing import Tuple, NamedTuple, Optional, Callable
2
+ from jax import Array
3
+ from functools import partial
4
+ import jax
5
+ import jax.numpy as jnp
6
+
7
+
8
+ def skew_eye(n: int, dtype: Optional[jnp.dtype] = None) -> Array:
9
+ r"""
10
+ The skew-symmetric identity matrix
11
+
12
+ :param n:
13
+ Number of rows in the output divided by 2.
14
+
15
+ :param dtype:
16
+ Optional dtype, default to floating point.
17
+
18
+ :return:
19
+ The skew-symmetric identity matrix of shape (2n, 2n), defined as
20
+
21
+ .. math::
22
+
23
+ J = \begin{pmatrix}
24
+ 0 & I \\ -I & 0
25
+ \end{pmatrix}
26
+ """
27
+ I = jnp.eye(n, dtype=dtype)
28
+ O = jnp.zeros((n, n), dtype=dtype)
29
+ return jnp.block([[O, I], [-I, O]])
30
+
31
+
32
+ def _pfaffian_direct(A: Array) -> Array:
33
+ n = A.shape[-1]
34
+ batch = A.shape[:-2]
35
+
36
+ if n % 2 == 1:
37
+ return jnp.zeros(batch, dtype=A.dtype)
38
+
39
+ # By convention, pfaffian of an empty matrix is 1
40
+ elif n == 0:
41
+ return jnp.ones(batch, dtype=A.dtype)
42
+
43
+ elif n == 2:
44
+ return A[..., 0, 1]
45
+
46
+ elif n == 4:
47
+ idx = jnp.triu_indices(n, 1)
48
+ A_upper = A[..., idx[0], idx[1]]
49
+ a, b, c, d, e, f = jnp.moveaxis(A_upper, -1, 0)
50
+ return a * f - b * e + d * c
51
+
52
+
53
+ def _householder(x: Array, n: Optional[int] = None) -> Tuple[Array, Array, Array]:
54
+ if n is None:
55
+ n = 0
56
+ x0 = x[0]
57
+ x = x.at[0].set(0)
58
+ else:
59
+ x0 = x[n]
60
+ x = jnp.where(jnp.arange(x.size) <= n, 0, x)
61
+
62
+ sigma = jnp.vdot(x, x)
63
+ norm_x = jnp.sqrt(x0.conj() * x0 + sigma)
64
+
65
+ phase = jnp.where(x0 == 0.0, 1.0, jnp.sign(x0))
66
+ alpha = -phase * norm_x
67
+
68
+ v = x.at[n].set(x0 - alpha)
69
+ v *= jax.lax.rsqrt(jnp.vdot(v, v))
70
+
71
+ cond = sigma == 0.0
72
+ v = jnp.where(cond, 0, v)
73
+ tau = jnp.where(cond, 0, 2)
74
+ alpha = jnp.where(cond, x0, alpha)
75
+
76
+ return v, tau, alpha
77
+
78
+
79
+ @jax.custom_jvp
80
+ def _slogpf_householder(A: Array) -> Array:
81
+ n = A.shape[0]
82
+
83
+ def body_fun(i, val):
84
+ A, sign, log = val
85
+ v, tau, alpha = _householder(A[:, i], i + 1)
86
+ w = tau * A @ v.conj()
87
+ vw = jnp.outer(v, w)
88
+ A += vw - vw.T
89
+
90
+ new_val = (1 - tau) * jnp.where(i % 2 == 0, -alpha, 1.0)
91
+ sign *= jnp.sign(new_val)
92
+ log += jnp.log(jnp.abs(new_val))
93
+ return A, sign, log
94
+
95
+ sign = jnp.array(1, dtype=A.dtype)
96
+ log = jnp.array(0, dtype=jnp.finfo(A.dtype).dtype)
97
+ init_val = (A, sign, log)
98
+ A, sign, log = jax.lax.fori_loop(0, n - 2, body_fun, init_val)
99
+
100
+ sign *= jnp.sign(A[n - 2, n - 1])
101
+ log += jnp.log(jnp.abs(A[n - 2, n - 1]))
102
+ return sign, log
103
+
104
+
105
+ @jax.custom_jvp
106
+ def _slogpf_householder_for(A: Array) -> Tuple[Array, Array]:
107
+ vals = []
108
+ for i in range(A.shape[0] - 2):
109
+ v, tau, alpha = _householder(A[1:, 0])
110
+ A = A[1:, 1:]
111
+ w = tau * A @ v.conj()
112
+ vw = jnp.outer(v, w)
113
+ A += vw - vw.T
114
+
115
+ vals.append(1 - tau)
116
+ if i % 2 == 0:
117
+ vals.append(-alpha)
118
+
119
+ vals.append(A[-2, -1])
120
+ vals = jnp.asarray(vals)
121
+
122
+ sign = jnp.prod(jnp.sign(vals))
123
+ log = jnp.sum(jnp.log(jnp.abs(vals)))
124
+ return sign, log
125
+
126
+
127
+ @jax.custom_jvp
128
+ def _slogpf_schur(A: Array) -> Tuple[Array, Array]:
129
+ T, Z = jax.scipy.linalg.schur(A)
130
+ vals = jnp.diag(T, k=1)[::2]
131
+ s, _ = jnp.linalg.slogdet(Z)
132
+
133
+ sign = s * jnp.prod(jnp.sign(vals))
134
+ log = jnp.sum(jnp.log(jnp.abs(vals)))
135
+ return sign, log
136
+
137
+
138
+ def _check_input(
139
+ A: Array, method: str
140
+ ) -> Tuple[Array, Callable[[Array], Tuple[Array, Array]]]:
141
+ if A.ndim < 2 or A.shape[-2] != A.shape[-1]:
142
+ raise ValueError(
143
+ f"The expected input is a square matrix or a batch of them, got input shape {A.shape}."
144
+ )
145
+ A = (A - jnp.swapaxes(A, -2, -1)) / 2
146
+
147
+ if method == "householder":
148
+ slogpf_fn = _slogpf_householder
149
+ elif method == "householder_for":
150
+ slogpf_fn = _slogpf_householder_for
151
+ elif method == "schur":
152
+ if jnp.issubdtype(A, jnp.complexfloating):
153
+ raise ValueError("The schur method is only available for real dtypes.")
154
+ slogpf_fn = _slogpf_schur
155
+ else:
156
+ raise ValueError(
157
+ f"Unknown pfaffian method '{method}'. "
158
+ "Supported methods include 'householder', 'householder_for' and 'schur'."
159
+ )
160
+
161
+ return A, slogpf_fn
162
+
163
+
164
+ class SlogpfResult(NamedTuple):
165
+ sign: Array
166
+ logabspf: Array
167
+
168
+
169
+ @partial(jax.jit, static_argnames=("method",))
170
+ def slogpf(A: Array, *, method: str = "householder") -> SlogpfResult:
171
+ """
172
+ Compute the sign and (natural) logarithm of the pfaffian of an array.
173
+
174
+ :param A:
175
+ An array with shape (..., n, n)
176
+
177
+ :param method:
178
+ The method used to compute the pfaffian. Options include
179
+
180
+ ``"householder"``:
181
+ Houserholder transformation, internally using ``jax.lax.fori_loop``
182
+ to balance between running and compiling time;
183
+
184
+ ``"householder_for"``:
185
+ Houserholder transformation, internally using jitted python for loops
186
+ to reduce the running time at a cost of much longer compiling time;
187
+
188
+ ``"schur""``:
189
+ Schur decomposition using
190
+ `jax.scipy.linalg.schur <https://docs.jax.dev/en/latest/_autosummary/jax.scipy.linalg.schur.html>`_,
191
+ but only available on CPU and real dtypes.
192
+
193
+ The default method is ``"householder"``.
194
+
195
+ :return:
196
+ A ``NamedTuple`` with the following attributes. Both attribtutes have the same
197
+ batch dimension as the input A.
198
+
199
+ ``sign``:
200
+ sign(A). For a real input, it's 1, 0, or -1. For a complex input, it's a
201
+ complex number with absolute value 1, or else 0.
202
+
203
+ ``logabspf``:
204
+ The natural log of the absolute value of the pfaffian.
205
+
206
+ .. tip::
207
+
208
+ The input A is always skew-symmetrized.
209
+
210
+ .. tip::
211
+
212
+ This function has ``jax.custom_jvp`` defined and is backward compatible for both real and complex dtypes.
213
+ """
214
+ A, slogpf_fn = _check_input(A, method)
215
+
216
+ n = A.shape[-1]
217
+ if n <= 4 or n % 2 == 1:
218
+ pfA = _pfaffian_direct(A)
219
+ return SlogpfResult(jnp.sign(pfA), jnp.log(jnp.abs(pfA)))
220
+ else:
221
+ batch = A.shape[:-2]
222
+ A = A.reshape(-1, n, n)
223
+
224
+ outputs = jax.vmap(slogpf_fn)(A)
225
+ return SlogpfResult(outputs[0].reshape(batch), outputs[1].reshape(batch))
226
+
227
+
228
+ @partial(jax.jit, static_argnames=("method",))
229
+ def pf(A: Array, *, method: str = "householder") -> Array:
230
+ """
231
+ Compute the pfaffian of an array
232
+
233
+ :param A:
234
+ An array with shape (..., n, n)
235
+
236
+ :param method:
237
+ The method used to compute the pfaffian. Options include
238
+
239
+ ``"householder"``:
240
+ Houserholder transformation, internally using ``jax.lax.fori_loop``
241
+ to balance between running and compiling time;
242
+
243
+ ``"householder_for"``:
244
+ Houserholder transformation, internally using jitted python for loops
245
+ to reduce the running time at a cost of much longer compiling time;
246
+
247
+ ``"schur""``:
248
+ Schur decomposition using
249
+ `jax.scipy.linalg.schur <https://docs.jax.dev/en/latest/_autosummary/jax.scipy.linalg.schur.html>`_,
250
+ but only available on CPU and real dtypes.
251
+
252
+ The default method is ``"householder"``.
253
+
254
+ :return:
255
+ The pfaffian of A with the same batch dimensions.
256
+
257
+ .. tip::
258
+
259
+ The input A is always be skew-symmetrized.
260
+
261
+ .. tip::
262
+
263
+ This function has ``jax.custom_jvp`` defined and is backward compatible for both real and complex dtypes.
264
+ """
265
+ A, slogpf_fn = _check_input(A, method)
266
+
267
+ n = A.shape[-1]
268
+ if n <= 4 or n % 2 == 1:
269
+ return _pfaffian_direct(A)
270
+ else:
271
+ batch = A.shape[:-2]
272
+ A = A.reshape(-1, n, n)
273
+
274
+ sign, log = jax.vmap(slogpf_fn)(A)
275
+ pfA = sign * jnp.exp(log)
276
+ return pfA.reshape(batch)
277
+
278
+
279
+ def _slogpf_jvp(
280
+ primals: Tuple[Array], tangents: Tuple[Array], method: str
281
+ ) -> Tuple[Tuple[Array, Array], Tuple[Array, Array]]:
282
+ (A,) = primals
283
+ (dA,) = tangents
284
+ A, slogpf_fn = _check_input(A, method)
285
+
286
+ sign, ans = slogpf_fn(A)
287
+ ans_dot = jnp.trace(jnp.linalg.solve(A, dA)) / 2
288
+
289
+ if jnp.issubdtype(A.dtype, jnp.complexfloating):
290
+ sign_dot = 1j * sign * ans_dot.imag
291
+ ans_dot = ans_dot.real
292
+ else:
293
+ sign_dot = jnp.zeros_like(sign)
294
+
295
+ return (sign, ans), (sign_dot, ans_dot)
296
+
297
+
298
+ _slogpf_householder.defjvp(partial(_slogpf_jvp, method="householder"))
299
+ _slogpf_householder_for.defjvp(partial(_slogpf_jvp, method="householder_for"))
300
+ _slogpf_schur.defjvp(partial(_slogpf_jvp, method="schur"))
@@ -0,0 +1,9 @@
1
+ Metadata-Version: 2.4
2
+ Name: lrux
3
+ Version: 0.1.0
4
+ Summary: Fast low-rank updates (LRU) of matrix determinants and pfaffians in JAX
5
+ Author-email: Ao Chen <chenao.phys@gmail.com>, Christopher Roth <christopher_roth@utexas.edu>
6
+ Requires-Python: >=3.8
7
+ License-File: LICENSE
8
+ Requires-Dist: jax>=0.4.4
9
+ Dynamic: license-file
@@ -0,0 +1,9 @@
1
+ lrux/__init__.py,sha256=nD_NrlBN23t7aiwiSdsuHEK8f4FJwt9t5Bmmqmrr1hA,190
2
+ lrux/det_lru.py,sha256=6TKJkZnErS8lLlYOc1Tnr-TayCJ7p0lzujGSssbO9fs,17630
3
+ lrux/pf_lru.py,sha256=K53igWzN8Q4Kq93hsS_TKCI7o8xOsXleRkRzegkOEw4,14036
4
+ lrux/pfaffian.py,sha256=_QaABw8kBB9Yy1YugdQJzdtppreP0wbSEqwvCS57otU,8590
5
+ lrux-0.1.0.dist-info/licenses/LICENSE,sha256=XK1r3QjKtOz1L8u1S5X9RLysQI1z7JWAvwf6mEgjvco,1064
6
+ lrux-0.1.0.dist-info/METADATA,sha256=HfVS0qJiXgjI9wfUElCwXCdu_rNJrw1QRDwOLG6dVyg,316
7
+ lrux-0.1.0.dist-info/WHEEL,sha256=_zCd3N1l69ArxyTb8rzEoP9TpbYXkqRFSNOD5OuxnTs,91
8
+ lrux-0.1.0.dist-info/top_level.txt,sha256=O8y6IfUNFFU5hU-DKZ6WGcK-AIh--pAvaoE6gxzkDyc,5
9
+ lrux-0.1.0.dist-info/RECORD,,
@@ -0,0 +1,5 @@
1
+ Wheel-Version: 1.0
2
+ Generator: setuptools (80.9.0)
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
5
+
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2025 Chen Ao
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
@@ -0,0 +1 @@
1
+ lrux