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 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)