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 +3 -0
- lrux/det_lru.py +563 -0
- lrux/pf_lru.py +429 -0
- lrux/pfaffian.py +300 -0
- lrux-0.1.0.dist-info/METADATA +9 -0
- lrux-0.1.0.dist-info/RECORD +9 -0
- lrux-0.1.0.dist-info/WHEEL +5 -0
- lrux-0.1.0.dist-info/licenses/LICENSE +21 -0
- lrux-0.1.0.dist-info/top_level.txt +1 -0
lrux/__init__.py
ADDED
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,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
|