splineax 0.2.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.
- splineax/__init__.py +30 -0
- splineax/operators/__init__.py +7 -0
- splineax/operators/_bcoo.py +78 -0
- splineax/operators/_bcsr.py +82 -0
- splineax/operators/_jacobian.py +610 -0
- splineax/operators/_operations.py +198 -0
- splineax/operators/_sparse.py +32 -0
- splineax/py.typed +0 -0
- splineax/solvers/__init__.py +22 -0
- splineax/solvers/_auto.py +114 -0
- splineax/solvers/_klu.py +576 -0
- splineax/solvers/_sparse.py +127 -0
- splineax/solvers/_spsolve.py +227 -0
- splineax-0.2.0.dist-info/METADATA +84 -0
- splineax-0.2.0.dist-info/RECORD +16 -0
- splineax-0.2.0.dist-info/WHEEL +4 -0
splineax/__init__.py
ADDED
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
from .operators import (
|
|
2
|
+
BCOOLinearOperator as BCOOLinearOperator,
|
|
3
|
+
)
|
|
4
|
+
from .operators import (
|
|
5
|
+
BCSRLinearOperator as BCSRLinearOperator,
|
|
6
|
+
)
|
|
7
|
+
from .operators import (
|
|
8
|
+
JacobianColoring as JacobianColoring,
|
|
9
|
+
)
|
|
10
|
+
from .operators import (
|
|
11
|
+
SparseJacobianLinearOperator as SparseJacobianLinearOperator,
|
|
12
|
+
)
|
|
13
|
+
from .operators import (
|
|
14
|
+
SparseJacobianLinearOperatorColoring as SparseJacobianLinearOperatorColoring,
|
|
15
|
+
)
|
|
16
|
+
from .solvers import (
|
|
17
|
+
KLU as KLU,
|
|
18
|
+
)
|
|
19
|
+
from .solvers import (
|
|
20
|
+
AbstractSparseLinearSolver as AbstractSparseLinearSolver,
|
|
21
|
+
)
|
|
22
|
+
from .solvers import (
|
|
23
|
+
AutoSparseLinearSolver as AutoSparseLinearSolver,
|
|
24
|
+
)
|
|
25
|
+
from .solvers import (
|
|
26
|
+
SparseLinearSolver as SparseLinearSolver,
|
|
27
|
+
)
|
|
28
|
+
from .solvers import (
|
|
29
|
+
Spsolve as Spsolve,
|
|
30
|
+
)
|
|
@@ -0,0 +1,7 @@
|
|
|
1
|
+
from ._bcoo import BCOOLinearOperator as BCOOLinearOperator
|
|
2
|
+
from ._bcsr import BCSRLinearOperator as BCSRLinearOperator
|
|
3
|
+
from ._jacobian import JacobianColoring as JacobianColoring
|
|
4
|
+
from ._jacobian import SparseJacobianLinearOperator as SparseJacobianLinearOperator
|
|
5
|
+
from ._jacobian import (
|
|
6
|
+
SparseJacobianLinearOperatorColoring as SparseJacobianLinearOperatorColoring,
|
|
7
|
+
)
|
|
@@ -0,0 +1,78 @@
|
|
|
1
|
+
import equinox as eqx
|
|
2
|
+
import jax
|
|
3
|
+
import jax.numpy as jnp
|
|
4
|
+
from jax.experimental.sparse import BCOO
|
|
5
|
+
from jaxtyping import Array, Inexact
|
|
6
|
+
from lineax import AbstractLinearOperator, is_symmetric
|
|
7
|
+
from lineax._tags import transpose_tags
|
|
8
|
+
|
|
9
|
+
from ._operations import (
|
|
10
|
+
register_sparse_operator,
|
|
11
|
+
sparse_as_matrix,
|
|
12
|
+
sparse_in_structure,
|
|
13
|
+
sparse_mv,
|
|
14
|
+
sparse_out_structure,
|
|
15
|
+
)
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class BCOOLinearOperator(AbstractLinearOperator):
|
|
19
|
+
"""Wraps a `jax.experimental.sparse.BCOO` array into a linear operator.
|
|
20
|
+
|
|
21
|
+
If the matrix has shape `(a, b)` then matrix-vector multiplication (`self.mv`) is
|
|
22
|
+
defined in the usual way: as accepting a vector of shape `(b,)` and returning a
|
|
23
|
+
vector of shape `(a,)`.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
matrix: Inexact[BCOO, "a b"]
|
|
27
|
+
tags: frozenset[object] = eqx.field(static=True)
|
|
28
|
+
|
|
29
|
+
def __init__(
|
|
30
|
+
self, matrix: Inexact[BCOO, "a b"], tags: object | frozenset[object] = ()
|
|
31
|
+
):
|
|
32
|
+
"""**Arguments:**
|
|
33
|
+
|
|
34
|
+
- `matrix`: a two-dimensional `BCOO` array. For an array with shape `(a, b)`
|
|
35
|
+
then this operator can perform matrix-vector products on a vector of shape
|
|
36
|
+
`(b,)` to return a vector of shape `(a,)`.
|
|
37
|
+
- `tags`: any tags indicating whether this matrix has any particular properties,
|
|
38
|
+
like symmetry or positive-definite-ness. Note that these properties are
|
|
39
|
+
unchecked and you may get incorrect values elsewhere if these tags are
|
|
40
|
+
wrong.
|
|
41
|
+
"""
|
|
42
|
+
if matrix.ndim != 2:
|
|
43
|
+
raise ValueError(
|
|
44
|
+
"`BCOOLinearOperator(matrix=...)` should be 2-dimensional."
|
|
45
|
+
)
|
|
46
|
+
if not jnp.issubdtype(matrix.dtype, jnp.inexact):
|
|
47
|
+
matrix = BCOO(
|
|
48
|
+
(matrix.data.astype(float), matrix.indices), shape=matrix.shape
|
|
49
|
+
)
|
|
50
|
+
self.matrix = matrix
|
|
51
|
+
self.tags = tags if isinstance(tags, frozenset) else frozenset([tags])
|
|
52
|
+
|
|
53
|
+
def mv(self, vector: Inexact[Array, " b"]) -> Inexact[Array, " a"]:
|
|
54
|
+
return sparse_mv(self.matrix, vector)
|
|
55
|
+
|
|
56
|
+
def as_matrix(self) -> Inexact[Array, "a b"]:
|
|
57
|
+
return sparse_as_matrix(self.matrix)
|
|
58
|
+
|
|
59
|
+
def transpose(self) -> "BCOOLinearOperator":
|
|
60
|
+
if is_symmetric(self):
|
|
61
|
+
return self
|
|
62
|
+
matrix_T: BCOO = self.matrix.T
|
|
63
|
+
return BCOOLinearOperator(matrix_T, transpose_tags(self.tags))
|
|
64
|
+
|
|
65
|
+
def in_structure(self) -> jax.ShapeDtypeStruct:
|
|
66
|
+
return sparse_in_structure(self)
|
|
67
|
+
|
|
68
|
+
def out_structure(self) -> jax.ShapeDtypeStruct:
|
|
69
|
+
return sparse_out_structure(self)
|
|
70
|
+
|
|
71
|
+
def _conj(self) -> "BCOOLinearOperator":
|
|
72
|
+
matrix = BCOO(
|
|
73
|
+
(self.matrix.data.conj(), self.matrix.indices), shape=self.matrix.shape
|
|
74
|
+
)
|
|
75
|
+
return BCOOLinearOperator(matrix, self.tags)
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
register_sparse_operator(BCOOLinearOperator)
|
|
@@ -0,0 +1,82 @@
|
|
|
1
|
+
import equinox as eqx
|
|
2
|
+
import jax
|
|
3
|
+
import jax.numpy as jnp
|
|
4
|
+
from jax.experimental.sparse import BCOO, BCSR
|
|
5
|
+
from jaxtyping import Array, Inexact
|
|
6
|
+
from lineax import AbstractLinearOperator, is_symmetric
|
|
7
|
+
from lineax._tags import transpose_tags
|
|
8
|
+
|
|
9
|
+
from ._operations import (
|
|
10
|
+
register_sparse_operator,
|
|
11
|
+
sparse_as_matrix,
|
|
12
|
+
sparse_in_structure,
|
|
13
|
+
sparse_mv,
|
|
14
|
+
sparse_out_structure,
|
|
15
|
+
)
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class BCSRLinearOperator(AbstractLinearOperator):
|
|
19
|
+
"""Wraps a `jax.experimental.sparse.BCSR` array into a linear operator.
|
|
20
|
+
|
|
21
|
+
If the matrix has shape `(a, b)` then matrix-vector multiplication (`self.mv`) is
|
|
22
|
+
defined in the usual way: as accepting a vector of shape `(b,)` and returning a
|
|
23
|
+
vector of shape `(a,)`.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
matrix: Inexact[BCSR, "a b"]
|
|
27
|
+
tags: frozenset[object] = eqx.field(static=True)
|
|
28
|
+
|
|
29
|
+
def __init__(
|
|
30
|
+
self, matrix: Inexact[BCSR, "a b"], tags: object | frozenset[object] = ()
|
|
31
|
+
):
|
|
32
|
+
"""**Arguments:**
|
|
33
|
+
|
|
34
|
+
- `matrix`: a two-dimensional `BCSR` array. For an array with shape `(a, b)`
|
|
35
|
+
then this operator can perform matrix-vector products on a vector of shape
|
|
36
|
+
`(b,)` to return a vector of shape `(a,)`.
|
|
37
|
+
- `tags`: any tags indicating whether this matrix has any particular properties,
|
|
38
|
+
like symmetry or positive-definite-ness. Note that these properties are
|
|
39
|
+
unchecked and you may get incorrect values elsewhere if these tags are
|
|
40
|
+
wrong.
|
|
41
|
+
"""
|
|
42
|
+
if matrix.ndim != 2:
|
|
43
|
+
raise ValueError(
|
|
44
|
+
"`BCSRLinearOperator(matrix=...)` should be 2-dimensional."
|
|
45
|
+
)
|
|
46
|
+
if not jnp.issubdtype(matrix.dtype, jnp.inexact):
|
|
47
|
+
matrix = BCSR(
|
|
48
|
+
(matrix.data.astype(jnp.float32), matrix.indices, matrix.indptr),
|
|
49
|
+
shape=matrix.shape,
|
|
50
|
+
)
|
|
51
|
+
self.matrix = matrix
|
|
52
|
+
self.tags = tags if isinstance(tags, frozenset) else frozenset([tags])
|
|
53
|
+
|
|
54
|
+
def mv(self, vector: Inexact[Array, " b"]) -> Inexact[Array, " a"]:
|
|
55
|
+
return sparse_mv(self.matrix, vector)
|
|
56
|
+
|
|
57
|
+
def as_matrix(self) -> Inexact[Array, "a b"]:
|
|
58
|
+
return sparse_as_matrix(self.matrix)
|
|
59
|
+
|
|
60
|
+
def transpose(self) -> "BCSRLinearOperator":
|
|
61
|
+
if is_symmetric(self):
|
|
62
|
+
return self
|
|
63
|
+
# `BCSR.transpose` is not implemented in JAX; round-trip through `BCOO`.
|
|
64
|
+
matrix_T_bcoo: BCOO = self.matrix.to_bcoo().T
|
|
65
|
+
matrix_T = BCSR.from_bcoo(matrix_T_bcoo)
|
|
66
|
+
return BCSRLinearOperator(matrix_T, transpose_tags(self.tags))
|
|
67
|
+
|
|
68
|
+
def in_structure(self) -> jax.ShapeDtypeStruct:
|
|
69
|
+
return sparse_in_structure(self)
|
|
70
|
+
|
|
71
|
+
def out_structure(self) -> jax.ShapeDtypeStruct:
|
|
72
|
+
return sparse_out_structure(self)
|
|
73
|
+
|
|
74
|
+
def _conj(self) -> "BCSRLinearOperator":
|
|
75
|
+
matrix = BCSR(
|
|
76
|
+
(self.matrix.data.conj(), self.matrix.indices, self.matrix.indptr),
|
|
77
|
+
shape=self.matrix.shape,
|
|
78
|
+
)
|
|
79
|
+
return BCSRLinearOperator(matrix, self.tags)
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
register_sparse_operator(BCSRLinearOperator)
|