diff --git a/CHANGELOG.md b/CHANGELOG.md index 35ded2a..84ced0a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,202 @@ and the project adheres to [Semantic Versioning](https://semver.org/). ## [Unreleased] +## [0.4.3] — 2026-09-12 + +### Removed + +- **`jax_pytree_class` is removed from the public API** (breaking). Pytree + registration is no longer a JAX-named decorator applied by hand, but a + backend-neutral registry: containers inherit `PyTreeNode` and each backend + installs its own tree protocol (see *Added*). No replacement decorator is + exported — a class becomes a pytree node by being a concrete `LinOp`, + `Functional`, or one of the three pytree `Space` types. Code that decorated its + own class should inherit `spacecore.backend.PyTreeNode` instead. +- **`SpectralLpNormFunctional` is removed** (breaking). A per-formula spectral + class duplicates the value, the gradient, the pytree methods and the validation + of its coordinate twin — it re-checked `p >= 1` that `LpNormFunctional` already + enforces. The Schatten `p`-norm is now the spectral *lift* of the coordinate + `p`-norm: + + ```python + # before + f = sc.SpectralLpNormFunctional(X, p) + # after + f = sc.spectralize(X, lambda s: sc.LpNormFunctional(s, p)) + ``` + + `NuclearNormFunctional(X)` is unchanged as a named constructor and now returns a + `SpectralFunctional` whose `base` is `LpNormFunctional(s, 1.0)`; code reading + `f.p` should read `f.base.p`. + +### Added + +- **Backend-neutral pytree registration.** `PyTreeNode` (the capability mixin + owning the `tree_flatten` / `tree_unflatten` contract), `PyTreeRegistry`, and + `BackendOps.install_pytree_protocol()` — a plugin hook, no-op by default, that a + backend implements to teach its transform machinery to see inside SpaceCore + containers. The registry wires the class × backend cross-product with + back-fill in both directions, so registration no longer depends on import order. + `LinOp` and `Functional` now carry the capability on their bases, so every + concrete subclass — including user-defined ones — is registered automatically. +- **Torch is now a transform-capable backend.** `TorchOps.install_pytree_protocol` + registers containers with `torch.utils._pytree`, so a Torch-backed operator can + flow through `torch.compile` and `torch.func` as a traced argument rather than an + opaque leaf. One registration also covers `torch.utils._cxx_pytree`. +- **`SpectralFunctional`, `spectralize`, `eigenvalue_space`** — lift **any** + symmetric coordinate functional onto a Jordan spectrum (Lewis's theorem), instead + of hand-writing a class per formula. `spectralize(X, NegativeEntropyFunctional)` + is the von Neumann entropy; `SquaredL2NormFunctional` lifts to the squared + Frobenius norm, `HuberFunctional` to its spectral analogue. The base functional + must be symmetric — eigenvalues have no canonical order, so a non-symmetric `f` + makes `f(lambda(X))` ill-defined rather than merely mis-differentiated. +- **`RealifiedFunctional` and `realify`** — view a complex-domain functional over + stacked real coordinates `(Re v, Im v)`, for optimizers that assume a real vector + space. The real gradient is `(Re g_v, Im g_v)` with `g_v` the *coordinate* + gradient, so the metric correction stays where ADR-010 put it. `realify` is a + no-op on an already-real domain. +- **`OpsRegistry`** — the backend registry extracted out of `Contextual`, which now + holds ambient policy only. Instantiable, so registration is testable without the + process-wide singleton; `register_ops` still raises `ContextConflictError` on a + duplicate family. +- **`BackendOps.complex_dtype`** — the inverse of `real_dtype`. Deriving it at the + call site is not portable: NumPy promotes `float32` against a Python complex to + `complex128`, while JAX and Torch give `complex64`. + +- **`ProductFunctional` and `make_functional_product`** — the functional algebra + becomes multiplicative. `F * G` is the pointwise product `F(x) * G(x)` on a + shared domain, with the product-rule Riesz gradient + + ``` + grad(F·G)(x) = conj(G(x))·grad F(x) + conj(F(x))·grad G(x) + ``` + + combined through the domain's own `scale`/`add` (a domain element may be a + pytree), and a `value_and_grad` that evaluates each factor once. Deliberately + **binary**: the product rule is a two-factor law, and `(F*G)*H` expresses the + n-ary case at the same cost. The conjugations are identities for the usual + real-valued factors. +- **`ConstantFunctional` and `make_constant_functional`** — the constant map + `x -> c`, with zero gradient. This is the embedding of a scalar into the + functional algebra, previously unrepresentable: the algebra had a zero element + and an affine shift but no constant node. `make_constant_functional` collapses + `c = 0` to `ZeroFunctional` so the additive identity keeps one representation. +- **`spacecore.opfamily`: `OperatorFamily`, `FunctionalScaledOperator`, + `make_functional_scaled_operator`** — `F * A` for a `Functional` and a `LinOp` + is the functional-weighted map `m(x) = F(x) A x`. That map is **not linear** + (both the scale and the direction move with `x`), so it is deliberately *not* a + `LinOp`; it is a point-indexed *family* `x -> A_x`, of which a `LinOp` is the + constant case. Freezing the point recovers linearity, and each point carries + two different operators: + + - `m.at(x)` — the frozen member `F(x) · A`, an ordinary `ScaledLinOp` that + composes, sums and has an adjoint; + - `m.linearize_at(x)` — the derivative `Dm(x)[h] = ·Ax + F(x)·A h`, + with metric adjoint `_Y·grad F(x) + conj(F(x))·A^# w`, for Newton-type + steps. + + They differ by exactly a rank-one term and coincide only when `F` is constant. + A `ConstantFunctional` weight collapses to a plain `ScaledLinOp`, so the linear + case is never forced through the non-linear type. The module is top-level + because it depends on both `linop` and `functional`, and neither depends on it. +- **`checked_method(out_scalar=True)` / `out_batched_scalar=True`** — the codomain + check for a `Functional`. `out_space=` names an attribute holding a `Space`, and + a functional's codomain is the scalar *field*, reported only as a string, so the + decorator that guards every `LinOp` output had nothing to bind to and the check + was hand-written per subclass. `out_scalar` asserts `shape == ()`; + `out_batched_scalar` asserts `(N,)` with `N` read from the input named by + `in_space`, which it therefore requires. +- **`Space.scalar_field`, `Space.declared_scalar_field`, `Space.check_scalar`** — + the field of scalars a space is closed under, *declared* rather than inferred + from the dtype. It defaults to `field`; `HermitianSpace` declares `"real"`, + because complex Hermitian matrices have complex entries but form a **real** + vector space (`i·H` is anti-Hermitian). + +### Changed + +- **`check_level` is a property of the bound object, not of `Context`** (breaking). + `Context(ops, dtype=..., check_level=...)` and the deprecated `enable_checks=` + argument are gone; `Context` is now exactly `(ops, dtype)`. Pass `check_level=` + to the space / operator / functional constructor, or set it ambiently with + `set_check_level` / `use_check_level`. Two contexts differing only in strictness + are now correctly the same context. +- **`spacecore._contextual` is now `spacecore.contextual`** (breaking), and + `Context` moved out of `spacecore.backend`. The top-level `spacecore.Context` + re-export is unchanged. +- **Ambient context and check level are scoped with `contextvars`.** + `use_context` / `use_check_level` install a `ContextVar` override unwound by + `Token`, so nesting is exact and concurrent threads or async tasks cannot clobber + one another. `set_context` / `set_check_level` still write the process-wide + baseline. The backend registry deliberately stays global — scoping it would make + a backend registered inside a `with` block vanish on exit. +- **`available_ops()` is memoized and returns a tuple** rather than a list. + Discovery attempts a real import per optional backend and scans entry-point + metadata; it ran three times per `import spacecore`. A tuple because the cached + value is shared. + +- **`Functional.__mul__` / `__rmul__` now dispatch on the operand type.** A scalar + still gives `ScaledFunctional`; a `Functional` now gives the pointwise product + and a `LinOp` the functional-weighted family. Previously both returned + `NotImplemented`, so `F * G` and `F * A` raised `TypeError`. Operands that are + neither scalar-like nor `Functional` nor `LinOp` still defer to the reflected + operation. +- **Functional outputs are checked as scalars at `standard` and above.** 17 + `value` and 3 `vvalue` implementations carry the new decorator flags, replacing + four hand-written `_checks_at_least("standard")` / `_check_scalar_shape` bodies + with one implementation. The level matches what those call sites already used — + `cheap` deliberately does not run it. Cost is ~0.5 µs per `value` call, and + nothing at `check_level="none"`, which still short-circuits before any check. + `MatrixFreeLinearFunctional.vvalue` keeps its own check: it permits several + leading batch axes, which is broader than the single-axis contract `vvalue` + documents. `RealifiedFunctional.value` is checked on its output only — its input + is validated against the *complex* domain by the inner functional. +- **`_check_scalar_shape` distinguishes single from batched output** in its error + message ("Expected scalar output" vs "Expected scalar batch output"). It said + "batch" unconditionally, which was harmless while it guarded four batched paths + and misleading now that it guards ~20 mostly single-element ones. +- **`HermitianSpace.scale` / `scale_batch` reject a non-real multiplier** + (breaking, at `standard` and above). `scale(1j, H)` returned a skew-Hermitian + array still typed as an element of `Herm(n)`; the error surfaced at some later + membership check, or never at `check_level="none"`. Only a *provably* non-real + multiplier is refused, so a traced scalar under `jax.jit` still passes. `field` + remains dtype-derived and still drives equality and repr: a real-dtype and a + complex-dtype `Herm(n)` are genuinely different spaces. + +### Fixed + +- **A present-but-broken optional backend no longer aborts `import spacecore`.** + `spacecore.backend` eagerly imported the JAX subpackage to reach + `jax_pytree_class`, so an installed-but-unimportable backend raising anything + other than `ModuleNotFoundError` — a shadowed `cupy`, a partially-installed + `jax` — propagated out of the import. Discovery is now the only path into a + backend package, and it warns and skips instead. +- **One broken backend now warns once, not once per discovery call.** + +- **`ComposedFunctional` now has a gradient.** `F.compose(A).grad(x)` raised + `NotImplementedError`: the node implemented `value` but not the chain rule, + though `LinOp.rapply` already provides the metric adjoint it needs. Added + `grad`, a fused `value_and_grad`, and a batched `vgrad`, with cores registered + in the `composed-functional` kernel set: + + ``` + grad(F o A)(x) = A^#(grad F(A x)) + ``` + + `rapply` **is** `A^#` (ADR-009), so no Riesz map is applied on top of it — + that would count the geometry twice — and no explicit conjugation appears, + because the adjoint identity absorbs it. `value_and_grad` applies `A` once and + shares the image, where the inherited default applied it twice. The typed + specializations in `make_functional_composed` (`InnerProductFunctional`, + `LinOpQuadraticForm`) were unaffected, and are used as a cross-check on the + generic node. +- **`scalar_eq` no longer swallows every exception.** It wrapped its comparison in + a bare `except Exception: return False`, so a raising `__eq__` was silently + reported as inequality. Narrowed to `TypeError` — the base class of JAX's + `TracerBoolConversionError` and of any "cannot reduce to a concrete bool" + failure — which keeps the intended verdict for an abstract scalar (undecidable, + so canonicalization is skipped and the expression tree stays unfolded but + correct) while letting a genuinely broken `__eq__` propagate. + ## [0.4.2] — 2026-07-01 ### Added diff --git a/README.md b/README.md index a95b8e9..ca6b159 100644 --- a/README.md +++ b/README.md @@ -180,7 +180,8 @@ level) adds named constructors over that machinery, with no new core types: `least_squares` for `½‖Ax−b‖²`, coordinate norms (`SquaredL2NormFunctional`, `LpNormFunctional`, `L1NormFunctional`), `NegativeEntropyFunctional`, `KLDivergenceFunctional`, `HuberFunctional`, the spectral -`SpectralLpNormFunctional`/`NuclearNormFunctional`, and the metric-aware +`NuclearNormFunctional` (and `spectralize` / `SpectralFunctional` to lift any +coordinate functional to a spectral one), and the metric-aware proximal primitive `generalized_shrinkage` with the wrappers `prox_l1`, `prox_l2sq`, and `project_nonneg`. Each objective's gradient is the metric (Riesz) gradient under the domain geometry, and the proximal step is taken in @@ -214,21 +215,29 @@ and [deviation catalog](https://pavlo3p.github.io/SpaceCore/design/backend_devia ## Validation Policy -A `Context` carries a `check_level` that determines how aggressively spaces, -operators, functionals, and solver preconditions validate their inputs. The -ordered levels are `CHECK_LEVELS = ("none", "cheap", "standard", "strict")`: -`cheap` covers shape/dtype/backend/tree-structure, `standard` adds membership -and Hermitian checks, and `strict` adds bounded expensive probes. Checks are -opt-in per context, so hot paths can run unvalidated while development and tests -run strict. +Every context-bound object — space, operator, functional — carries a +`check_level` that determines how aggressively it validates its inputs, together +with solver preconditions. The ordered levels are +`CHECK_LEVELS = ("none", "cheap", "standard", "strict")`: `cheap` covers +shape/dtype/backend/tree-structure, `standard` adds membership and Hermitian +checks, and `strict` adds bounded expensive probes. Checks are opt-in per +object, so hot paths can run unvalidated while development and tests run strict. + +The level is a property of the bound object rather than of the `Context`: a +context fixes backend ops and dtype, while two objects sharing that backend and +dtype may legitimately validate at different strictness. Pass `check_level=` to +a constructor, or move the ambient default with `sc.set_check_level(...)` / +`sc.use_check_level(...)`. ```python import numpy as np import spacecore as sc -ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="standard") -X = sc.DenseCoordinateSpace((2,), ctx) -A = sc.DenseLinOp(ctx.asarray([[2.0, 0.0], [0.0, 3.0]]), X, X, ctx) +ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) +X = sc.DenseCoordinateSpace((2,), ctx, check_level="standard") +A = sc.DenseLinOp( + ctx.asarray([[2.0, 0.0], [0.0, 3.0]]), X, X, ctx, check_level="standard" +) try: A.apply(ctx.asarray([1.0, 2.0, 3.0])) # wrong shape diff --git a/bench/_dashboard.py b/bench/_dashboard.py index a2c3c27..ca85f11 100644 --- a/bench/_dashboard.py +++ b/bench/_dashboard.py @@ -56,7 +56,7 @@ import webbrowser from pathlib import Path from statistics import median -from typing import Iterable +from typing import Any, Iterable from ._io import _metadata from ._probes import ProbeResult @@ -173,6 +173,7 @@ def render_dashboard( results: Iterable[ProbeResult], out_path: str | Path, baseline: Iterable[ProbeResult] | None = None, + meta: dict[str, Any] | None = None, ) -> Path: """Render an interactive HTML dashboard for a bench run. @@ -213,7 +214,11 @@ def render_dashboard( ) summary = _summary(rows) - meta = _metadata() + # Prefer the metadata recorded in the source artifact (the machine that + # *ran* the benchmark); fall back to this machine only when the artifact + # carried none. Rendering on a different host must not overwrite the + # processor/platform/versions of the run. + meta = meta or _metadata() overall = _build_overall(results, diagnoses) backends = sorted({r["backend"] for r in rows}) diff --git a/bench/_operations.py b/bench/_operations.py index 9ee411f..b36cdb5 100644 --- a/bench/_operations.py +++ b/bench/_operations.py @@ -52,25 +52,25 @@ def benchmark_check_level(check_level: str): """Build probe contexts with one explicit benchmark check level.""" token = _ACTIVE_CHECK_LEVEL.set(check_level) try: - yield + with sc.use_check_level(check_level): + yield finally: _ACTIVE_CHECK_LEVEL.reset(token) -def _backend_ctx(backend: str, *, check_level: str | None = None) -> sc.Context: +def _backend_ctx(backend: str) -> sc.Context: """Build a SpaceCore ``Context`` for the requested backend.""" - check_level = check_level or _ACTIVE_CHECK_LEVEL.get() if backend == "numpy": - return sc.Context(sc.NumpyOps(), dtype=np.float64, check_level=check_level) + return sc.Context(sc.NumpyOps(), dtype=np.float64) if backend == "jax": from tests._helpers import jax_real_dtype - return sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level=check_level) + return sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) if backend == "torch": from tests._helpers import torch_real_dtype td = torch_real_dtype() - return sc.Context(sc.TorchOps(), dtype=td, check_level=check_level) + return sc.Context(sc.TorchOps(), dtype=td) raise ValueError(f"unknown backend {backend!r}") @@ -119,9 +119,8 @@ def _make_space_add(backend: str, seed: int, size: int) -> ProbeCase: ctx = _backend_ctx(backend) space, x, x_np = _dense_vector(ctx, size, seed) _, y, y_np = _dense_vector(ctx, size, seed + 100) - unchecked_ctx = _backend_ctx(backend, check_level="none") - unchecked_space = sc.DenseCoordinateSpace((size,), unchecked_ctx) - ux, uy = unchecked_ctx.asarray(x_np), unchecked_ctx.asarray(y_np) + unchecked_space = sc.DenseCoordinateSpace((size,), ctx, check_level="none") + ux, uy = ctx.asarray(x_np), ctx.asarray(y_np) return ProbeCase( bare_label=f"{backend}: x + y", sc_label="DenseCoordinateSpace.add", @@ -137,9 +136,8 @@ def _make_space_scale(backend: str, seed: int, size: int) -> ProbeCase: ctx = _backend_ctx(backend) space, x, x_np = _dense_vector(ctx, size, seed) alpha = float(_rng(seed + 1).standard_normal()) - unchecked_ctx = _backend_ctx(backend, check_level="none") - unchecked_space = sc.DenseCoordinateSpace((size,), unchecked_ctx) - ux = unchecked_ctx.asarray(x_np) + unchecked_space = sc.DenseCoordinateSpace((size,), ctx, check_level="none") + ux = ctx.asarray(x_np) return ProbeCase( bare_label=f"{backend}: alpha * x", sc_label="DenseCoordinateSpace.scale", @@ -155,9 +153,8 @@ def _make_space_inner(backend: str, seed: int, size: int) -> ProbeCase: ctx = _backend_ctx(backend) space, x, x_np = _dense_vector(ctx, size, seed) _, y, y_np = _dense_vector(ctx, size, seed + 100) - unchecked_ctx = _backend_ctx(backend, check_level="none") - unchecked_space = sc.DenseCoordinateSpace((size,), unchecked_ctx) - ux, uy = unchecked_ctx.asarray(x_np), unchecked_ctx.asarray(y_np) + unchecked_space = sc.DenseCoordinateSpace((size,), ctx, check_level="none") + ux, uy = ctx.asarray(x_np), ctx.asarray(y_np) return ProbeCase( bare_label=f"{backend}: vdot(x, y)", sc_label="DenseCoordinateSpace.inner", @@ -172,9 +169,8 @@ def _make_space_inner(backend: str, seed: int, size: int) -> ProbeCase: def _make_space_norm(backend: str, seed: int, size: int) -> ProbeCase: ctx = _backend_ctx(backend) space, x, x_np = _dense_vector(ctx, size, seed) - unchecked_ctx = _backend_ctx(backend, check_level="none") - unchecked_space = sc.DenseCoordinateSpace((size,), unchecked_ctx) - ux = unchecked_ctx.asarray(x_np) + unchecked_space = sc.DenseCoordinateSpace((size,), ctx, check_level="none") + ux = ctx.asarray(x_np) return ProbeCase( bare_label=f"{backend}: linalg.norm(x)", sc_label="DenseCoordinateSpace.norm", @@ -227,10 +223,9 @@ def _make_dense_apply(backend: str, seed: int, size: int) -> ProbeCase: a, a_np = _dense_matrix(ctx, size, seed) _, x, x_np = _dense_vector(ctx, size, seed + 7) op = sc.DenseLinOp(a, space, space, ctx) - unchecked_ctx = _backend_ctx(backend, check_level="none") - unchecked_space = sc.DenseCoordinateSpace((size,), unchecked_ctx) - ua, ux = unchecked_ctx.asarray(a_np), unchecked_ctx.asarray(x_np) - unchecked_op = sc.DenseLinOp(ua, unchecked_space, unchecked_space, unchecked_ctx) + unchecked_space = sc.DenseCoordinateSpace((size,), ctx, check_level="none") + ua, ux = ctx.asarray(a_np), ctx.asarray(x_np) + unchecked_op = sc.DenseLinOp(ua, unchecked_space, unchecked_space, ctx) return ProbeCase( bare_label=f"{backend}: A @ x", sc_label="DenseLinOp.apply", @@ -248,10 +243,9 @@ def _make_dense_rapply(backend: str, seed: int, size: int) -> ProbeCase: a, a_np = _dense_matrix(ctx, size, seed) _, y, y_np = _dense_vector(ctx, size, seed + 8) op = sc.DenseLinOp(a, space, space, ctx) - unchecked_ctx = _backend_ctx(backend, check_level="none") - unchecked_space = sc.DenseCoordinateSpace((size,), unchecked_ctx) - ua, uy = unchecked_ctx.asarray(a_np), unchecked_ctx.asarray(y_np) - unchecked_op = sc.DenseLinOp(ua, unchecked_space, unchecked_space, unchecked_ctx) + unchecked_space = sc.DenseCoordinateSpace((size,), ctx, check_level="none") + ua, uy = ctx.asarray(a_np), ctx.asarray(y_np) + unchecked_op = sc.DenseLinOp(ua, unchecked_space, unchecked_space, ctx) return ProbeCase( bare_label=f"{backend}: A.T.conj() @ y", sc_label="DenseLinOp.rapply", @@ -271,10 +265,9 @@ def _make_dense_vapply(backend: str, seed: int, size: int) -> ProbeCase: xs_np = np.asarray(rng.standard_normal((8, size)), dtype=_np_dtype(ctx)) xs = ctx.asarray(xs_np) op = sc.DenseLinOp(a, space, space, ctx) - unchecked_ctx = _backend_ctx(backend, check_level="none") - unchecked_space = sc.DenseCoordinateSpace((size,), unchecked_ctx) - ua, uxs = unchecked_ctx.asarray(a_np), unchecked_ctx.asarray(xs_np) - unchecked_op = sc.DenseLinOp(ua, unchecked_space, unchecked_space, unchecked_ctx) + unchecked_space = sc.DenseCoordinateSpace((size,), ctx, check_level="none") + ua, uxs = ctx.asarray(a_np), ctx.asarray(xs_np) + unchecked_op = sc.DenseLinOp(ua, unchecked_space, unchecked_space, ctx) return ProbeCase( bare_label=f"{backend}: xs @ A.T", sc_label="DenseLinOp.vapply", @@ -1078,9 +1071,7 @@ def _make_matrix_free_diagonal_rapply(backend: str, seed: int, size: int) -> Pro def _make_matrix_free_fft_apply(backend: str, seed: int, size: int) -> ProbeCase: # FFT operators here use a complex context so apply/rapply round-trip works. - ctx = sc.Context( - sc.NumpyOps(), dtype=np.complex128, check_level=_ACTIVE_CHECK_LEVEL.get() - ) + ctx = sc.Context(sc.NumpyOps(), dtype=np.complex128) space = sc.DenseCoordinateSpace((size,), ctx) rng = _rng(seed) x_np = np.asarray(rng.standard_normal(size), dtype=np.complex128) @@ -1410,9 +1401,7 @@ def _make_dense_apply_device(backend: str, device: str, seed: int, size: int) -> # correctness gate keeps a float32-width tolerance for MPS. import torch - ctx = sc.Context( - sc.TorchOps(), dtype=torch.float32, check_level=_ACTIVE_CHECK_LEVEL.get() - ) + ctx = sc.Context(sc.TorchOps(), dtype=torch.float32) else: ctx = _backend_ctx(backend) space = sc.DenseCoordinateSpace((size,), ctx) diff --git a/bench/dashboard.py b/bench/dashboard.py index f993d49..5694bd2 100644 --- a/bench/dashboard.py +++ b/bench/dashboard.py @@ -115,6 +115,7 @@ def main(argv: list[str] | None = None) -> int: micro_rows = [] macro_rows = [] + run_meta = None for raw in args.input: path = Path(raw) if not path.exists(): @@ -125,6 +126,8 @@ def main(argv: list[str] | None = None) -> int: except json.JSONDecodeError as err: print(f"failed to parse {path}: {err}", file=sys.stderr) return 2 + if run_meta is None and isinstance(payload, dict) and payload.get("meta"): + run_meta = payload["meta"] kind = _classify_payload(payload) if kind == "micro": micro_rows.extend(_load_micro(path)) @@ -144,7 +147,7 @@ def main(argv: list[str] | None = None) -> int: out_path = render_dashboard( micro_rows, args.out, - macro_results=macro_rows or None, + meta=run_meta, ) print(f"wrote {out_path}") if args.open: diff --git a/bench/macro/cg_poisson.py b/bench/macro/cg_poisson.py index 045be93..04df031 100644 --- a/bench/macro/cg_poisson.py +++ b/bench/macro/cg_poisson.py @@ -212,23 +212,21 @@ def _factory(backend: str, device: str, seed: int, size_params: dict[str, Any]) maxiter = int(size_params["maxiter"]) lambda_ = float(size_params["lambda"]) - # Two contexts, two callables (the runner does NOT switch context). - ctx_none = _backend_ctx(backend, check_level="none") - ctx_cheap = _backend_ctx(backend, check_level="cheap") - np_dtype = _np_dtype(ctx_none) + # Two check levels, two callables (the runner does NOT switch context). + ctx = _backend_ctx(backend) + np_dtype = _np_dtype(ctx) # Pre-compute the right-hand side once in NumPy, ship to each backend. b_np = _make_rhs_np(n, np_dtype) # SpaceCore space and operator for each check level. - space_none = sc.DenseCoordinateSpace((n, n), ctx_none) - space_cheap = sc.DenseCoordinateSpace((n, n), ctx_cheap) + space_none = sc.DenseCoordinateSpace((n, n), ctx, check_level="none") + space_cheap = sc.DenseCoordinateSpace((n, n), ctx, check_level="cheap") mode_callables: dict[ModeName, Callable[[], Any]] = {} if backend == "numpy": - b_arr_none = ctx_none.asarray(b_np) - b_arr_cheap = ctx_cheap.asarray(b_np) + b_arr = ctx.asarray(b_np) def bare_call(b=b_np, lam=lambda_, mi=maxiter): return _numpy_cg(lambda u: _numpy_apply(u, lam), b, mi) @@ -241,13 +239,13 @@ def apply(u): apply_none = sc_apply_none_factory() apply_cheap = sc_apply_none_factory() - op_none = sc.MatrixFreeLinOp(apply_none, apply_none, space_none, space_none, ctx_none) - op_cheap = sc.MatrixFreeLinOp(apply_cheap, apply_cheap, space_cheap, space_cheap, ctx_cheap) + op_none = sc.MatrixFreeLinOp(apply_none, apply_none, space_none, space_none, ctx) + op_cheap = sc.MatrixFreeLinOp(apply_cheap, apply_cheap, space_cheap, space_cheap, ctx) - def sc_none_call(op=op_none, b=b_arr_none, mi=maxiter): + def sc_none_call(op=op_none, b=b_arr, mi=maxiter): return sc.cg(op, b, maxiter=mi, tol=0.0, atol=0.0) - def sc_cheap_call(op=op_cheap, b=b_arr_cheap, mi=maxiter): + def sc_cheap_call(op=op_cheap, b=b_arr, mi=maxiter): return sc.cg(op, b, maxiter=mi, tol=0.0, atol=0.0) mode_callables["bare"] = bare_call @@ -261,27 +259,26 @@ def sc_cheap_call(op=op_cheap, b=b_arr_cheap, mi=maxiter): import jax.numpy as jnp jax_dtype = jnp.asarray(b_np).dtype - jnp.asarray(b_np, dtype=ctx_none.dtype if hasattr(ctx_none, "dtype") else jax_dtype) - # Re-cast in the active jax dtype for both contexts. - b_arr_none = ctx_none.asarray(b_np) - b_arr_cheap = ctx_cheap.asarray(b_np) + jnp.asarray(b_np, dtype=ctx.dtype if hasattr(ctx, "dtype") else jax_dtype) + # Re-cast in the active jax dtype. + b_arr = ctx.asarray(b_np) apply_jax = _jax_apply_factory(lambda_) - def bare_call(b=b_arr_none, mi=maxiter): + def bare_call(b=b_arr, mi=maxiter): return _jax_cg(apply_jax, b, mi) # SpaceCore public path uses the same eager apply. - op_none = sc.MatrixFreeLinOp(apply_jax, apply_jax, space_none, space_none, ctx_none) + op_none = sc.MatrixFreeLinOp(apply_jax, apply_jax, space_none, space_none, ctx) op_cheap_apply = _jax_apply_factory(lambda_) op_cheap = sc.MatrixFreeLinOp( - op_cheap_apply, op_cheap_apply, space_cheap, space_cheap, ctx_cheap + op_cheap_apply, op_cheap_apply, space_cheap, space_cheap, ctx ) - def sc_none_call(op=op_none, b=b_arr_none, mi=maxiter): + def sc_none_call(op=op_none, b=b_arr, mi=maxiter): return sc.cg(op, b, maxiter=mi, tol=0.0, atol=0.0) - def sc_cheap_call(op=op_cheap, b=b_arr_cheap, mi=maxiter): + def sc_cheap_call(op=op_cheap, b=b_arr, mi=maxiter): return sc.cg(op, b, maxiter=mi, tol=0.0, atol=0.0) # Lowered: jit the apply so each matvec is a fused kernel. @@ -291,10 +288,10 @@ def sc_cheap_call(op=op_cheap, b=b_arr_cheap, mi=maxiter): # to time the first call; the runner already separates that. Here we # let it count compile cost as compile_time_ns on the lowered mode. op_lowered = sc.MatrixFreeLinOp( - jit_apply, jit_apply, space_none, space_none, ctx_none + jit_apply, jit_apply, space_none, space_none, ctx ) - def sc_lowered_call(op=op_lowered, b=b_arr_none, mi=maxiter): + def sc_lowered_call(op=op_lowered, b=b_arr, mi=maxiter): return sc.cg(op, b, maxiter=mi, tol=0.0, atol=0.0) mode_callables["bare"] = bare_call @@ -304,23 +301,22 @@ def sc_lowered_call(op=op_lowered, b=b_arr_none, mi=maxiter): elif backend == "torch": - b_arr_none = ctx_none.asarray(b_np) - b_arr_cheap = ctx_cheap.asarray(b_np) + b_arr = ctx.asarray(b_np) bare_apply = _torch_apply_factory(lambda_) - def bare_call(b=b_arr_none, mi=maxiter): + def bare_call(b=b_arr, mi=maxiter): return _torch_cg(bare_apply, b, mi) apply_none = _torch_apply_factory(lambda_) apply_cheap = _torch_apply_factory(lambda_) - op_none = sc.MatrixFreeLinOp(apply_none, apply_none, space_none, space_none, ctx_none) - op_cheap = sc.MatrixFreeLinOp(apply_cheap, apply_cheap, space_cheap, space_cheap, ctx_cheap) + op_none = sc.MatrixFreeLinOp(apply_none, apply_none, space_none, space_none, ctx) + op_cheap = sc.MatrixFreeLinOp(apply_cheap, apply_cheap, space_cheap, space_cheap, ctx) - def sc_none_call(op=op_none, b=b_arr_none, mi=maxiter): + def sc_none_call(op=op_none, b=b_arr, mi=maxiter): return sc.cg(op, b, maxiter=mi, tol=0.0, atol=0.0) - def sc_cheap_call(op=op_cheap, b=b_arr_cheap, mi=maxiter): + def sc_cheap_call(op=op_cheap, b=b_arr, mi=maxiter): return sc.cg(op, b, maxiter=mi, tol=0.0, atol=0.0) mode_callables["bare"] = bare_call diff --git a/bench/macro/density_pipeline.py b/bench/macro/density_pipeline.py index e3cf0d1..2d86e0f 100644 --- a/bench/macro/density_pipeline.py +++ b/bench/macro/density_pipeline.py @@ -20,7 +20,7 @@ * ``bare`` — pure backend ops on raw arrays. * ``spacecore_public_none`` / ``spacecore_public_cheap`` — :class:`spacecore.HermitianSpace` and :class:`spacecore.StackedSpace`, - with the corresponding :attr:`Context.check_level`. + constructed with the corresponding ``check_level``. * ``spacecore_lowered`` — JAX-jitted SpaceCore path; on NumPy / Torch it aliases the public-none callable. @@ -262,35 +262,32 @@ def _factory( # Build a single set of operand arrays on each backend / context. rng = _rng(seed) - # Public-cheap context drives the operand dtype + asarray; reuse the + # The backend context drives the operand dtype + asarray; reuse the # same NumPy operand pool for every mode so cross-mode error - # comparisons are meaningful. - ctx_cheap = _backend_ctx(backend, check_level="cheap") - ctx_none = _backend_ctx(backend, check_level="none") - np_dtype = _np_dtype(ctx_cheap) + # comparisons are meaningful. Check levels are per-object now, so a + # single context serves every mode. + ctx = _backend_ctx(backend) + np_dtype = _np_dtype(ctx) rho_np = _make_density_batch(rng, batch, D, np_dtype) eye_d_np = (np.eye(d) / d).astype(np_dtype) # Pre-materialize backend arrays once (operand construction is *not* - # timed). - rho_bare = ctx_cheap.asarray(rho_np) - eye_d_bare = ctx_cheap.asarray(eye_d_np) - rho_none = ctx_none.asarray(rho_np) - rho_cheap = ctx_cheap.asarray(rho_np) - eye_d_none = ctx_none.asarray(eye_d_np) - eye_d_cheap = ctx_cheap.asarray(eye_d_np) + # timed). Array conversion is level-independent, so all modes share + # the same operands. + rho_arr = ctx.asarray(rho_np) + eye_d_arr = ctx.asarray(eye_d_np) # SpaceCore stacked spaces (one per check level). - herm_full_none = sc.HermitianSpace(D, ctx=ctx_none) - herm_pt_none = sc.HermitianSpace(d0, ctx=ctx_none) - stacked_full_none = sc.StackedSpace(herm_full_none, batch, ctx_none) - stacked_pt_none = sc.StackedSpace(herm_pt_none, batch, ctx_none) + herm_full_none = sc.HermitianSpace(D, ctx=ctx, check_level="none") + herm_pt_none = sc.HermitianSpace(d0, ctx=ctx, check_level="none") + stacked_full_none = sc.StackedSpace(herm_full_none, batch, ctx, check_level="none") + stacked_pt_none = sc.StackedSpace(herm_pt_none, batch, ctx, check_level="none") - herm_full_cheap = sc.HermitianSpace(D, ctx=ctx_cheap) - herm_pt_cheap = sc.HermitianSpace(d0, ctx=ctx_cheap) - stacked_full_cheap = sc.StackedSpace(herm_full_cheap, batch, ctx_cheap) - stacked_pt_cheap = sc.StackedSpace(herm_pt_cheap, batch, ctx_cheap) + herm_full_cheap = sc.HermitianSpace(D, ctx=ctx, check_level="cheap") + herm_pt_cheap = sc.HermitianSpace(d0, ctx=ctx, check_level="cheap") + stacked_full_cheap = sc.StackedSpace(herm_full_cheap, batch, ctx, check_level="cheap") + stacked_pt_cheap = sc.StackedSpace(herm_pt_cheap, batch, ctx, check_level="cheap") # Build bare callables per backend. if backend == "numpy": @@ -298,12 +295,12 @@ def bare_cb(): return _bare_pipeline_numpy(rho_np, d0, d) elif backend == "jax": jax_pipeline = _bare_pipeline_jax_factory(d0, d) - rho_jax = rho_bare # already a jax array via ctx.asarray + rho_jax = rho_arr # already a jax array via ctx.asarray def bare_cb(): return jax_pipeline(rho_jax) elif backend == "torch": - torch_pipeline = _bare_pipeline_torch_factory(d0, d, eye_d_bare) - rho_torch = rho_bare + torch_pipeline = _bare_pipeline_torch_factory(d0, d, eye_d_arr) + rho_torch = rho_arr def bare_cb(): return torch_pipeline(rho_torch) else: @@ -311,15 +308,15 @@ def bare_cb(): # SpaceCore public callables. sc_none_pipeline = _sc_pipeline_factory( - ctx_none, stacked_full_none, herm_pt_none, stacked_pt_none, d0, d, eye_d_none + ctx, stacked_full_none, herm_pt_none, stacked_pt_none, d0, d, eye_d_arr ) sc_cheap_pipeline = _sc_pipeline_factory( - ctx_cheap, stacked_full_cheap, herm_pt_cheap, stacked_pt_cheap, d0, d, eye_d_cheap + ctx, stacked_full_cheap, herm_pt_cheap, stacked_pt_cheap, d0, d, eye_d_arr ) def sc_none_cb(): - return sc_none_pipeline(rho_none) + return sc_none_pipeline(rho_arr) def sc_cheap_cb(): - return sc_cheap_pipeline(rho_cheap) + return sc_cheap_pipeline(rho_arr) # Lowered callable. if backend == "jax": @@ -329,7 +326,7 @@ def sc_cheap_cb(): # implementation through JAX tracing. sc_lowered_jit = jax.jit(sc_none_pipeline) def sc_lowered_cb(): - return sc_lowered_jit(rho_none) + return sc_lowered_jit(rho_arr) else: # For NumPy / Torch the lowered path matches public_none. sc_lowered_cb = sc_none_cb diff --git a/bench/macro/jax_full_loop.py b/bench/macro/jax_full_loop.py index d7dd25b..210339c 100644 --- a/bench/macro/jax_full_loop.py +++ b/bench/macro/jax_full_loop.py @@ -130,21 +130,20 @@ def bare_callable(): return {"x": x, "residual": residual} # --- SpaceCore eager paths (public_none / public_cheap). Same callable - # shape, separate Context per check_level. We build the full LinOp once - # in setup; the timed callable just rebuilds the carries and steps. - ctx_none = _backend_ctx("jax", check_level="none") - ctx_cheap = _backend_ctx("jax", check_level="cheap") + # shape, one shared Context with a per-object check_level. We build the + # full LinOp once in setup; the timed callable just rebuilds the carries + # and steps. + ctx = _backend_ctx("jax") - space_none = sc.DenseCoordinateSpace((n,), ctx_none) - space_cheap = sc.DenseCoordinateSpace((n,), ctx_cheap) + space_none = sc.DenseCoordinateSpace((n,), ctx, check_level="none") + space_cheap = sc.DenseCoordinateSpace((n,), ctx, check_level="cheap") - a_arr_none = ctx_none.asarray(a_np) - a_arr_cheap = ctx_cheap.asarray(a_np) - b_arr_none = ctx_none.asarray(b_np) - b_arr_cheap = ctx_cheap.asarray(b_np) + # Array conversion is level-independent — one copy serves both variants. + a_arr = ctx.asarray(a_np) + b_arr = ctx.asarray(b_np) - op_none = sc.DenseLinOp(a_arr_none, space_none, space_none, ctx_none) - op_cheap = sc.DenseLinOp(a_arr_cheap, space_cheap, space_cheap, ctx_cheap) + op_none = sc.DenseLinOp(a_arr, space_none, space_none, ctx, check_level="none") + op_cheap = sc.DenseLinOp(a_arr, space_cheap, space_cheap, ctx, check_level="cheap") def _public_cg_step(op, space, b_arr): """Eager CG: one Python-level loop with SpaceCore public ops.""" @@ -166,18 +165,16 @@ def _public_cg_step(op, space, b_arr): return {"x": x, "residual": residual} def spacecore_public_none_callable(): - return _public_cg_step(op_none, space_none, b_arr_none) + return _public_cg_step(op_none, space_none, b_arr) def spacecore_public_cheap_callable(): - return _public_cg_step(op_cheap, space_cheap, b_arr_cheap) + return _public_cg_step(op_cheap, space_cheap, b_arr) # --- SpaceCore lowered path: jax.jit a function that uses the SC LinOp's # apply inside a jax.lax.scan. The trace lowers ``op.apply`` to the same # underlying ``a @ x`` JAX op without going through any Python-level # SpaceCore validation per iteration. - a_arr_lowered = ctx_none.asarray(a_np) - b_arr_lowered = ctx_none.asarray(b_np) - op_lowered = sc.DenseLinOp(a_arr_lowered, space_none, space_none, ctx_none) + op_lowered = sc.DenseLinOp(a_arr, space_none, space_none, ctx, check_level="none") def _lowered_cg_loop(b): x0 = space_none.zeros() @@ -204,7 +201,7 @@ def body(carry, _): _lowered_cg_jit = jax.jit(_lowered_cg_loop) def spacecore_lowered_callable(): - x, residual = _lowered_cg_jit(b_arr_lowered) + x, residual = _lowered_cg_jit(b_arr) return {"x": x, "residual": residual} def reference_metric_extractor(result: Any) -> dict[str, float]: @@ -301,21 +298,23 @@ def bare_callable(): return {"x": x, "objective": objective} # --- SpaceCore eager public paths. - ctx_none = _backend_ctx("jax", check_level="none") - ctx_cheap = _backend_ctx("jax", check_level="cheap") + ctx = _backend_ctx("jax") - domain_none = sc.DenseCoordinateSpace((n,), ctx_none) - codomain_none = sc.DenseCoordinateSpace((m,), ctx_none) - domain_cheap = sc.DenseCoordinateSpace((n,), ctx_cheap) - codomain_cheap = sc.DenseCoordinateSpace((m,), ctx_cheap) + domain_none = sc.DenseCoordinateSpace((n,), ctx, check_level="none") + codomain_none = sc.DenseCoordinateSpace((m,), ctx, check_level="none") + domain_cheap = sc.DenseCoordinateSpace((n,), ctx, check_level="cheap") + codomain_cheap = sc.DenseCoordinateSpace((m,), ctx, check_level="cheap") - a_arr_none = ctx_none.asarray(a_np) - a_arr_cheap = ctx_cheap.asarray(a_np) - b_arr_none = ctx_none.asarray(b_np) - b_arr_cheap = ctx_cheap.asarray(b_np) + # Array conversion is level-independent — one copy serves both variants. + a_arr = ctx.asarray(a_np) + b_arr = ctx.asarray(b_np) - op_none = sc.DenseLinOp(a_arr_none, domain_none, codomain_none, ctx_none) - op_cheap = sc.DenseLinOp(a_arr_cheap, domain_cheap, codomain_cheap, ctx_cheap) + op_none = sc.DenseLinOp( + a_arr, domain_none, codomain_none, ctx, check_level="none" + ) + op_cheap = sc.DenseLinOp( + a_arr, domain_cheap, codomain_cheap, ctx, check_level="cheap" + ) def _public_pdhg_step(op, domain, codomain, b_arr): x = domain.zeros() @@ -337,16 +336,16 @@ def _public_pdhg_step(op, domain, codomain, b_arr): return {"x": x, "objective": objective} def spacecore_public_none_callable(): - return _public_pdhg_step(op_none, domain_none, codomain_none, b_arr_none) + return _public_pdhg_step(op_none, domain_none, codomain_none, b_arr) def spacecore_public_cheap_callable(): - return _public_pdhg_step(op_cheap, domain_cheap, codomain_cheap, b_arr_cheap) + return _public_pdhg_step(op_cheap, domain_cheap, codomain_cheap, b_arr) # --- SpaceCore lowered path: jit a function that uses op.apply/op.rapply # inside a jax.lax.scan body. - a_arr_lowered = ctx_none.asarray(a_np) - b_arr_lowered = ctx_none.asarray(b_np) - op_lowered = sc.DenseLinOp(a_arr_lowered, domain_none, codomain_none, ctx_none) + op_lowered = sc.DenseLinOp( + a_arr, domain_none, codomain_none, ctx, check_level="none" + ) def _lowered_pdhg_loop(b): x0 = domain_none.zeros() @@ -373,7 +372,7 @@ def body(carry, _): _lowered_pdhg_jit = jax.jit(_lowered_pdhg_loop) def spacecore_lowered_callable(): - x, objective = _lowered_pdhg_jit(b_arr_lowered) + x, objective = _lowered_pdhg_jit(b_arr) return {"x": x, "objective": objective} def reference_metric_extractor(result: Any) -> dict[str, float]: diff --git a/bench/macro/operator_stress.py b/bench/macro/operator_stress.py index b7ec6d7..795cddb 100644 --- a/bench/macro/operator_stress.py +++ b/bench/macro/operator_stress.py @@ -175,13 +175,12 @@ def _factory( chain_depth = int(size_params["chain_depth"]) rng = _rng(seed) - # SpaceCore contexts and space — built once per payload. - ctx_none = _backend_ctx(backend, check_level="none") - ctx_cheap = _backend_ctx(backend, check_level="cheap") - np_dtype = _np_dtype(ctx_none) + # SpaceCore context and spaces — built once per payload. + ctx = _backend_ctx(backend) + np_dtype = _np_dtype(ctx) - space_none = sc.DenseCoordinateSpace((d,), ctx_none) - space_cheap = sc.DenseCoordinateSpace((d,), ctx_cheap) + space_none = sc.DenseCoordinateSpace((d,), ctx, check_level="none") + space_cheap = sc.DenseCoordinateSpace((d,), ctx, check_level="cheap") # Pre-generate every matrix and vector in NumPy. A_np = [ @@ -200,25 +199,21 @@ def _factory( y_np = np.asarray(rng.standard_normal(d), dtype=np_dtype) # Convert per-backend operand arrays once. - A_none = [ctx_none.asarray(M) for M in A_np] - B_none = [ctx_none.asarray(M) for M in B_np] - A_cheap = [ctx_cheap.asarray(M) for M in A_np] - B_cheap = [ctx_cheap.asarray(M) for M in B_np] - x_none = ctx_none.asarray(x_np) - y_none = ctx_none.asarray(y_np) - x_cheap = ctx_cheap.asarray(x_np) - y_cheap = ctx_cheap.asarray(y_np) + A_arr = [ctx.asarray(M) for M in A_np] + B_arr = [ctx.asarray(M) for M in B_np] + x_arr = ctx.asarray(x_np) + y_arr = ctx.asarray(y_np) # Bare-mode operands — backend-native, no SpaceCore objects. - x_bare = _to_backend(backend, x_np, ctx_none) - y_bare = _to_backend(backend, y_np, ctx_none) - A_dense_bare = _build_dense_matrix(backend, ctx_none, A_np, B_np, alphas) + x_bare = _to_backend(backend, x_np, ctx) + y_bare = _to_backend(backend, y_np, ctx) + A_dense_bare = _build_dense_matrix(backend, ctx, A_np, B_np, alphas) A_dense_bare_H = _bare_conj_transpose(backend, A_dense_bare) matmul = _bare_matmul(backend) # SpaceCore expression for the public-API paths. - expr_none = _build_sc_expression(ctx_none, space_none, A_none, B_none, alphas) - expr_cheap = _build_sc_expression(ctx_cheap, space_cheap, A_cheap, B_cheap, alphas) + expr_none = _build_sc_expression(ctx, space_none, A_arr, B_arr, alphas) + expr_cheap = _build_sc_expression(ctx, space_cheap, A_arr, B_arr, alphas) # Lowered path: collapse the expression to a single dense matrix once. A_dense_lowered = expr_none.to_dense() @@ -248,7 +243,7 @@ def bare_call() -> dict[str, Any]: "backend": backend, } - # ----- spacecore public (check_level varies via context) ----- + # ----- spacecore public (check_level varies per object) ----- def _sc_public_call(expr: sc.LinOp, x_in: Any, y_in: Any) -> dict[str, Any]: apply_out = x_in rapply_out = y_in @@ -265,10 +260,10 @@ def _sc_public_call(expr: sc.LinOp, x_in: Any, y_in: Any) -> dict[str, Any]: } def sc_public_none_call() -> dict[str, Any]: - return _sc_public_call(expr_none, x_none, y_none) + return _sc_public_call(expr_none, x_arr, y_arr) def sc_public_cheap_call() -> dict[str, Any]: - return _sc_public_call(expr_cheap, x_cheap, y_cheap) + return _sc_public_call(expr_cheap, x_arr, y_arr) # ----- spacecore lowered: precomputed dense matrix, bare matmuls ----- def sc_lowered_call() -> dict[str, Any]: diff --git a/bench/macro/pdhg.py b/bench/macro/pdhg.py index d2ef3c5..ffc39f8 100644 --- a/bench/macro/pdhg.py +++ b/bench/macro/pdhg.py @@ -137,18 +137,16 @@ def _build_numpy_payload( payload_meta: dict[str, Any], ) -> MacroPayload: m, n = a_np.shape - ctx_none = _backend_ctx("numpy", check_level="none") - ctx_cheap = _backend_ctx("numpy", check_level="cheap") - dom_none = sc.DenseCoordinateSpace((n,), ctx_none) - cod_none = sc.DenseCoordinateSpace((m,), ctx_none) - dom_cheap = sc.DenseCoordinateSpace((n,), ctx_cheap) - cod_cheap = sc.DenseCoordinateSpace((m,), ctx_cheap) - - a_arr_none = ctx_none.asarray(a_np) - a_arr_cheap = ctx_cheap.asarray(a_np) - ctx_none.asarray(b_np) - op_none = sc.DenseLinOp(a_arr_none, dom_none, cod_none, ctx_none) - op_cheap = sc.DenseLinOp(a_arr_cheap, dom_cheap, cod_cheap, ctx_cheap) + ctx = _backend_ctx("numpy") + dom_none = sc.DenseCoordinateSpace((n,), ctx, check_level="none") + cod_none = sc.DenseCoordinateSpace((m,), ctx, check_level="none") + dom_cheap = sc.DenseCoordinateSpace((n,), ctx, check_level="cheap") + cod_cheap = sc.DenseCoordinateSpace((m,), ctx, check_level="cheap") + + # Array conversion is level-independent — one copy serves both variants. + a_arr = ctx.asarray(a_np) + op_none = sc.DenseLinOp(a_arr, dom_none, cod_none, ctx, check_level="none") + op_cheap = sc.DenseLinOp(a_arr, dom_cheap, cod_cheap, ctx, check_level="cheap") a_local = a_np b_local = b_np @@ -234,24 +232,23 @@ def _build_jax_payload( import jax.numpy as jnp m, n = a_np.shape - ctx_none = _backend_ctx("jax", check_level="none") - ctx_cheap = _backend_ctx("jax", check_level="cheap") - dom_none = sc.DenseCoordinateSpace((n,), ctx_none) - cod_none = sc.DenseCoordinateSpace((m,), ctx_none) - dom_cheap = sc.DenseCoordinateSpace((n,), ctx_cheap) - cod_cheap = sc.DenseCoordinateSpace((m,), ctx_cheap) - - np_dtype = _np_dtype(ctx_none) + ctx = _backend_ctx("jax") + dom_none = sc.DenseCoordinateSpace((n,), ctx, check_level="none") + cod_none = sc.DenseCoordinateSpace((m,), ctx, check_level="none") + dom_cheap = sc.DenseCoordinateSpace((n,), ctx, check_level="cheap") + cod_cheap = sc.DenseCoordinateSpace((m,), ctx, check_level="cheap") + + np_dtype = _np_dtype(ctx) a_typed = np.asarray(a_np, dtype=np_dtype) b_typed = np.asarray(b_np, dtype=np_dtype) - a_jax_none = ctx_none.asarray(a_typed) - a_jax_cheap = ctx_cheap.asarray(a_typed) - b_jax = ctx_none.asarray(b_typed) - op_none = sc.DenseLinOp(a_jax_none, dom_none, cod_none, ctx_none) - op_cheap = sc.DenseLinOp(a_jax_cheap, dom_cheap, cod_cheap, ctx_cheap) + # Array conversion is level-independent — one copy serves both variants. + a_jax = ctx.asarray(a_typed) + b_jax = ctx.asarray(b_typed) + op_none = sc.DenseLinOp(a_jax, dom_none, cod_none, ctx, check_level="none") + op_cheap = sc.DenseLinOp(a_jax, dom_cheap, cod_cheap, ctx, check_level="cheap") - a_local = a_jax_none + a_local = a_jax b_local = b_jax one_plus_sigma = 1.0 + sigma zeros_x = jnp.zeros((n,), dtype=a_local.dtype) @@ -362,24 +359,23 @@ def _build_torch_payload( import torch m, n = a_np.shape - ctx_none = _backend_ctx("torch", check_level="none") - ctx_cheap = _backend_ctx("torch", check_level="cheap") - dom_none = sc.DenseCoordinateSpace((n,), ctx_none) - cod_none = sc.DenseCoordinateSpace((m,), ctx_none) - dom_cheap = sc.DenseCoordinateSpace((n,), ctx_cheap) - cod_cheap = sc.DenseCoordinateSpace((m,), ctx_cheap) - - np_dtype = _np_dtype(ctx_none) + ctx = _backend_ctx("torch") + dom_none = sc.DenseCoordinateSpace((n,), ctx, check_level="none") + cod_none = sc.DenseCoordinateSpace((m,), ctx, check_level="none") + dom_cheap = sc.DenseCoordinateSpace((n,), ctx, check_level="cheap") + cod_cheap = sc.DenseCoordinateSpace((m,), ctx, check_level="cheap") + + np_dtype = _np_dtype(ctx) a_typed = np.asarray(a_np, dtype=np_dtype) b_typed = np.asarray(b_np, dtype=np_dtype) - a_torch_none = ctx_none.asarray(a_typed) - a_torch_cheap = ctx_cheap.asarray(a_typed) - b_torch = ctx_none.asarray(b_typed) - op_none = sc.DenseLinOp(a_torch_none, dom_none, cod_none, ctx_none) - op_cheap = sc.DenseLinOp(a_torch_cheap, dom_cheap, cod_cheap, ctx_cheap) + # Array conversion is level-independent — one copy serves both variants. + a_torch = ctx.asarray(a_typed) + b_torch = ctx.asarray(b_typed) + op_none = sc.DenseLinOp(a_torch, dom_none, cod_none, ctx, check_level="none") + op_cheap = sc.DenseLinOp(a_torch, dom_cheap, cod_cheap, ctx, check_level="cheap") - a_local = a_torch_none + a_local = a_torch b_local = b_torch one_plus_sigma = 1.0 + sigma torch_dtype = a_local.dtype diff --git a/bench/macro/power_lanczos.py b/bench/macro/power_lanczos.py index 3583b1e..cc1c33b 100644 --- a/bench/macro/power_lanczos.py +++ b/bench/macro/power_lanczos.py @@ -30,7 +30,7 @@ The operator is built once per ``(backend, size, seed)`` triple in the factory body. Timed callables only invoke the solver — they do not allocate the potential or convert dtypes. For the JAX paths, the public -context is built with the matching ``check_level`` so the runner +space is built with the matching ``check_level`` so the runner honours the cheap-check overhead. """ from __future__ import annotations @@ -284,7 +284,7 @@ def _bare_lanczos( def _build_sc_callable_power( backend: str, - ctx_check: str, + check_level: str, v_np: np.ndarray, d: int, K: int, @@ -292,8 +292,8 @@ def _build_sc_callable_power( jit: bool, ) -> Callable[[], Any]: """Build the SpaceCore power-iteration callable for one check level.""" - ctx = _backend_ctx(backend, check_level=ctx_check) - space = sc.DenseCoordinateSpace((d,), ctx) + ctx = _backend_ctx(backend) + space = sc.DenseCoordinateSpace((d,), ctx, check_level=check_level) v_ctx = ctx.asarray(v_np) x0_ctx = ctx.asarray(x0_np) @@ -321,7 +321,7 @@ def call_jit() -> Any: def _build_sc_callable_lanczos( backend: str, - ctx_check: str, + check_level: str, v_np: np.ndarray, d: int, K: int, @@ -329,8 +329,8 @@ def _build_sc_callable_lanczos( jit: bool, ) -> Callable[[], Any]: """Build the SpaceCore lanczos_smallest callable for one check level.""" - ctx = _backend_ctx(backend, check_level=ctx_check) - space = sc.DenseCoordinateSpace((d,), ctx) + ctx = _backend_ctx(backend) + space = sc.DenseCoordinateSpace((d,), ctx, check_level=check_level) v_ctx = ctx.asarray(v_np) x0_ctx = ctx.asarray(x0_np) @@ -386,9 +386,9 @@ def _factory_power( d = int(size_params["d"]) K = int(size_params["K"]) - # Use the public-none context to derive a canonical dtype for the - # operator data; the cheap-checked context shares the same dtype. - base_ctx = _backend_ctx(backend, check_level="none") + # Derive a canonical dtype for the operator data; both check levels + # share the same context and dtype. + base_ctx = _backend_ctx(backend) np_dtype = _np_dtype(base_ctx) v_np = _make_potential(d, seed, np_dtype) @@ -450,7 +450,7 @@ def _factory_lanczos( d = int(size_params["d"]) K = int(size_params["K"]) - base_ctx = _backend_ctx(backend, check_level="none") + base_ctx = _backend_ctx(backend) np_dtype = _np_dtype(base_ctx) v_np = _make_potential(d, seed, np_dtype) diff --git a/bench/macro/qot_barycenter.py b/bench/macro/qot_barycenter.py index 60ff685..debebee 100644 --- a/bench/macro/qot_barycenter.py +++ b/bench/macro/qot_barycenter.py @@ -244,10 +244,9 @@ def _factory( epsilon = float(size_params["epsilon"]) D = d0 * d - # Resolve contexts up front (one per mode that uses SpaceCore). - ctx_none = _backend_ctx(backend, check_level="none") - ctx_cheap = _backend_ctx(backend, check_level="cheap") - np_dtype = _np_dtype(ctx_none) + # Resolve the backend context up front (check levels are per-object now). + ctx = _backend_ctx(backend) + np_dtype = _np_dtype(ctx) problem = _build_problem_np(S, d0, d, seed, np_dtype) # Pre-build the per-state Kronecker-sum tensor (U kron I_d + I_d0 kron V). @@ -306,23 +305,24 @@ def bare_callable(): else: raise ValueError(f"unknown backend {backend!r}") - # SpaceCore public-API path: build a HermitianSpace per check-level - # and convert operands into the matching context. - def _build_sc_callable(ctx: sc.Context) -> Callable[[], Any]: - space = sc.HermitianSpace(D, ctx=ctx) - C_sc = ctx.asarray(problem["C"]) - kron_sc = ctx.asarray(kron_np) - # Pre-asarray each target row so the timed step is allocation-free. - targets_sc = [ctx.asarray(problem["targets"][s]) for s in range(S)] - ops = ctx.ops + # SpaceCore public-API path: build a HermitianSpace per check-level. + # Operand conversion is level-independent, so it happens once here. + C_sc = ctx.asarray(problem["C"]) + kron_sc = ctx.asarray(kron_np) + # Pre-asarray each target row so the timed step is allocation-free. + targets_sc = [ctx.asarray(problem["targets"][s]) for s in range(S)] + ops = ctx.ops + + def _build_sc_callable(check_level: str) -> Callable[[], Any]: + space = sc.HermitianSpace(D, ctx=ctx, check_level=check_level) def call(): return _sc_step(space, C_sc, kron_sc, targets_sc, epsilon, d0, d, ops) return call - public_none_callable = _build_sc_callable(ctx_none) - public_cheap_callable = _build_sc_callable(ctx_cheap) + public_none_callable = _build_sc_callable("none") + public_cheap_callable = _build_sc_callable("cheap") # Lowered path. On JAX this is the jit-compiled bare step. Elsewhere # we alias the public_none callable, which is the contract default. diff --git a/docs/source/api/context.rst b/docs/source/api/context.rst index e5d1267..443c231 100644 --- a/docs/source/api/context.rst +++ b/docs/source/api/context.rst @@ -1,12 +1,16 @@ Context API =========== -``Context`` packages backend operations, default dtype, and runtime validation -policy. Spaces and operators store a normalized context and use it for array -construction and checks. +``Context`` packages backend operations and the default dtype. Spaces, +operators, and functionals store a normalized context and use it for array +construction and conversion. -Use ``check_level="none"``, ``"cheap"``, ``"standard"``, or ``"strict"``. -The exported ``spacecore.CheckLevel`` literal is available for annotations. +Validation policy is *not* part of a context: ``check_level`` belongs to the +bound object, so two objects on the same backend and dtype may validate at +different strictness. Use ``check_level="none"``, ``"cheap"``, ``"standard"``, +or ``"strict"`` on the object, or move the ambient default with +``set_check_level`` / ``use_check_level``. The exported ``spacecore.CheckLevel`` +literal is available for annotations. See :doc:`../design/checking_policy`. Context ------- @@ -14,9 +18,9 @@ Context .. autosummary:: :nosignatures: - spacecore.backend.Context + spacecore.Context -.. autoclass:: spacecore.backend.Context +.. autoclass:: spacecore.Context :members: Context helpers @@ -27,18 +31,30 @@ Context helpers spacecore.get_context spacecore.set_context + spacecore.use_context + spacecore.get_check_level + spacecore.set_check_level + spacecore.use_check_level spacecore.normalize_context spacecore.normalize_ops spacecore.resolve_context_priority spacecore.register_ops -* ``get_context`` and ``set_context`` manage the global default context. +* ``get_context`` and ``set_context`` manage the global default context, and + ``use_context`` overrides it for a block, scoped to the current thread or + async task. +* ``get_check_level`` / ``set_check_level`` / ``use_check_level`` do the same for + the ambient validation level applied to newly constructed bound objects. * ``normalize_context`` turns backend names, families, concrete contexts, or ``None`` into a context. * ``resolve_context_priority`` chooses a common context for constructors. * ``register_ops`` adds a custom backend implementation. .. autofunction:: spacecore.get_context .. autofunction:: spacecore.set_context +.. autofunction:: spacecore.use_context +.. autofunction:: spacecore.get_check_level +.. autofunction:: spacecore.set_check_level +.. autofunction:: spacecore.use_check_level .. autofunction:: spacecore.normalize_context .. autofunction:: spacecore.normalize_ops .. autofunction:: spacecore.resolve_context_priority diff --git a/docs/source/api/functionals.rst b/docs/source/api/functionals.rst index 21933e0..193a016 100644 --- a/docs/source/api/functionals.rst +++ b/docs/source/api/functionals.rst @@ -20,6 +20,69 @@ Base and composition * ``ComposedFunctional`` represents pullback ``f o A`` for a linear operator ``A``. * ``make_functional_composed`` constructs the same pullback with simplifications. +Algebra +------- + +Lazy nodes for combining functionals. The operator overloads on ``Functional`` +(``a * F``, ``F + G``, ``F - G``, ``-F``, ``F * G``) delegate to the ``make_*`` +factories, which apply local, *structural* canonicalization — they read node +types, never values. + +.. autosummary:: + :nosignatures: + + spacecore.functional.ScaledFunctional + spacecore.functional.SumFunctional + spacecore.functional.ShiftedFunctional + spacecore.functional.ZeroFunctional + spacecore.functional.ConstantFunctional + spacecore.functional.ProductFunctional + spacecore.functional.make_scaled_functional + spacecore.functional.make_functional_sum + spacecore.functional.make_shifted_functional + spacecore.functional.make_constant_functional + spacecore.functional.make_functional_product + +* ``ScaledFunctional`` is ``a * F``; the Riesz gradient scales by ``conj(a)``. +* ``SumFunctional`` is ``F_1 + ... + F_n`` on a shared domain. +* ``ShiftedFunctional`` is the affine shift ``F + c`` (gradient unchanged). +* ``ZeroFunctional`` is the additive identity, recognized by the canonicalizers. +* ``ConstantFunctional`` is ``x -> c``, the embedding of a scalar into the + algebra; ``make_constant_functional`` collapses ``c = 0`` to ``ZeroFunctional``. +* ``ProductFunctional`` is the pointwise product ``F(x) * G(x)``, with the + product-rule Riesz gradient + :math:`\overline{G(x)}\, \nabla F(x) + \overline{F(x)}\, \nabla G(x)`. Because + functionals are scalar-valued, scaling is the constant-factor case of a + product: ``make_functional_product`` folds a ``ConstantFunctional`` factor back + into a ``ScaledFunctional``, and a ``ZeroFunctional`` factor to zero. + +Operator families (``F · A``) +----------------------------- + +.. autosummary:: + :nosignatures: + + spacecore.OperatorFamily + spacecore.FunctionalScaledOperator + spacecore.make_functional_scaled_operator + +``F * A`` for a ``Functional`` and a ``LinOp`` is the functional-weighted map +:math:`m(x) = F(x)\,Ax`. This is **not** linear — both the scale and the +direction move with ``x`` — so it is *not* a ``LinOp``; it is an +``OperatorFamily``, a point-indexed family :math:`x \mapsto A_x`. + +Each point carries two different linear operators, and they are not the same: + +* ``m.at(x)`` — the **frozen member** ``F(x) · A``, an ordinary ``LinOp`` that + composes, sums, and has an adjoint. Freezing the point is what recovers + linearity. +* ``m.linearize_at(x)`` — the **derivative** :math:`Dm(x)`, for Newton-type + steps. They coincide only when ``F`` is constant; otherwise they differ by a + rank-one term. + +A ``ConstantFunctional`` weight collapses to an ordinary ``ScaledLinOp``, so the +linear case is never forced through the non-linear type. + Linear functionals ------------------ @@ -59,8 +122,12 @@ metric (Riesz) gradients under the domain geometry. spacecore.functional.SquaredL2NormFunctional spacecore.functional.LpNormFunctional spacecore.functional.L1NormFunctional - spacecore.functional.SpectralLpNormFunctional + spacecore.functional.SpectralFunctional + spacecore.functional.spectralize + spacecore.functional.eigenvalue_space spacecore.functional.NuclearNormFunctional + spacecore.functional.RealifiedFunctional + spacecore.functional.realify spacecore.functional.NegativeEntropyFunctional spacecore.functional.KLDivergenceFunctional spacecore.functional.HuberFunctional @@ -68,8 +135,14 @@ metric (Riesz) gradients under the domain geometry. * ``least_squares`` builds the ``scale ||A x - b||^2`` objective as a ``LinOpQuadraticForm``. * ``SquaredL2NormFunctional`` is ``1/2 ||x||_X^2`` (gradient ``x``, clean shrinkage prox). * ``LpNormFunctional`` / ``L1NormFunctional`` are coordinate ``p``-norms. -* ``SpectralLpNormFunctional`` / ``NuclearNormFunctional`` are the Schatten ``p``-norm - and nuclear norm of a Jordan spectrum (e.g. Hermitian eigenvalues). +* ``SpectralFunctional`` / ``spectralize`` lift **any** symmetric coordinate + functional onto a Jordan spectrum (Lewis): the Schatten ``p``-norm is + ``spectralize(X, lambda s: LpNormFunctional(s, p))``, the von Neumann entropy + is ``spectralize(X, NegativeEntropyFunctional)``. ``eigenvalue_space`` builds + the real space the spectrum lives in. +* ``NuclearNormFunctional`` is the named Schatten-1 case. +* ``RealifiedFunctional`` / ``realify`` present a complex-domain functional over + stacked real coordinates, for real-only optimizers. * ``NegativeEntropyFunctional`` and ``KLDivergenceFunctional`` are the entropy objectives. * ``HuberFunctional`` is the separable Huber loss. @@ -137,11 +210,20 @@ Autodoc .. autofunction:: spacecore.functional.L1NormFunctional -.. autoclass:: spacecore.functional.SpectralLpNormFunctional +.. autoclass:: spacecore.functional.SpectralFunctional :members: +.. autofunction:: spacecore.functional.spectralize + +.. autofunction:: spacecore.functional.eigenvalue_space + .. autofunction:: spacecore.functional.NuclearNormFunctional +.. autoclass:: spacecore.functional.RealifiedFunctional + :members: + +.. autofunction:: spacecore.functional.realify + .. autoclass:: spacecore.functional.NegativeEntropyFunctional :members: diff --git a/docs/source/design/checking_policy.rst b/docs/source/design/checking_policy.rst index 58b7bff..cb85282 100644 --- a/docs/source/design/checking_policy.rst +++ b/docs/source/design/checking_policy.rst @@ -1,18 +1,23 @@ Checking policy =============== -SpaceCore uses ``Context.check_level`` as its public runtime-validation policy. -The public type is ``spacecore.CheckLevel``, a literal type with four ordered -values: ``"none"``, ``"cheap"``, ``"standard"``, and ``"strict"``. A literal -keeps context construction simple and makes invalid spellings visible to static +SpaceCore uses ``ContextBound.check_level`` as its public runtime-validation +policy. The public type is ``spacecore.CheckLevel``, a literal type with four +ordered values: ``"none"``, ``"cheap"``, ``"standard"``, and ``"strict"``. A +literal keeps construction simple and makes invalid spellings visible to static type checkers without introducing a separate policy object. +The level belongs to the *bound object* — space, operator, functional — and not +to the :class:`~spacecore.Context`. A context fixes backend ops and dtype; two +objects sharing both may still need different strictness, so the policy is set +per object, with an ambient default for objects that do not name one. + .. code-block:: python import spacecore as sc - ctx = sc.Context(sc.NumpyOps(), dtype="float64", check_level="standard") - X = sc.DenseCoordinateSpace((3,), ctx=ctx) + ctx = sc.Context(sc.NumpyOps(), dtype="float64") + X = sc.DenseCoordinateSpace((3,), ctx=ctx, check_level="standard") x = X.ctx.asarray([1.0, 2.0, 3.0]) X.check_member(x) @@ -53,9 +58,11 @@ Choosing a level * Performance-sensitive trusted code: use ``"cheap"`` or ``"none"``. * User-facing libraries: usually use ``"standard"``. -The process-wide default context remains ``"none"`` for compatibility. A -direct ``Context(...)`` defaults to ``"standard"``, matching the previous -direct-constructor default. +An object constructed without an explicit ``check_level`` takes the ambient +default, read with :func:`spacecore.get_check_level`. Move that default +process-wide with :func:`spacecore.set_check_level`, or for a block with +:func:`spacecore.use_check_level`, which is scoped to the current thread or +async task. Where checks run ---------------- @@ -71,25 +78,26 @@ When a context is inferred from several source objects, SpaceCore selects the least expensive source level. For example, combining ``"strict"`` and ``"cheap"`` contexts produces a ``"cheap"`` inferred policy. -Migration from ``enable_checks`` --------------------------------- - -``enable_checks`` remains as a deprecated compatibility keyword: +Migration from context-carried levels +------------------------------------- -* ``enable_checks=True`` maps to ``check_level="standard"``; -* ``enable_checks=False`` maps to ``check_level="none"``; -* passing both keywords raises ``TypeError``. +Before 0.4.3 the level rode on the ``Context``, and the long-deprecated +``enable_checks=`` Boolean was still accepted there. Both are gone: ``Context`` +is now exactly ``(ops, dtype)``, and passing either keyword to it raises +``TypeError``. .. code-block:: python - # New spelling - ctx = sc.Context(sc.NumpyOps(), check_level="standard") + # 0.4.3 and later: the level is set on the object + ctx = sc.Context(sc.NumpyOps()) + X = sc.DenseCoordinateSpace((3,), ctx=ctx, check_level="standard") - # Deprecated equivalent - legacy_ctx = sc.Context(sc.NumpyOps(), enable_checks=True) + # ...or moved for everything constructed in a block + with sc.use_check_level("strict"): + Y = sc.DenseCoordinateSpace((3,), ctx=ctx) -``ctx.enable_checks`` remains a deprecated Boolean view and is true for -``cheap``, ``standard``, and ``strict`` contexts. +The level is always one of the four literals; the Boolean spelling is not +accepted in its place. Read an object's level back from ``obj.check_level``. Implementation convention ------------------------- diff --git a/docs/source/release_notes.rst b/docs/source/release_notes.rst index 0c76c5d..c907561 100644 --- a/docs/source/release_notes.rst +++ b/docs/source/release_notes.rst @@ -1,6 +1,97 @@ Release notes ============= +Version 0.4.3 +------------- + +Released 2026-09-12. SpaceCore 0.4.3 moves validation policy off the context and +onto the bound object, makes pytree registration backend-neutral, and extends the +functional algebra with products, constants, and the operator-family type that +``Functional * LinOp`` needs. + +Removed +~~~~~~~ + +* ``jax_pytree_class`` is removed from the public API (breaking). Pytree + registration is now expressed through the backend-neutral ``PyTreeNode`` + mixin, so a JAX-named decorator no longer sits in a backend-agnostic surface. +* ``SpectralLpNormFunctional`` is removed (breaking). ``spectralize`` lifts any + coordinate functional to the spectrum of a Jordan-algebra space, which + subsumes the per-formula spectral class. + +Added +~~~~~ + +* Backend-neutral pytree registration. ``PyTreeNode`` carries the capability, + and ``TorchOps.install_pytree_protocol`` makes Torch a transform-capable + backend alongside JAX. +* ``SpectralFunctional``, ``spectralize``, and ``eigenvalue_space`` lift any + coordinate functional onto a spectrum; ``RealifiedFunctional`` and ``realify`` + view a complex-domain functional over its real coordinates, so optimizers that + assume a real vector space can consume it without hand-written Wirtinger + bookkeeping. +* ``ProductFunctional`` / ``make_functional_product`` and ``ConstantFunctional`` + / ``make_constant_functional`` complete the functional algebra with a + pointwise product carrying the product-rule Riesz gradient, and with the + embedding of a scalar as a functional. +* ``spacecore.opfamily``: ``OperatorFamily`` and ``FunctionalScaledOperator`` + model ``F * A``, the functional-weighted map ``m(x) = F(x) A x``. That map is + not linear, so it is deliberately not a ``LinOp``; ``m.at(x)`` freezes a member + and ``m.linearize_at(x)`` gives the derivative for Newton-type steps. +* ``Space.scalar_field`` / ``declared_scalar_field`` / ``check_scalar`` declare + the field a space is closed under rather than inferring it from the dtype, and + ``checked_method`` gains ``out_scalar`` / ``out_batched_scalar`` so a + functional's scalar codomain is checked by the same decorator that guards + every ``LinOp`` output. +* ``OpsRegistry`` extracts the backend registry out of ``Contextual``, and + ``BackendOps.complex_dtype`` completes the pair with ``real_dtype``. + +Changed +~~~~~~~ + +* ``check_level`` is a property of the bound object, not of ``Context`` + (breaking). ``Context(ops, dtype=..., check_level=...)`` and the deprecated + ``enable_checks=`` argument are gone -- a ``Context`` is now exactly + ``(ops, dtype)``. Pass ``check_level=`` to a space, operator, or functional, or + move the ambient default with ``set_check_level`` / ``use_check_level``. See + :doc:`design/checking_policy`. +* ``spacecore._contextual`` is now the public ``spacecore.contextual`` + (breaking). +* Ambient context and check level are scoped with ``contextvars``, so an + override is visible to the current thread and async task rather than to the + whole process. +* ``Functional.__mul__`` / ``__rmul__`` dispatch on the operand type, and + functional outputs are checked as scalars at ``standard`` and above. +* ``available_ops()`` is memoized and returns a tuple. + +Fixed +~~~~~ + +* ``ComposedFunctional`` now has a gradient: ``F.compose(A).grad(x)`` previously + raised. +* A present-but-broken optional backend no longer aborts ``import spacecore``, + and one broken backend warns once rather than once per discovery call. +* ``scalar_eq`` no longer swallows every exception; it narrows to ``TypeError`` + so a genuinely broken ``__eq__`` propagates. + +Version 0.4.2 +------------- + +Released 2026-07-01. SpaceCore 0.4.2 adds the Jordan spectral primitives, a lazy +functional algebra, and a compiled convergence-aware optimizer driver. + +Added +~~~~~ + +* Jordan spectral primitives on every Jordan-algebra space: ``trace``, + ``determinant``, and ``unit``, with ``TraceFunctional`` over them. +* ``Functional.value_and_grad`` evaluates a functional and its Riesz gradient in + one pass. +* Structure-preserving tree spectra, so a ``TreeSpace`` spectral decomposition + keeps the tree shape. +* ``minimize_optax``, a compiled convergence-aware driver for the Optax + optimizers. + Version 0.4.1 ------------- diff --git a/pyproject.toml b/pyproject.toml index 9ecb753..c64b718 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -106,6 +106,13 @@ target-version = "py311" extend-exclude = ["*.ipynb"] [tool.ruff.lint] +# Name the enabled rule set explicitly rather than inheriting ruff's implicit +# default. That default has grown between releases -- 0.16 folded I, UP, B, S, +# SIM, PL and RUF into it -- so an unpinned `pip install ruff` in CI silently +# changed which rules were enforced, turning a green tree red without a single +# source edit. Listing the rules keeps `ruff check .` reproducible across ruff +# versions; widening the set is then a deliberate commit, not a release surprise. +select = ["E4", "E7", "E9", "F"] ignore = [ "D100", "D104", diff --git a/scripts/jit_audit.py b/scripts/jit_audit.py index b880e04..df41be4 100644 --- a/scripts/jit_audit.py +++ b/scripts/jit_audit.py @@ -16,7 +16,7 @@ def _ctx(): import spacecore as sc dtype = np.float64 if jax.config.read("jax_enable_x64") else np.float32 - return sc.Context(sc.JaxOps(), dtype=dtype, enable_checks=False) + return sc.Context(sc.JaxOps(), dtype=dtype) def _spd_operator(n: int): @@ -125,6 +125,12 @@ def main() -> None: if args.log_compiles: jax.config.update("jax_log_compiles", True) + # Validation policy lives on context-bound objects, not on the Context, so + # the audit lowers the ambient default instead of disabling checks per + # object: argument validation would otherwise add non-traced Python work + # that muddies the retrace counts this script measures. + sc.set_check_level("none") + A2 = _spd_operator(2) A3 = _spd_operator(3) x2a, x2b = _same_shape_inputs(A2) diff --git a/spacecore/__init__.py b/spacecore/__init__.py index eb5efee..e65d33e 100644 --- a/spacecore/__init__.py +++ b/spacecore/__init__.py @@ -2,20 +2,9 @@ from ._version import __version__ -from .backend import CHECK_LEVELS, CheckLevel, Context, BackendOps, NumpyOps, jax_pytree_class - -try: - from .backend import JaxOps as JaxOps -except ImportError: - pass -try: - from .backend import CuPyOps as CuPyOps -except ImportError: - pass -try: - from .backend import TorchOps as TorchOps -except ImportError: - pass +from .contextual import Context +from .backend import CHECK_LEVELS, CheckLevel, BackendOps, NumpyOps +from .backend._optional import available_ops as _available_ops from .linop import ( BlockDiagonalLinOp, BlockMatrixLinOp, @@ -37,8 +26,14 @@ make_sum, ) from .functional import ( + RealifiedFunctional, + realify, + SpectralFunctional, + eigenvalue_space, + spectralize, ComposedFunctional, Functional, + ConstantFunctional, HuberFunctional, InnerProductFunctional, KLDivergenceFunctional, @@ -49,16 +44,18 @@ MatrixFreeLinearFunctional, NegativeEntropyFunctional, NuclearNormFunctional, + ProductFunctional, QuadraticForm, ScaledFunctional, ShiftedFunctional, - SpectralLpNormFunctional, SquaredL2NormFunctional, SumFunctional, ZeroFunctional, generalized_shrinkage, least_squares, + make_constant_functional, make_functional_composed, + make_functional_product, make_functional_sum, make_scaled_functional, make_shifted_functional, @@ -66,6 +63,11 @@ prox_l1, prox_l2sq, ) +from .opfamily import ( + FunctionalScaledOperator, + OperatorFamily, + make_functional_scaled_operator, +) from .linalg import ( CGResult, ExpmMultiplyResult, @@ -116,23 +118,34 @@ from .types import DenseArray, SparseArray, ArrayLike from ._checks import checked_method -from ._contextual import ( +from .contextual import ( ContextBound, set_context, get_context, + set_check_level, + get_check_level, + use_check_level, + use_context, resolve_context_priority, register_ops, normalize_ops, normalize_context, ) +# Re-export each optional backend whose dependency is importable (JaxOps, +# CuPyOps, TorchOps). Mirrors the backend package: one loop over the single +# source of truth in ``backend._optional`` (a missing dependency is skipped, a +# broken install warns and is skipped) instead of per-backend try/except. +_optional_ops = _available_ops() +for _cls in _optional_ops: + globals().setdefault(_cls.__name__, _cls) + __all__ = [ "__version__", "Context", "CheckLevel", "CHECK_LEVELS", "BackendOps", - "jax_pytree_class", "NumpyOps", "LinOp", "ComposedLinOp", @@ -153,6 +166,7 @@ "SumToSingleLinOp", "StackedLinOp", "ComposedFunctional", + "FunctionalScaledOperator", "Functional", "LinearFunctional", "InnerProductFunctional", @@ -160,12 +174,17 @@ "QuadraticForm", "LinOpQuadraticForm", "ScaledFunctional", + "ConstantFunctional", + "ProductFunctional", "ShiftedFunctional", "SumFunctional", "ZeroFunctional", "make_functional_composed", "make_functional_sum", "make_scaled_functional", + "make_constant_functional", + "make_functional_scaled_operator", + "make_functional_product", "make_shifted_functional", "least_squares", "SquaredL2NormFunctional", @@ -174,8 +193,13 @@ "NegativeEntropyFunctional", "KLDivergenceFunctional", "HuberFunctional", - "SpectralLpNormFunctional", + "RealifiedFunctional", + "SpectralFunctional", "NuclearNormFunctional", + "OperatorFamily", + "eigenvalue_space", + "realify", + "spectralize", "generalized_shrinkage", "prox_l1", "prox_l2sq", @@ -228,15 +252,14 @@ "ContextBound", "set_context", "get_context", + "set_check_level", + "get_check_level", + "use_check_level", + "use_context", "resolve_context_priority", "register_ops", "normalize_ops", "normalize_context", ] -if "JaxOps" in globals(): - __all__.append("JaxOps") -if "TorchOps" in globals(): - __all__.append("TorchOps") -if "CuPyOps" in globals(): - __all__.append("CuPyOps") +__all__ += [_cls.__name__ for _cls in _optional_ops if _cls.__name__ != "NumpyOps"] diff --git a/spacecore/_batching.py b/spacecore/_batching.py index 2cde17f..bb7b8f8 100644 --- a/spacecore/_batching.py +++ b/spacecore/_batching.py @@ -18,10 +18,16 @@ def _check_batched(space: Any, xs: Any) -> None: def _check_scalar_shape(values: Any, shape: tuple[int, ...]) -> None: - """Raise if scalar output does not have ``shape``.""" + """Raise if scalar output does not have ``shape``. + + ``shape`` is ``()`` for a single evaluation and the leading batch shape for a + batched one; the message distinguishes the two, since this now guards every + ``Functional.value``/``vvalue`` rather than only the batched paths. + """ value_shape = tuple(getattr(values, "shape", ())) if value_shape != shape: - raise ValueError(f"Expected scalar batch output with shape {shape}, got {value_shape}.") + kind = "scalar output" if shape == () else "scalar batch output" + raise ValueError(f"Expected {kind} with shape {shape}, got {value_shape}.") def _leading_batch_size(space: Any, xs: Any) -> int: diff --git a/spacecore/_checks.py b/spacecore/_checks.py index 1701eed..963ba6f 100644 --- a/spacecore/_checks.py +++ b/spacecore/_checks.py @@ -5,6 +5,7 @@ from ._check_policy import ( CheckLevel, + check_level_at_least, enabled_to_level, normalize_check_level, require_mutually_exclusive, @@ -45,6 +46,8 @@ def checked_method( arg_positions: int | tuple[int, ...] | None = None, in_batched: bool = False, out_batched: bool = False, + out_scalar: bool = False, + out_batched_scalar: bool = False, ) -> Callable[[Callable[..., Any]], Callable[..., Any]]: """ Build a decorator that validates method inputs and outputs against spaces. @@ -68,6 +71,16 @@ def checked_method( Validate inputs as leading-axis batches instead of single elements. out_batched : bool, optional Validate outputs as leading-axis batches instead of single elements. + out_scalar : bool, optional + Validate that the output is scalar-shaped (``shape == ()``). This is the + codomain check for a :class:`~spacecore.functional.Functional`, whose + codomain is the scalar field rather than a :class:`Space` object and so + cannot be expressed through ``out_space``. + out_batched_scalar : bool, optional + Validate that the output is a vector of scalars, one per input element: + ``shape == (N,)`` for a leading batch of size ``N``. The batch size is + read from the *input* named by ``in_space`` at the first checked + position, so this flag requires ``in_space``. Returns ------- @@ -83,10 +96,23 @@ def checked_method( The validated path resolves ``in_space``/``out_space`` once per call (rather than per argument position) and uses the per-instance ``_check_member`` cache populated by :class:`spacecore.space.Space`. + + The scalar-output checks run at ``standard`` and above, matching the + hand-written ``_checks_at_least("standard")`` gating they replace. Shape is + static under ``jax.jit``, so they are safe to run while tracing. """ + require_mutually_exclusive( + "out_scalar", out_scalar or None, "out_batched_scalar", out_batched_scalar or None + ) + if out_batched_scalar and in_space is None: + # TypeError, matching ``require_mutually_exclusive``: both are misuse of + # the decorator's argument combination, not a bad runtime value. + raise TypeError("out_batched_scalar requires in_space to locate the batch size.") positions = _as_positions(arg_pos, arg_positions) # Resolve positions to a single-element fast path or a tuple iteration. single_pos = positions[0] if len(positions) == 1 else None + # The batch size for out_batched_scalar comes from the first checked input. + batch_pos = positions[0] def decorate(method: Callable[..., Any]) -> Callable[..., Any]: @wraps(method) @@ -130,6 +156,19 @@ def wrapper(self: Any, *args: Any, **kwargs: Any) -> Any: else: out_target._check_member(y) + if (out_scalar or out_batched_scalar) and check_level_at_least(level, "standard"): + from ._batching import _check_scalar_shape + + if out_scalar: + _check_scalar_shape(y, ()) + else: + from ._batching import _leading_batch_size + + batch_target = self if in_space == "self" else getattr(self, in_space) + _check_scalar_shape( + y, (_leading_batch_size(batch_target, args[batch_pos]),) + ) + return y return wrapper diff --git a/spacecore/_contextual/_policies.py b/spacecore/_contextual/_policies.py deleted file mode 100644 index 1bd0f62..0000000 --- a/spacecore/_contextual/_policies.py +++ /dev/null @@ -1,17 +0,0 @@ -from __future__ import annotations - - -class ContextError(RuntimeError): - pass - - -class ContextInferenceError(ContextError): - pass - - -class ContextConflictError(ContextError): - pass - - -class UnknownBackendError(ContextError): - pass diff --git a/spacecore/_contextual/_state.py b/spacecore/_contextual/_state.py deleted file mode 100644 index 42244c6..0000000 --- a/spacecore/_contextual/_state.py +++ /dev/null @@ -1,446 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING, Dict, Any, Iterable, Tuple -from warnings import warn - -from .._check_policy import ( - CheckLevel, - minimum_check_level, - normalize_check_level, - require_mutually_exclusive, -) -from ..types import DType -from ..backend._family import BackendFamily -from ..backend._ops import BackendOps -from ..backend.numpy import NumpyOps -from ._policies import ( - ContextConflictError, - ContextInferenceError, - UnknownBackendError, -) - -try: - from ..backend.jax import JaxOps -except ImportError: - JaxOps = None -try: - from ..backend.cupy import CuPyOps -except ImportError: - CuPyOps = None -try: - from ..backend.torch import TorchOps -except ImportError: - TorchOps = None - -if TYPE_CHECKING: - from ..backend._context import Context - - -def _context_type() -> type[Context]: - from ..backend._context import Context - - return Context - - -class Contextual: - """Resolve contexts and backend registrations.""" - - _default_ctx: Context - _available_ops: Dict[str, type[BackendOps]] - _default_dtype: DType | None = None - _default_check_level: CheckLevel = "none" - - def __init__(self) -> None: - Context = _context_type() - ops = NumpyOps() - self.default_ctx = Context( - ops=ops, - dtype=ops.sanitize_dtype(self._default_dtype), - check_level=self._default_check_level, - ) - - self._available_ops = { - self._backend_key(NumpyOps): NumpyOps, - } - if JaxOps is not None: - self._available_ops[self._backend_key(JaxOps)] = JaxOps - if CuPyOps is not None: - self._available_ops[self._backend_key(CuPyOps)] = CuPyOps - if TorchOps is not None: - self._available_ops[self._backend_key(TorchOps)] = TorchOps - - def normalize_context( - self, - ctx: Context | BackendFamily | str | None = None, - dtype: Any = None, - enable_checks: bool | None = None, - *, - check_level: CheckLevel | None = None, - ) -> Context: - Context = _context_type() - require_mutually_exclusive("check_level", check_level, "enable_checks", enable_checks) - if ctx is None: - if dtype is not None or enable_checks is not None or check_level is not None: - warn( - "Provided context is None; dtype and check policy parameters are ignored.", - UserWarning, - ) - return self.default_ctx - if isinstance(ctx, Context): - if dtype is not None or enable_checks is not None or check_level is not None: - warn( - "Provided concrete context; dtype and check policy parameters are ignored.", - UserWarning, - ) - return Context( - ops=ctx.ops, - dtype=ctx.ops.sanitize_dtype(ctx.dtype), - check_level=ctx.check_level, - ) - if isinstance(ctx, (str, BackendFamily)): - ctx = self._backend_key(ctx) - ops = self.get_ops(ctx) - return self.ctx_from_ops( - ops, - dtype=dtype, - enable_checks=enable_checks, - check_level=check_level, - ) - else: - raise TypeError(f"Expected Context, BackendFamily, str, or None, got {type(ctx)}.") - - def ctx_from_ops( - self, - ops: BackendOps, - dtype: DType | None = None, - enable_checks: bool | None = None, - *, - check_level: CheckLevel | None = None, - ) -> Context: - Context = _context_type() - dtype = ops.sanitize_dtype(dtype) - level = normalize_check_level( - check_level, - enable_checks=enable_checks, - default=self._default_check_level, - warn_legacy=enable_checks is not None, - ) - return Context(ops=ops, dtype=dtype, check_level=level) - - @property - def default_ctx(self) -> Context: - return self._default_ctx - - @default_ctx.setter - def default_ctx(self, ctx: Context | BackendFamily | str | None = None) -> None: - ctx = self.normalize_context(ctx) - self._default_ctx = ctx - - def get_ops( - self, name: str | BackendFamily | BackendOps | type[BackendOps] | Context - ) -> BackendOps: - name = self._backend_key(name) - if name not in self.available_ops: - allowed = ", ".join(k for k in self.available_ops.keys()) - raise UnknownBackendError(f"Unknown backend: {name!r}. Expected one of: {allowed}") - return self.available_ops[name]() - - @property - def available_ops(self) -> Dict[str, type[BackendOps]]: - return self._available_ops - - def register_ops(self, ops: type[BackendOps]) -> type[BackendOps]: - if not isinstance(ops, type) or not issubclass(ops, BackendOps): - raise TypeError(f"Expected type[BackendOps], got {type(ops)!r}") - else: - family = self._backend_key(ops) - if family in self.available_ops.keys(): - raise ContextConflictError(f"BackendOps {family} is already registered.") - self._available_ops[family] = ops - return ops - - def infer_context( - self, - x: Any, - enable_checks: bool | None = None, - *, - check_level: CheckLevel | None = None, - ) -> Context | None: - """Infer context from `.ctx` first, then registered backend arrays.""" - Context = _context_type() - if isinstance(x, Context): - return x - - ctx = getattr(x, "ctx", None) - if isinstance(ctx, Context): - return ctx - - matched: list[BackendOps] = [] - for name, ops in self.available_ops.items(): - try: - ops = ops() - if ops.is_array(x): - matched.append(ops) - except Exception: - # Keep inference conservative. - continue - - if not matched: - return None - if len(matched) > 1: - raise ContextInferenceError( - f"Ambiguous backend inference for object of type {type(x)!r}: {matched!r}." - ) - - ops = matched[0] - try: - dtype = ops.get_dtype(x) - except Exception: - dtype = getattr(x, "dtype", self.default_ctx.dtype) - - return self.ctx_from_ops( - ops, - dtype, - enable_checks, - check_level=check_level, - ) - - def infer_contexts(self, values: Iterable[Any]) -> Tuple[Context, ...]: - out: list[Context] = [] - for x in values: - ctx = self.infer_context(x) - if ctx is not None: - out.append(ctx) - return tuple(out) - - def are_compatible_contexts(self, *ctxs: Context) -> bool: - if len(ctxs) < 2: - return True - first = ctxs[0] - return all(ctx.ops.family == first.ops.family for ctx in ctxs[1:]) - - def are_compatible_values(self, *values: Any) -> bool: - return self.are_compatible_contexts(*self.infer_contexts(values)) - - def are_compatible_ops(self, *ops: BackendOps) -> bool: - if not ops: - return True - first = ops[0] - return all(op.family == first.family for op in ops) - - def enforce_convert_policy( - self, x: Any, to: Context | BackendFamily | str | None = None - ) -> Tuple[Any, Context]: - """Resolve the target context for ``x``.""" - self.infer_context(x) - ctx = self.normalize_context(to) - return x, ctx - - def _backend_key(self, x: str | BackendFamily | BackendOps | type[BackendOps] | Context) -> str: - Context = _context_type() - if isinstance(x, Context): - return self._backend_key(x.ops) - if isinstance(x, BackendOps): - return self._backend_key(x.family) - if isinstance(x, type) and issubclass(x, BackendOps): - return self._backend_key(x._family) - if isinstance(x, BackendFamily): - return x.value.lower() - if isinstance(x, str): - key = x.lower() - return "torch" if key == "pytorch" else key - raise TypeError(f"Unsupported backend key source: {type(x)!r}") - - def resolve_context_priority( - self, - priority_ctx: Context | BackendFamily | str | None = None, - *other_ctx: object, - ) -> Context: - """Resolve explicit context first, then compatible inferred contexts.""" - if priority_ctx is not None: - return self.normalize_context(priority_ctx) - - inferred = self.infer_contexts(other_ctx) - if not inferred: - return self.default_ctx - - if not self.are_compatible_contexts(*inferred): - fams = tuple(ctx.ops.family for ctx in inferred) - raise ValueError(f"Incompatible inferred contexts: {fams!r}") - - first = inferred[0] - ops = type(first.ops)() - dtype = self._join_dtypes(ops, *(ctx.dtype for ctx in inferred)) - - return self.ctx_from_ops( - ops=ops, - dtype=dtype, - check_level=minimum_check_level(tuple(ctx.check_level for ctx in inferred)), - ) - - def _join_dtypes(self, ops: BackendOps, *dtypes: DType | None) -> DType | None: - clean = [ops.sanitize_dtype(dt) for dt in dtypes if dt is not None] - if not clean: - return ops.sanitize_dtype(None) - - # Promote through the operands' OWN backend namespace. NumPy's - # ``result_type`` cannot interpret a torch/jax dtype, so joining the - # inferred contexts of a non-NumPy operator (for example a - # ``BlockDiagonalLinOp`` built with ``from_operators``, which infers a - # ``TreeSpace`` and joins its leaf dtypes) would otherwise raise - # ``TypeError: Cannot interpret 'torch.float64' as a data type``. - joined = ops.xp.result_type(*clean) - return ops.sanitize_dtype(joined) - - -_contextual: Contextual | None = None - - -def _state() -> Contextual: - """Return the process-wide contextual singleton.""" - global _contextual - if _contextual is None: - _contextual = Contextual() - return _contextual - - -def set_context( - ctx: Context | BackendFamily | str | None = None, - dtype: Any = None, - enable_checks: bool | None = None, - *, - check_level: CheckLevel | None = None, -) -> None: - """ - Set the process-wide default SpaceCore context. - - Parameters - ---------- - ctx : Context, BackendFamily, str, or None, optional - Context or backend specification. - dtype : Any, optional - Default dtype override. - enable_checks : bool or None, optional - Deprecated Boolean validation override. - check_level : CheckLevel or None, optional - Validation policy override for backend-name contexts. - """ - state = _state() - state.default_ctx = state.normalize_context( - ctx, - dtype=dtype, - enable_checks=enable_checks, - check_level=check_level, - ) - - -def get_context() -> Context: - """ - Return the current process-wide default SpaceCore context. - - Returns - ------- - Context - Active process-wide default context. - """ - return _state().default_ctx - - -def resolve_context_priority( - priority_ctx: Context | BackendFamily | str | None = None, - *other_ctx: object, -) -> Context: - """ - Resolve the context assigned to a newly created object. - - Parameters - ---------- - priority_ctx : Context, BackendFamily, str, or None, optional - Explicit context that takes precedence when provided. - *other_ctx : object - Objects or contexts used as fallback context sources. - - Returns - ------- - Context - Resolved context. - """ - return _state().resolve_context_priority(priority_ctx, *other_ctx) - - -def register_ops(ops: type[BackendOps]) -> type[BackendOps]: - """ - Register a backend operations implementation. - - Parameters - ---------- - ops : type of BackendOps - Backend operations class to register. - - Returns - ------- - type of BackendOps - Registered backend operations class. - """ - return _state().register_ops(ops) - - -def normalize_context( - ctx: Context | BackendFamily | str | None = None, - dtype: Any = None, - enable_checks: bool | None = None, - *, - check_level: CheckLevel | None = None, -) -> Context: - """ - Normalize a context specification through the process-wide state. - - Parameters - ---------- - ctx : Context, BackendFamily, str, or None, optional - Context or backend specification. - dtype : Any, optional - Default dtype override. - enable_checks : bool or None, optional - Deprecated Boolean validation override. - check_level : CheckLevel or None, optional - Validation policy override for backend-name contexts. - - Returns - ------- - Context - Normalized context. - """ - return _state().normalize_context( - ctx, - dtype=dtype, - enable_checks=enable_checks, - check_level=check_level, - ) - - -def normalize_ops(ops: str | BackendFamily | BackendOps | type[BackendOps] | Context) -> BackendOps: - """ - Normalize backend operations through the process-wide state. - - Parameters - ---------- - ops : str, BackendFamily, BackendOps, type of BackendOps, or Context - Backend operations specification. - - Returns - ------- - BackendOps - Normalized backend operations singleton. - """ - if isinstance(ops, BackendOps): - return ops - return _state().get_ops(ops) - - -def enforce_convert_policy( - x: Any, - to: Context | BackendFamily | str | None = None, -) -> tuple[Any, Context]: - """Resolve a conversion target context.""" - return _state().enforce_convert_policy(x, to) diff --git a/spacecore/_errors.py b/spacecore/_errors.py new file mode 100644 index 0000000..80ea110 --- /dev/null +++ b/spacecore/_errors.py @@ -0,0 +1,29 @@ +"""Errors shared by the backend registry and the context layer. + +These live at the top level, imported by nothing, because both +:mod:`spacecore.backend` (which owns the backend-ops registry) and +:mod:`spacecore.contextual` (which owns ambient policy and re-exports them as +public API) need to raise them. Defining them in either package would force the +other to import it, and ``contextual`` already depends on ``backend``. + +The ``Context*`` names are kept for backward compatibility: they were public from +``spacecore.contextual`` before the registry was extracted, and renaming them +would break callers catching them. +""" +from __future__ import annotations + + +class ContextError(RuntimeError): + """Base class for context resolution and backend registration failures.""" + + +class ContextInferenceError(ContextError): + """A value's backend could not be inferred unambiguously.""" + + +class ContextConflictError(ContextError): + """A backend family is already registered under a different implementation.""" + + +class UnknownBackendError(ContextError): + """No backend is registered under the requested family name.""" diff --git a/spacecore/_lazy_algebra.py b/spacecore/_lazy_algebra.py new file mode 100644 index 0000000..01b03be --- /dev/null +++ b/spacecore/_lazy_algebra.py @@ -0,0 +1,181 @@ +"""Canonicalization primitives shared by the lazy-algebra factories. + +The LinOp and Functional algebras build the same kind of expression tree — lazy +*sums*, *scalar multiples*, and *zero* elements with only local canonicalization +(flatten nested sums, drop zeros, unwrap singletons, fold nested scalings). The +node *semantics* differ (``apply``/``rapply`` for operators, ``value``/``grad`` +for functionals) and stay in each domain; only this canonicalization skeleton is +shared here, so ``make_sum`` / ``make_scaled`` are not duplicated between +:mod:`spacecore.linop` and :mod:`spacecore.functional`. + +Each generic is parameterized by a tiny "node vocabulary" (predicates and node +factories) so a caller stays fully in control of the concrete node types. +""" +from __future__ import annotations + +from numbers import Number +from typing import Any, Callable, Sequence, TypeVar + +T = TypeVar("T") + + +def is_scalar_like(value: Any) -> bool: + """Return whether ``value`` can be used as a scalar multiplier.""" + if isinstance(value, Number): + return True + shape = getattr(value, "shape", None) + if shape is not None: + return tuple(shape) == () + return getattr(value, "ndim", None) == 0 + + +def conjugate_scalar(value: Any) -> Any: + """Return the scalar conjugate when the value supports conjugation. + + Duck-typed rather than backend-dispatched: Python numbers, NumPy/CuPy + scalars, and 0-d Torch/JAX arrays all expose ``conjugate`` or ``conj``, and + a value that exposes neither is its own conjugate as far as this module is + concerned. + """ + if hasattr(value, "conjugate"): + return value.conjugate() + if hasattr(value, "conj"): + return value.conj() + return value + + +def scalar_eq(a: Any, b: Any) -> bool: + """Return whether two scalar-likes are *recognizably* equal, as a real ``bool``. + + Recognizably: the comparison must be decidable from the values themselves. + A concrete pair answers truthfully; an **abstract** (traced) scalar has no + value at trace time, so no comparison about it can be decided and this + returns ``False`` — "not recognizably equal", the conservative answer. The + callers are canonicalizers (:func:`fold_scaled`, the ``__eq__`` of scaled + nodes), and for them ``False`` means "skip the simplification", which is + always sound. Under ``jax.jit`` a traced coefficient therefore builds a + larger-but-correct expression tree rather than a folded one. + + That verdict is deliberate, so the catch is narrowed to :class:`TypeError` — + the base class of JAX's ``TracerBoolConversionError`` and of any other + "cannot reduce to a concrete bool" failure. Every other exception means the + operand's ``__eq__`` is itself broken and propagates instead of being + silently reported as inequality. + + Two matching NaN scalars compare equal (mirrors ``equal_nan=True``), so a + NaN-scaled node equals itself. Always returns a genuine Python ``bool`` — a + 0-d backend-array ``==`` would yield ``np.bool_``, which leaks through the + ``and`` combinator of any container's ``__eq__``. + """ + try: + if bool(a == b): + return True + # ``x != x`` is True only for NaN (including a complex value with a NaN + # component), so this matches NaN against NaN. + return bool(a != a) and bool(b != b) + except TypeError: + return False + + +def is_recognizably_nonreal(value: Any) -> bool: + """Return whether ``value`` is *provably* a non-real scalar. + + The polarity matters: this answers "can we see an imaginary part?", not "is + this real?". A traced scalar has no value to inspect, so it is reported as + ``False`` (not provably non-real) and callers let it through. Rejecting on + ``True`` therefore never fires spuriously under ``jax.jit``; it only rejects + a concrete value that genuinely carries an imaginary part. + + Used by spaces whose scalar field is narrower than their entry dtype (see + :attr:`spacecore.space.Space.scalar_field`), where an imaginary multiplier + would silently produce an element outside the space. + """ + try: + # A NaN in any component makes ``!=`` uninformative (``nan != nan`` is + # True for a *real* NaN), so nothing is provable and we let it through. + if bool(value != value): + return False + return bool(value != conjugate_scalar(value)) + except TypeError: + return False + + +def flatten_sum( + terms: Sequence[T], + *, + is_sum: Callable[[T], bool], + parts: Callable[[T], Sequence[T]], +) -> tuple[T, ...]: + """Flatten nested sum nodes into a flat tuple of terms. + + Parameters + ---------- + terms: + Operands to flatten (already validated by the caller). + is_sum: + Predicate that is ``True`` for a sum node. + parts: + Accessor returning the child terms of a sum node. + """ + flat: list[T] = [] + for term in terms: + if is_sum(term): + flat.extend(flatten_sum(parts(term), is_sum=is_sum, parts=parts)) + else: + flat.append(term) + return tuple(flat) + + +def finalize_sum( + terms: Sequence[T], + *, + is_zero: Callable[[T], bool], + make_zero: Callable[[], T], + make_sum_node: Callable[[tuple[T, ...]], T], +) -> T: + """Drop zero terms, then collapse to zero, unwrap a singleton, or build a node. + + The caller is expected to have already flattened and validated ``terms``. + """ + nonzero = tuple(term for term in terms if not is_zero(term)) + if not nonzero: + return make_zero() + if len(nonzero) == 1: + return nonzero[0] + return make_sum_node(nonzero) + + +def fold_scaled( + scalar: Any, + operand: T, + *, + is_zero: Callable[[T], bool], + unwrap_scaled: Callable[[T], "tuple[Any, T] | None"], + make_zero: Callable[[], T], + make_scaled_node: Callable[[Any, T], T], +) -> T: + """Canonicalize ``scalar * operand``. + + ``0 -> zero``, ``1 -> operand``, a zero operand passes through, nested + scalings fold into one coefficient, otherwise a scaled node is built. + ``unwrap_scaled(operand)`` returns ``(inner_scalar, inner_operand)`` for a + scaled node, or ``None`` otherwise. + """ + if scalar_eq(scalar, 0): + return make_zero() + if scalar_eq(scalar, 1): + return operand + if is_zero(operand): + return operand + nested = unwrap_scaled(operand) + if nested is not None: + inner_scalar, inner = nested + return fold_scaled( + scalar * inner_scalar, + inner, + is_zero=is_zero, + unwrap_scaled=unwrap_scaled, + make_zero=make_zero, + make_scaled_node=make_scaled_node, + ) + return make_scaled_node(scalar, operand) diff --git a/spacecore/_version.py b/spacecore/_version.py index 00b4de4..18493bb 100644 --- a/spacecore/_version.py +++ b/spacecore/_version.py @@ -1,3 +1,3 @@ """Single source of truth for the SpaceCore package version.""" -__version__ = "0.4.2" +__version__ = "0.4.3" diff --git a/spacecore/backend/__init__.py b/spacecore/backend/__init__.py index 0b9def0..6da8e01 100644 --- a/spacecore/backend/__init__.py +++ b/spacecore/backend/__init__.py @@ -1,41 +1,42 @@ """Backend contexts and operation implementations.""" from .._check_policy import CHECK_LEVELS, CheckLevel -from ._context import Context +from ._container import PyTreeNode, PyTreeRegistry, registry as pytree_registry from ._ops import BackendOps from ._family import BackendFamily -from .jax._pytree import jax_pytree_class from .numpy import NumpyOps +from ._optional import available_ops +from ._registry import OpsRegistry, backend_key, registry as ops_registry -try: - from .jax import JaxOps as JaxOps -except ImportError: - pass -try: - from .cupy import CuPyOps as CuPyOps -except ModuleNotFoundError as exc: - if exc.name != "cupy": - raise - -try: - from .torch import TorchOps as TorchOps -except ModuleNotFoundError as exc: - if exc.name != "torch": - raise +# Bind each optional backend whose dependency is importable (JaxOps, CuPyOps, +# TorchOps). ``_optional.available_ops`` centralizes the guarded import — a +# missing dependency is skipped, a broken install warns and is skipped — so the +# per-backend try/except blocks collapse to this loop over its single source of +# truth. ``NumpyOps`` is always first and already imported above. +# +# Each bound backend also installs its tree protocol with the container registry. +# This runs while ``spacecore.backend`` is importing, i.e. before ``linop``, +# ``functional`` and ``space`` define their container classes, so those classes +# register against an already-present protocol. A backend that arrives later +# still back-fills every class defined so far — the registry is order-free by +# construction — and a backend with no transform machinery inherits a no-op. +_ops = available_ops() +for _cls in _ops: + globals().setdefault(_cls.__name__, _cls) + _cls.install_pytree_protocol() __all__ = [ - "Context", "CheckLevel", "CHECK_LEVELS", "BackendFamily", "BackendOps", - "jax_pytree_class", "NumpyOps", + "OpsRegistry", + "PyTreeNode", + "PyTreeRegistry", + "backend_key", + "ops_registry", + "pytree_registry", + "available_ops", + *(_cls.__name__ for _cls in _ops if _cls.__name__ != "NumpyOps"), ] - -if "JaxOps" in globals(): - __all__.append("JaxOps") -if "CuPyOps" in globals(): - __all__.append("CuPyOps") -if "TorchOps" in globals(): - __all__.append("TorchOps") diff --git a/spacecore/backend/_container.py b/spacecore/backend/_container.py new file mode 100644 index 0000000..e46188f --- /dev/null +++ b/spacecore/backend/_container.py @@ -0,0 +1,215 @@ +"""Backend-neutral container protocol and the registry that wires it to backends. + +A SpaceCore ``LinOp``/``Functional``/``Space`` is a *structured container*: it +decomposes into ``(children, aux)`` — dynamic array-bearing children plus static +metadata — and rebuilds from that pair. Every backend with function-transform +machinery needs to know that structure in order to see inside our objects: +``jax.jit``/``grad``/``vmap`` and ``lax.while_loop`` consult ``jax.tree_util``'s +pytree registry, ``torch.compile``/``torch.func`` consult Torch's. There is **no +shared cross-framework registry** — registering with one does nothing for the +other — so supporting N backends means N registration calls per class. + +That is the bookkeeping no pytree library hands you, and it is all this module +does. Two sets grow independently and in an order nobody controls: + +* **container classes**, as ``linop``/``functional``/``space`` modules import; +* **backends with a tree protocol**, as each ``BackendOps`` is loaded. + +:class:`PyTreeRegistry` keeps both and, whenever either side gains a member, +wires it against everything currently on the other side (*two-axis back-fill*). +Registration therefore does not depend on import order, which is what lets a +backend loaded late still pick up classes defined early, and vice versa. + +The flatten protocol itself lives on :class:`PyTreeNode`, the capability mixin +concrete containers inherit; the per-backend translation lives in each backend's +adapter. + +This module sits in :mod:`spacecore.backend` — the backend-*neutral* abstraction +layer, beside the :class:`~spacecore.backend.BackendOps` contract whose +``install_pytree_protocol`` feeds the registry — and imports **nothing**, not even +from its own package. No individual backend is named here; each one reaches in +from its own adapter. So the concept stays backend-agnostic while living next to +the layer it serves, and the module is safe to import from anywhere without +risking a cycle. + +Note on global state: ``registry`` is a process-wide singleton, but a *monotonic* +one — classes and backends are only ever added, wiring is idempotent, and the +final state is a function of what got imported rather than of the order it got +imported in. That grow-only property is what makes a shared instance safe here; +do not add removal or reset to it (tests construct their own instance instead). +""" +from __future__ import annotations + +import threading +import warnings +from abc import ABC, abstractmethod +from typing import Any, Callable, Self + +#: Registers one class with one backend's pytree system. Supplied by that +#: backend's adapter, which owns the translation from our ``tree_flatten`` / +#: ``tree_unflatten`` pair into whatever shape the backend expects. +TreeRegistrar = Callable[[type], None] + + +class PyTreeRegistry: + """Backend-keyed registry wiring container classes into backend tree protocols. + + Maintains the cross-product of registered classes and registered backends, + guaranteeing each ``(backend, class)`` pair is wired exactly once regardless + of the order the two arrive in. + + Instantiable so tests can exercise the wiring with a fake registrar and no + backend installed; library code uses the module-level :data:`registry`. + """ + + def __init__(self) -> None: + self._classes: list[type] = [] + self._seen: set[type] = set() + self._backends: dict[str, TreeRegistrar] = {} + self._done: set[tuple[str, type]] = set() + # Registration normally happens during imports (already serialized), but + # a backend may install its protocol lazily at runtime; an RLock keeps + # the check-then-wire in _wire atomic without costing anything at import. + self._lock = threading.RLock() + + def register_class(self, cls: type) -> None: + """Record a container class and wire it into every backend known so far. + + Parameters + ---------- + cls : type + Concrete container class providing ``tree_flatten``/``tree_unflatten``. + """ + with self._lock: + if cls not in self._seen: + self._seen.add(cls) + self._classes.append(cls) + for name, registrar in self._backends.items(): + self._wire(name, registrar, cls) + + def register_backend(self, name: str, registrar: TreeRegistrar) -> None: + """Record a backend's tree protocol and back-fill every class known so far. + + Parameters + ---------- + name : str + Backend family name, e.g. ``"jax"`` or ``"torch"``. + registrar : callable + Registers a single class with that backend's pytree system. + """ + with self._lock: + self._backends[name] = registrar + for cls in self._classes: + self._wire(name, registrar, cls) + + def registered_classes(self) -> tuple[type, ...]: + """Return the container classes recorded so far, in registration order.""" + with self._lock: + return tuple(self._classes) + + def backends(self) -> tuple[str, ...]: + """Return the names of the backends whose tree protocol is installed.""" + with self._lock: + return tuple(self._backends) + + def is_wired(self, name: str, cls: type) -> bool: + """Return whether ``cls`` has been wired into backend ``name``.""" + with self._lock: + return (name, cls) in self._done + + def _wire(self, name: str, registrar: TreeRegistrar, cls: type) -> None: + """Register ``cls`` with backend ``name`` at most once. + + A failing registrar must never abort ``import spacecore`` — an unusable + transform integration is a far smaller problem than an unimportable + library, and hard-failing here would reintroduce exactly the eager-import + fragility this design removes. The pair is marked done even on failure so + a benign cause (the class already registered with that backend by other + means, which several backends report by raising) cannot produce repeated + attempts or warning spam. + """ + key = (name, cls) + if key in self._done: + return + self._done.add(key) + try: + registrar(cls) + except Exception as exc: # noqa: BLE001 - see docstring: import must survive + warnings.warn( + f"Could not register {cls.__module__}.{cls.__qualname__} with the " + f"{name!r} pytree protocol: {type(exc).__name__}: {exc}. " + f"{name} transforms will treat instances as opaque leaves.", + RuntimeWarning, + stacklevel=3, + ) + + +#: Process-wide registry used by :class:`PyTreeNode` and the backend adapters. +registry = PyTreeRegistry() + + +class PyTreeNode(ABC): + """Capability: a structured container exposed to backends' transform systems. + + Owns the flatten protocol the registry depends on, and auto-registers every + concrete subclass with :data:`registry`. Compose it alongside the other + capability mixins; it is deliberately independent of ``ContextBound`` — being + bound to a backend context and being a flattenable container are orthogonal + capabilities, and keeping them apart means adding a backend's tree protocol + never reaches into the context/check-policy layer. + + Implementations must round-trip: ``cls.tree_unflatten(*reversed(obj.tree_flatten()))`` + reconstructs an equal object. Children are the dynamic, array-bearing parts a + transform may trace or map over; ``aux`` is static metadata compared by + equality when a backend reassembles the structure. + """ + + __slots__ = () + + @abstractmethod + def tree_flatten(self) -> tuple[tuple[Any, ...], Any]: + """Return ``(children, aux)`` for this container. + + Returns + ------- + tuple + ``children`` — dynamic, array-bearing members, in a stable order; + ``aux`` — static metadata sufficient, with ``children``, to rebuild. + """ + + @classmethod + @abstractmethod + def tree_unflatten(cls, aux: Any, children: tuple[Any, ...]) -> Self: + """Rebuild an instance from ``aux`` and ``children``. + + Parameters + ---------- + aux : object + The static metadata returned by :meth:`tree_flatten`. + children : tuple + The dynamic members, possibly transformed or traced. + """ + + def __init_subclass__(cls, **kwargs: Any) -> None: + """Register subclasses that provide the flatten protocol. + + Gates on the *protocol method* rather than on :func:`inspect.isabstract` + because the two ask different questions. ``isabstract`` asks "is anything + still unimplemented?"; the registry only needs "is this class flattenable + yet?". They disagree for a class that supplies ``tree_flatten`` but stays + abstract for an unrelated reason — say a ``LinOp`` subclass that factors + out flattening while leaving ``apply`` to its own subclasses. Such a class + is perfectly registrable: a registrar records the type and consults the + two protocol methods, nothing else. Registering it is also harmless, since + it is never instantiated. + + (``inspect.isabstract`` *does* report correctly from inside + ``__init_subclass__``. ``__abstractmethods__`` is not populated until + ``ABCMeta.__new__`` returns, but CPython detects that and falls back to + scanning for abstract methods manually. The choice here is about which + question to ask, not a workaround for a broken one.) + """ + super().__init_subclass__(**kwargs) + if getattr(cls.tree_flatten, "__isabstractmethod__", False): + return + registry.register_class(cls) diff --git a/spacecore/backend/_ops.py b/spacecore/backend/_ops.py index 82af281..7143be8 100644 --- a/spacecore/backend/_ops.py +++ b/spacecore/backend/_ops.py @@ -59,6 +59,29 @@ def has_native_vmap(self) -> bool: """Whether ``vmap`` is implemented by the backend rather than a Python loop.""" return False + @classmethod + def install_pytree_protocol(cls) -> None: + """Register this backend's tree protocol with the container registry. + + Called once per available backend as it is bound (see + :mod:`spacecore.backend`). A backend that implements this teaches its + function-transform machinery to see inside SpaceCore containers, so + instances can flow through that backend's tracing, batching and + differentiation rather than being treated as opaque leaves. + + The default is a no-op, and that is the correct answer for a backend with + no transform machinery — NumPy and CuPy have no pytree registry to + register with. + + Implementations translate + :meth:`~spacecore.backend._container.PyTreeNode.tree_flatten` / + ``tree_unflatten`` into whatever shape the backend expects and hand the + result to ``spacecore.backend._container.registry.register_backend``. Import the + backend's own modules *inside* this method, never at module scope, so an + absent or broken dependency cannot break ``import spacecore``. + """ + return + def free_memory_bytes(self) -> int | None: """Return currently free device memory in bytes, or ``None`` if unknown. @@ -297,6 +320,29 @@ def _to_axis_tuple(self, axis: int | Sequence[int] | None) -> int | tuple[int, . return axis return tuple(axis) + def _as_array(self, x: Any) -> DenseArray: + """Return ``x`` as a backend array, promoting a host scalar to 0-d. + + SpaceCore's convention is that every array-valued method returns a + backend *array*, never a host scalar. NumPy (and CuPy) reductions + return ``np.float64``-style scalars for which :meth:`is_array` is + False and :meth:`get_dtype` raises; JAX and Torch already return 0-d + arrays. Wrapping here makes the four backends agree. + + The ``is_array`` guard keeps this free on backends that never produce + host scalars, and keeps JAX tracers (which are ``jax.Array`` + instances) untouched, so the call is safe inside ``jit``. + """ + return x if self.is_array(x) else self.xp.asarray(x) + + def _take_along_axis(self, x: DenseArray, indices: DenseArray, axis: int) -> DenseArray: + """Gather along ``axis`` using per-position ``indices``. + + Used by the spectral-ordering normalization. Torch spells this + ``take_along_dim`` and overrides this method. + """ + return self.xp.take_along_axis(x, indices, axis=axis) + def _permute_dims(self, x: DenseArray, axes: Sequence[int]) -> DenseArray: axes = tuple(axes) if hasattr(self.xp, "permute_dims"): @@ -392,6 +438,37 @@ def real_dtype(self, dtype: DType) -> DType: return self.sanitize_dtype("float64") return self.sanitize_dtype("float32" if itemsize <= 8 else "float64") + def complex_dtype(self, dtype: DType) -> DType: + """ + Return the complex dtype with the same precision as ``dtype``. + + The inverse of :meth:`real_dtype`, and the pairing a caller needs when + moving between a complex space and a real view of it. Backend dtype + policy belongs here rather than at the call site: promoting through + ``result_type(dtype, 1j)`` is *not* portable — NumPy widens ``float32`` + to ``complex128`` because a Python ``complex`` is double precision, while + JAX and Torch give ``complex64``. + + Parameters + ---------- + dtype: + Backend or portable dtype specifier. + + Returns + ------- + DType + ``dtype`` itself when it is already complex; otherwise the complex + dtype whose real component is ``dtype`` (``complex64`` for float32, + ``complex128`` for float64). + """ + dtype = self.sanitize_dtype(dtype) + if self.is_complex_dtype(dtype): + return dtype + for candidate in (self.xp.complex64, self.xp.complex128): + if self.real_dtype(candidate) == dtype: + return candidate + return self.xp.result_type(dtype, self.xp.asarray(1j)) + def get_dtype(self, x: Any) -> DType: """Return x.dtype after verifying x is a backend array.""" if self.is_array(x): @@ -482,14 +559,30 @@ def eye(self, n: int, m: int | None = None, dtype: DType | None = None) -> Dense """Create an identity-like matrix (delegates to xp.eye).""" return self.xp.eye(n, m, dtype=self._dtype_arg(dtype)) + # -- Flattening order -------------------------------------------------- + # + # SpaceCore flattens in **C order** (row-major: the last axis varies + # fastest), on every backend, always. This is the vec convention that the + # coordinate layer is built on -- ``CoordinateSpace.flatten``, + # ``LinOp.to_matrix``, and the Kronecker identities in ``kernels/`` all + # assume it, and vec(A x) = (I ⊗ A) vec(x) rather than (Aᵀ ⊗ I) is exactly + # this choice. It is a *declared* convention rather than a forwarded + # ``order=`` argument because Torch has no Fortran-order reshape at all: + # offering the parameter would mean offering it on three backends of four. + # + # The order is over *index* positions, not memory layout, so it is + # unaffected by whether the input is C-contiguous, Fortran-contiguous, or + # a non-contiguous view -- a transposed array flattens by its logical + # index order, not its buffer. + def ravel(self, x: DenseArray) -> DenseArray: - """Flatten x to one dimension.""" + """Flatten x to one dimension in C order (last axis varies fastest).""" if hasattr(self.xp, "ravel"): return self.xp.ravel(x) return self.reshape(x, (-1,)) def reshape(self, x: DenseArray, shape: Tuple[int, ...] | int) -> DenseArray: - """Reshape x (delegates to xp.reshape).""" + """Reshape x in C order (delegates to xp.reshape; see the note above).""" shape_arg = (shape,) if isinstance(shape, int) else shape return self.xp.reshape(x, shape_arg) @@ -769,12 +862,14 @@ def sum( keepdims: bool = False, dtype: DType | None = None, ) -> DenseArray: - """Sum over given axes (delegates to xp.sum).""" - return self.xp.sum( - x, - axis=self._to_axis_tuple(axis), - dtype=self._dtype_arg(dtype), - keepdims=keepdims, + """Sum over given axes (delegates to xp.sum). Returns a 0-d array when total.""" + return self._as_array( + self.xp.sum( + x, + axis=self._to_axis_tuple(axis), + dtype=self._dtype_arg(dtype), + keepdims=keepdims, + ) ) def mean( @@ -783,8 +878,8 @@ def mean( axis: int | Sequence[int] | None = None, keepdims: bool = False, ) -> DenseArray: - """Mean over given axes (delegates to xp.mean).""" - return self.xp.mean(x, axis=self._to_axis_tuple(axis), keepdims=keepdims) + """Mean over given axes (delegates to xp.mean). Returns a 0-d array when total.""" + return self._as_array(self.xp.mean(x, axis=self._to_axis_tuple(axis), keepdims=keepdims)) def min( self, @@ -792,8 +887,8 @@ def min( axis: int | Sequence[int] | None = None, keepdims: bool = False, ) -> DenseArray: - """Minimum over given axes (delegates to xp.min).""" - return self.xp.min(x, axis=self._to_axis_tuple(axis), keepdims=keepdims) + """Minimum over given axes (delegates to xp.min). Returns a 0-d array when total.""" + return self._as_array(self.xp.min(x, axis=self._to_axis_tuple(axis), keepdims=keepdims)) def max( self, @@ -801,8 +896,8 @@ def max( axis: int | Sequence[int] | None = None, keepdims: bool = False, ) -> DenseArray: - """Maximum over given axes (delegates to xp.max).""" - return self.xp.max(x, axis=self._to_axis_tuple(axis), keepdims=keepdims) + """Maximum over given axes (delegates to xp.max). Returns a 0-d array when total.""" + return self._as_array(self.xp.max(x, axis=self._to_axis_tuple(axis), keepdims=keepdims)) def prod( self, @@ -811,47 +906,54 @@ def prod( keepdims: bool = False, dtype: DType | None = None, ) -> DenseArray: - """Product over given axes (delegates to xp.prod).""" - return self.xp.prod( - x, - axis=self._to_axis_tuple(axis), - dtype=self._dtype_arg(dtype), - keepdims=keepdims, + """Product over given axes (delegates to xp.prod). Returns a 0-d array when total.""" + return self._as_array( + self.xp.prod( + x, + axis=self._to_axis_tuple(axis), + dtype=self._dtype_arg(dtype), + keepdims=keepdims, + ) ) def trace(self, x: DenseArray) -> DenseArray: - """Trace of a matrix (delegates to xp.trace when available).""" - if hasattr(self.xp, "trace"): - return self.xp.trace(x) - return self.sum(self.diagonal(x)) + """Trace over the **trailing two axes**; batched input traces each matrix. + + Deliberately *not* ``xp.trace(x)``: NumPy and JAX default to + ``axis1=0, axis2=1``, which on a ``(batch, n, n)`` stack traces across + the batch axis and silently returns a length-``n`` vector, while + ``torch.trace`` rejects anything but a 2-D input. SpaceCore's batched + convention is leading batch axes, so the trailing two are the matrix. + """ + return self.sum(self.diagonal(x), axis=-1) def argsort(self, x: DenseArray, axis: int = -1) -> DenseArray: """Return indices that sort ``x`` along an axis.""" - return self.xp.argsort(x, axis=axis) + return self._as_array(self.xp.argsort(x, axis=axis)) def sort(self, x: DenseArray, axis: int = -1) -> DenseArray: """Sort x along an axis (delegates to xp.sort).""" - return self.xp.sort(x, axis=axis) + return self._as_array(self.xp.sort(x, axis=axis)) def argmin(self, x: DenseArray, axis: int | None = None, keepdims: bool = False) -> DenseArray: - """Return indices of minima along an axis.""" - return self.xp.argmin(x, axis=axis, keepdims=keepdims) + """Return indices of minima along an axis. Returns a 0-d array when total.""" + return self._as_array(self.xp.argmin(x, axis=axis, keepdims=keepdims)) def argmax(self, x: DenseArray, axis: int | None = None, keepdims: bool = False) -> DenseArray: - """Return indices of maxima along an axis.""" - return self.xp.argmax(x, axis=axis, keepdims=keepdims) + """Return indices of maxima along an axis. Returns a 0-d array when total.""" + return self._as_array(self.xp.argmax(x, axis=axis, keepdims=keepdims)) def vdot(self, x: DenseArray, y: DenseArray) -> DenseArray: """Return ``sum(conj(x) * y)`` over flattened inputs. Matches NumPy, JAX, and Torch ``vdot`` semantics. ``DenseLinOp.rapply`` - relies on this convention for complex inputs. + relies on this convention for complex inputs. Returns a 0-d array. """ x_flat = self.ravel(x) y_flat = self.ravel(y) if hasattr(self.xp, "vdot"): - return self.xp.vdot(x_flat, y_flat) - return self.xp.vecdot(x_flat, y_flat) + return self._as_array(self.xp.vdot(x_flat, y_flat)) + return self._as_array(self.xp.vecdot(x_flat, y_flat)) def matmul( self, @@ -870,15 +972,74 @@ def einsum(self, subscripts: str, *operands: DenseArray) -> DenseArray: """Einstein summation (delegates to xp.einsum).""" return self.xp.einsum(subscripts, *operands) + # -- Spectral presentation conventions --------------------------------- + # + # A decomposition is determined only up to a choice of ordering and, for + # each vector, a sign (real) or unit phase (complex). Four backends are + # free to make those choices differently -- and do: Torch returns complex + # Hermitian eigenvectors that are -1 times NumPy's and JAX's. SpaceCore + # therefore *declares* the presentation and normalizes to it here rather + # than inheriting whatever the underlying library returned: + # + # * eigenvalues ascending, singular values descending; + # * each eigenvector / singular vector scaled so that its entry of + # largest magnitude is real and positive. + # + # NOT normalized, and not normalizable: the basis of a degenerate + # eigenspace (multiplicity > 1) is arbitrary, so callers must not depend + # on it. Ties in "entry of largest magnitude" are broken by first + # occurrence, which all backends agree on. + # + # Everything below is pure array arithmetic -- no data-dependent Python + # branch, no host sync -- so it stays safe under jit, vmap and compile. + # Cost is O(n^2) on top of an O(n^3) decomposition. + + def _column_phase(self, vectors: DenseArray) -> DenseArray: + """Return the unit phase of each column's largest entry, shaped ``(..., 1, k)``.""" + magnitude = self.abs(vectors) + pivot_row = self.xp.argmax(magnitude, axis=-2) + pivot_row = self.expand_dims(pivot_row, -2) + rows = self.reshape(self.arange(self.shape(vectors)[-2]), (-1, 1)) + selector = self.astype(rows == pivot_row, self.get_dtype(vectors)) + pivot = self.sum(vectors * selector, axis=-2, keepdims=True) + scale = self.abs(pivot) + # A unit-norm vector cannot be all zeros; the guard only keeps the + # division defined for degenerate inputs rather than changing a result. + safe_scale = self.where(scale > 0, scale, self.ones_like(scale)) + return self.where(scale > 0, pivot / safe_scale, self.ones_like(pivot)) + + def _gauge_columns(self, vectors: DenseArray) -> DenseArray: + """Scale each column so its largest-magnitude entry is real positive.""" + return vectors / self._column_phase(vectors) + + def _order_eigenpairs( + self, + eigenvalues: DenseArray, + eigenvectors: DenseArray, + ) -> tuple[DenseArray, DenseArray]: + """Sort eigenpairs ascending by eigenvalue and gauge the eigenvectors.""" + order = self.xp.argsort(self.real(eigenvalues), axis=-1) + eigenvalues = self._take_along_axis(eigenvalues, order, -1) + eigenvectors = self._take_along_axis(eigenvectors, self.expand_dims(order, -2), -1) + return eigenvalues, self._gauge_columns(eigenvectors) + def eigh( self, x: DenseArray, backend_kwargs: dict[str, Any] | None = None, ) -> tuple[DenseArray, DenseArray]: - """Eigenpairs of a Hermitian dense matrix (delegates to xp.linalg.eigh).""" + """Eigenpairs of a Hermitian dense matrix, in SpaceCore's presentation. + + Eigenvalues ascending; each eigenvector scaled so its largest-magnitude + entry is real and positive. See the conventions note above -- this is + deliberately not raw ``xp.linalg.eigh`` output. + """ if self.is_sparse(x): raise TypeError("eigh requires a dense array; sparse input is not supported.") - return self.xp.linalg.eigh(x, **({} if backend_kwargs is None else backend_kwargs)) + eigenvalues, eigenvectors = self.xp.linalg.eigh( + x, **({} if backend_kwargs is None else backend_kwargs) + ) + return self._order_eigenpairs(eigenvalues, eigenvectors) def norm( self, @@ -887,8 +1048,8 @@ def norm( axis: int | Sequence[int] | None = None, keepdims: bool = False, ) -> DenseArray: - """Vector or matrix norm (delegates to xp.linalg.norm).""" - return self.xp.linalg.norm(x, ord=ord, axis=axis, keepdims=keepdims) + """Vector or matrix norm (delegates to xp.linalg.norm). Returns a 0-d array when total.""" + return self._as_array(self.xp.linalg.norm(x, ord=ord, axis=axis, keepdims=keepdims)) def solve( self, @@ -904,8 +1065,11 @@ def eigvalsh( A: DenseArray, backend_kwargs: dict[str, Any] | None = None, ) -> DenseArray: - """Eigenvalues of a Hermitian dense matrix (delegates to xp.linalg.eigvalsh).""" - return self.xp.linalg.eigvalsh(A, **({} if backend_kwargs is None else backend_kwargs)) + """Eigenvalues of a Hermitian dense matrix, ascending (SpaceCore convention).""" + eigenvalues = self.xp.linalg.eigvalsh( + A, **({} if backend_kwargs is None else backend_kwargs) + ) + return self.sort(self.real(eigenvalues), axis=-1) def svd( self, @@ -913,12 +1077,47 @@ def svd( full_matrices: bool = True, backend_kwargs: dict[str, Any] | None = None, ) -> tuple[DenseArray, DenseArray, DenseArray]: - """Singular value decomposition (delegates to xp.linalg.svd).""" - return self.xp.linalg.svd( + """Singular value decomposition, in SpaceCore's presentation. + + Singular values descending; each of the leading ``k = min(m, n)`` + columns of ``U`` scaled so its largest-magnitude entry is real and + positive, with the compensating phase applied to the matching row of + ``Vh`` so that ``U @ diag(s) @ Vh`` still reconstructs ``A`` exactly. + + With ``full_matrices=True`` the trailing columns of ``U`` (and rows of + ``Vh``) span a null space whose basis is arbitrary and backend-specific + -- they are left untouched, and are not comparable across backends. + Pass ``full_matrices=False`` when you need agreement. + """ + U, s, Vh = self.xp.linalg.svd( A, full_matrices=full_matrices, **({} if backend_kwargs is None else backend_kwargs), ) + return self._order_svd(U, s, Vh) + + def _order_svd( + self, + U: DenseArray, + s: DenseArray, + Vh: DenseArray, + ) -> tuple[DenseArray, DenseArray, DenseArray]: + """Sort singular triples descending and gauge them, preserving ``U s Vh``.""" + k = self.shape(s)[-1] + order = self.xp.argsort(-s, axis=-1) + s = self._take_along_axis(s, order, -1) + + column_order = self.expand_dims(order, -2) + U_main = self._take_along_axis(U[..., :, :k], column_order, -1) + Vh_main = self._take_along_axis(Vh[..., :k, :], self.expand_dims(order, -1), -2) + + phase = self._column_phase(U_main) + U_main = U_main / phase + Vh_main = Vh_main * self.swapaxes(phase, -1, -2) + + U = self.concatenate([U_main, U[..., :, k:]], axis=-1) + Vh = self.concatenate([Vh_main, Vh[..., k:, :]], axis=-2) + return U, s, Vh def cholesky( self, @@ -987,8 +1186,13 @@ def diag(self, x: DenseArray) -> DenseArray: return self.xp.diag(x) def diagonal(self, x: DenseArray) -> DenseArray: - """Return the main diagonal of x (delegates to xp.diagonal).""" - return self.xp.diagonal(x) + """Return the main diagonal over the **trailing two axes**. + + Same correction as :meth:`trace`: the array libraries default to + ``axis1=0, axis2=1``, which on a ``(batch, n, n)`` stack diagonalizes + across the batch axis. Identical to the library default for 2-D input. + """ + return self.xp.diagonal(x, axis1=-2, axis2=-1) def tril(self, x: DenseArray) -> DenseArray: """Lower triangle of x (delegates to xp.tril).""" diff --git a/spacecore/backend/_optional.py b/spacecore/backend/_optional.py new file mode 100644 index 0000000..b5d7adb --- /dev/null +++ b/spacecore/backend/_optional.py @@ -0,0 +1,215 @@ +from __future__ import annotations + +import importlib +import warnings +from functools import lru_cache +from importlib import metadata + +from ._ops import BackendOps +from .numpy import NumpyOps + +# (submodule of spacecore.backend, ops attribute, dependency import name) +_OPTIONAL_OPS: tuple[tuple[str, str, str], ...] = ( + (".jax", "JaxOps", "jax"), + (".cupy", "CuPyOps", "cupy"), + (".torch", "TorchOps", "torch"), +) + +# Packaging entry-point group that external packages use to advertise a backend. +_ENTRY_POINT_GROUP = "spacecore.backends" + + +def _entry_points(group: str) -> list: + """Return the entry points in ``group``, across ``importlib.metadata`` APIs.""" + eps = metadata.entry_points() + select = getattr(eps, "select", None) + if select is not None: # Python 3.10+: EntryPoints.select(group=...) + return list(select(group=group)) + return list(eps.get(group, [])) # Python < 3.10: dict-of-lists API + + +def discover_entry_point_ops() -> list[type[BackendOps]]: + """Discover backend ops classes advertised by installed packages. + + A third-party package registers a SpaceCore backend by declaring a + ``spacecore.backends`` entry point pointing at a :class:`BackendOps` + subclass, e.g. in its ``pyproject.toml``:: + + [project.entry-points."spacecore.backends"] + mlx = "spacecore_mlx:MLXOps" + + This makes the backend layer *open for extension* — a new backend is added by + installing a package, not by editing SpaceCore. Each entry point is loaded + defensively: a load failure, or a target that is not a ``BackendOps`` + subclass, is skipped with a :class:`UserWarning` rather than aborting + ``import spacecore`` — the same non-fatal contract as :func:`import_backend`. + """ + discovered: list[type[BackendOps]] = [] + for ep in _entry_points(_ENTRY_POINT_GROUP): + name = getattr(ep, "name", ep) + try: + obj = ep.load() + except Exception as exc: # noqa: BLE001 - a broken plugin must not be fatal + warnings.warn( + f"SpaceCore backend entry point {name!r} failed to load " + f"({type(exc).__name__}: {exc}); skipping it.", + stacklevel=2, + ) + continue + if not (isinstance(obj, type) and issubclass(obj, BackendOps)): + warnings.warn( + f"SpaceCore backend entry point {name!r} does not point at a " + f"BackendOps subclass (got {obj!r}); skipping it.", + stacklevel=2, + ) + continue + discovered.append(obj) + return discovered + + +def _backend_absent(exc: ModuleNotFoundError, dep: str) -> bool: + """Return True iff the failure is the backend dependency itself being missing.""" + return exc.name == dep + + +def import_backend(module: str, dep: str): + """Import an optional backend submodule of ``spacecore.backend``. + + Parameters + ---------- + module : str + Submodule to import, relative to ``spacecore.backend`` (for example + ``".jax"``). + dep : str + Import name of the optional third-party dependency the submodule + needs (for example ``"jax"``). + + Returns + ------- + module or None + The imported module, or ``None`` when ``dep`` is not installed. + + Notes + ----- + An *absent* dependency (``dep`` not installed) is turned into ``None`` + silently — that is the normal "backend not present" case. Any *other* + import failure means the backend is installed but cannot load: broken + against a mismatched runtime, shadowed by a namespace shim, a partial + install, or a missing transitive dependency. Those are also non-fatal + (an optional backend must never abort ``import spacecore``) but they are + surfaced with a :class:`UserWarning` rather than hidden, so a genuinely + broken install is diagnosable instead of invisible. + """ + try: + return importlib.import_module(module, package=__package__) + except ImportError as exc: + if isinstance(exc, ModuleNotFoundError) and _backend_absent(exc, dep): + return None + warnings.warn( + f"Optional backend {module!r} is installed but failed to import " + f"({type(exc).__name__}: {exc}); skipping it.", + stacklevel=2, + ) + return None + + +@lru_cache(maxsize=1) +def available_ops() -> tuple[type[BackendOps], ...]: + """Backend ops classes whose optional dependency is importable. + + NumPy is always present. Each built-in optional backend is attempted; a + missing dependency is skipped, while any *other* import failure warns and is + skipped. External backends advertised via ``spacecore.backends`` entry points + (see :func:`discover_entry_point_ops`) are appended last; a built-in family + takes precedence, so a plugin cannot shadow (for example) ``"numpy"``. + + **Cached for the process.** Discovery is not free — it attempts a real import + per optional backend and scans packaging metadata for entry points — and it is + called from several places during ``import spacecore``. Without the cache a + present-but-broken backend emits its warning once per call, so one broken + install reads as several distinct problems. + + The result is a tuple, not a list, precisely because it is shared: a cached + mutable sequence would let one caller's edit reach every other caller. + + Availability is a property of the environment and does not change within a + process, so the cache never needs invalidating in normal use. Registering a + backend at runtime goes to :class:`~spacecore.backend.OpsRegistry` and does not + pass through here. Tests that monkeypatch the import machinery or the entry + points must call ``available_ops.cache_clear()``. + + Returns + ------- + tuple of type of BackendOps + Available ops classes, always starting with :class:`NumpyOps`. + """ + ops: list[type[BackendOps]] = [NumpyOps] + for module, attr, dep in _OPTIONAL_OPS: + mod = import_backend(module, dep) + if mod is None: + continue + cls = getattr(mod, attr, None) + if cls is not None: + ops.append(cls) + seen = {o._family for o in ops} + for cls in discover_entry_point_ops(): + if cls._family not in seen: + ops.append(cls) + seen.add(cls._family) + return tuple(ops) + + +def available_families() -> tuple[str, ...]: + """Backend family names whose optional dependency is importable. + + Returns + ------- + tuple of str + Lowercase family names (for example ``("numpy", "jax")``), always + starting with ``"numpy"``. + """ + return tuple(o._family for o in available_ops()) + + +def is_available(family: str) -> bool: + """Whether a backend ``family`` can be imported in this environment. + + Parameters + ---------- + family : str + Lowercase backend family name (for example ``"jax"``). + + Returns + ------- + bool + ``True`` if the family's ops class is importable. + """ + return family in available_families() + + +def require_backend(family: str) -> type[BackendOps]: + """Return the ops class for ``family`` or raise an actionable error. + + Parameters + ---------- + family : str + Lowercase backend family name (for example ``"jax"``). + + Returns + ------- + type of BackendOps + The ops class for ``family``. + + Raises + ------ + ImportError + If the family's optional dependency is not installed, with a message + pointing at the corresponding ``spacecore`` extra. + """ + for o in available_ops(): + if o._family == family: + return o + raise ImportError( + f"The {family!r} backend requires an optional dependency that is not " + f"installed. Install it with: pip install spacecore[{family}]" + ) diff --git a/spacecore/backend/_registry.py b/spacecore/backend/_registry.py new file mode 100644 index 0000000..24f820a --- /dev/null +++ b/spacecore/backend/_registry.py @@ -0,0 +1,193 @@ +"""Registry of backend implementations, keyed by backend family. + +The single answer to "which ``BackendOps`` classes exist in this process". It is +seeded at construction from :func:`~spacecore.backend._optional.available_ops` +— what is *importable* in this environment — and can then be extended at runtime +via :meth:`OpsRegistry.register`, which is how a third-party backend joins +without an entry point. + +Separation of concerns: this module answers *what backends exist*; +:mod:`spacecore.contextual` answers *which one is currently in effect*. Those are +different kinds of state — a registry is process-wide and shared by every thread, +while ambient policy is scoped and swappable — and holding them in one object +made the registry impossible to test without the ambient singleton. Compare +:mod:`spacecore.backend._container`, which applies the same split to pytree +registration. + +Two deliberate differences from :class:`~spacecore.backend._container.PyTreeRegistry`: + +* **Duplicate registration raises** here rather than being idempotent. Two + different classes claiming ``"numpy"`` is a genuine conflict a caller should + hear about, whereas re-wiring an already-wired ``(backend, class)`` pair is a + no-op by construction. +* **Removal is supported** (:meth:`unregister`). ``PyTreeRegistry`` must stay + monotonic because its two-axis back-fill depends on it; this registry is a + plain keyed lookup with no derived state, so removing an entry is safe. It + exists for tests and for backends registered dynamically. +""" +from __future__ import annotations + +import threading +from types import MappingProxyType +from typing import Any, Mapping + +from .._errors import ContextConflictError, UnknownBackendError +from ._family import BackendFamily +from ._ops import BackendOps + +#: Anything that can name a backend: a family string, the enum, an ops instance +#: or class, or a Context (duck-typed via ``.ops`` to avoid importing contextual). +BackendKeyLike = Any + + +def backend_key(x: BackendKeyLike) -> str: + """Normalize any backend-naming value to its lowercase family key. + + Accepts a family string (``"pytorch"`` is folded to ``"torch"``), a + :class:`BackendFamily`, a :class:`BackendOps` instance or subclass, or any + object exposing a ``.ops`` attribute (a ``Context``; duck-typed so this + module need not import :mod:`spacecore.contextual`, which would be a cycle). + """ + if isinstance(x, BackendOps): + return backend_key(x.family) + if isinstance(x, type) and issubclass(x, BackendOps): + return backend_key(x._family) + if isinstance(x, BackendFamily): + return x.value.lower() + if isinstance(x, str): + key = x.lower() + return "torch" if key == "pytorch" else key + ops = getattr(x, "ops", None) + if ops is not None and isinstance(ops, BackendOps): + return backend_key(ops) + raise TypeError(f"Unsupported backend key source: {type(x)!r}") + + +class OpsRegistry: + """Process-wide map from backend family name to its ``BackendOps`` class. + + Instantiable so tests can exercise registration without touching the shared + instance; library code uses the module-level :data:`registry`. + + Parameters + ---------- + seed : iterable of type of BackendOps, optional + Classes to register at construction, typically the result of + ``available_ops()``. + """ + + def __init__(self, seed: Any = ()) -> None: + self._ops: dict[str, type[BackendOps]] = {} + self._lock = threading.RLock() + for cls in seed: + self.register(cls) + + def register(self, ops: type[BackendOps]) -> type[BackendOps]: + """Register a backend class under its family name. + + Parameters + ---------- + ops : type of BackendOps + Backend class to register. + + Returns + ------- + type of BackendOps + The class, so this can be used as a decorator. + + Raises + ------ + TypeError + If ``ops`` is not a ``BackendOps`` subclass. + ContextConflictError + If the family is already registered — including by this same class. + Re-registration is treated as a conflict rather than a no-op because + it almost always means two implementations are competing for one + family name. + """ + if not isinstance(ops, type) or not issubclass(ops, BackendOps): + raise TypeError(f"Expected type[BackendOps], got {type(ops)!r}") + family = backend_key(ops) + with self._lock: + if family in self._ops: + raise ContextConflictError(f"BackendOps {family} is already registered.") + self._ops[family] = ops + return ops + + def unregister(self, key: BackendKeyLike) -> type[BackendOps] | None: + """Remove a family's registration and return it, or ``None`` if absent. + + Safe because this registry holds no derived state (contrast + :class:`~spacecore.backend._container.PyTreeRegistry`, whose back-fill + requires monotonicity). Intended for tests and dynamically scoped + backends, not for routine use. + """ + with self._lock: + return self._ops.pop(backend_key(key), None) + + def get_class(self, key: BackendKeyLike) -> type[BackendOps]: + """Return the registered class for ``key``. + + Raises + ------ + UnknownBackendError + If no backend is registered under that family. + """ + family = backend_key(key) + with self._lock: + cls = self._ops.get(family) + if cls is None: + allowed = ", ".join(self._ops) + raise UnknownBackendError( + f"Unknown backend: {family!r}. Expected one of: {allowed}" + ) + return cls + + def get(self, key: BackendKeyLike) -> BackendOps: + """Return a fresh instance of the backend registered under ``key``.""" + return self.get_class(key)() + + def classes(self) -> Mapping[str, type[BackendOps]]: + """Return a read-only view of family name to registered class.""" + with self._lock: + return MappingProxyType(dict(self._ops)) + + def families(self) -> tuple[str, ...]: + """Return the registered family names, in registration order.""" + with self._lock: + return tuple(self._ops) + + def __contains__(self, key: BackendKeyLike) -> bool: + try: + return backend_key(key) in self._ops + except TypeError: + return False + + def match(self, x: Any) -> list[BackendOps]: + """Return every registered backend that recognizes ``x`` as one of its arrays. + + The reverse lookup behind context inference. Instantiation or the + ``is_array`` probe failing is treated as "not a match" rather than an + error, keeping inference conservative: a backend that cannot answer + simply does not claim the value. + """ + matched: list[BackendOps] = [] + for cls in tuple(self.classes().values()): + try: + ops = cls() + if ops.is_array(x): + matched.append(ops) + except Exception: + continue + return matched + + +def _seeded_registry() -> OpsRegistry: + """Build the process registry from whatever backends are importable.""" + from ._optional import available_ops + + return OpsRegistry(available_ops()) + + +#: Process-wide backend registry used by :mod:`spacecore.contextual`. +registry = _seeded_registry() diff --git a/spacecore/backend/jax/__init__.py b/spacecore/backend/jax/__init__.py index 5f48e04..51f43bc 100644 --- a/spacecore/backend/jax/__init__.py +++ b/spacecore/backend/jax/__init__.py @@ -1,6 +1,11 @@ -"""JAX backend implementation and pytree registration helpers.""" +"""JAX backend implementation. -from ._pytree import jax_pytree_class as jax_pytree_class +Nothing here is imported eagerly by :mod:`spacecore.backend`; the package is +reached only through ``_optional.available_ops``, whose guarded import skips a +missing dependency and warns on a broken one. Pytree registration now lives on +:meth:`JaxOps.install_pytree_protocol`, driven by the backend-neutral registry in +:mod:`spacecore.backend._container`. +""" try: from ._ops import JaxOps as JaxOps @@ -8,7 +13,7 @@ if exc.name != "jax": raise -__all__ = ["jax_pytree_class"] +__all__ = [] if "JaxOps" in globals(): __all__.append("JaxOps") diff --git a/spacecore/backend/jax/_ops.py b/spacecore/backend/jax/_ops.py index c6db99b..57d1924 100644 --- a/spacecore/backend/jax/_ops.py +++ b/spacecore/backend/jax/_ops.py @@ -75,6 +75,21 @@ class JaxOps(BackendOps): def __init__(self) -> None: super().__init__() + @classmethod + def install_pytree_protocol(cls) -> None: + """Register SpaceCore containers as JAX pytree nodes. + + JAX's ``register_pytree_node_class`` consumes exactly the protocol + :class:`~spacecore.backend._container.PyTreeNode` defines — ``tree_flatten`` + returning ``(children, aux)`` and a ``tree_unflatten(aux, children)`` + classmethod — so it *is* the registrar and no translation is needed. + """ + import jax + + from .._container import registry + + registry.register_backend("jax", jax.tree_util.register_pytree_node_class) + def sanitize_dtype(self, dtype: DType | None) -> DType: """ Normalize a dtype specifier using JAX. diff --git a/spacecore/backend/jax/_pytree.py b/spacecore/backend/jax/_pytree.py deleted file mode 100644 index 63fbc58..0000000 --- a/spacecore/backend/jax/_pytree.py +++ /dev/null @@ -1,32 +0,0 @@ -from __future__ import annotations - -from typing import TypeVar - -T = TypeVar("T", bound=type) - - -def jax_pytree_class(klass: T) -> T: - """ - Mark a class as a JAX PyTree node, if JAX is available. - - Safe to import without JAX installed. - - Parameters - ---------- - klass : type - Class implementing JAX pytree methods. - - Returns - ------- - type - Registered class when JAX is available, otherwise ``klass`` unchanged. - """ - try: - from jax import tree_util - except Exception: - return klass - try: - tree_util.register_pytree_node_class(klass) - except Exception: - pass - return klass diff --git a/spacecore/backend/torch/_ops.py b/spacecore/backend/torch/_ops.py index f2b94f3..bf19379 100644 --- a/spacecore/backend/torch/_ops.py +++ b/spacecore/backend/torch/_ops.py @@ -77,6 +77,43 @@ class TorchOps(EagerControlFlowMixin, BackendOps): def __init__(self) -> None: super().__init__() + @classmethod + def install_pytree_protocol(cls) -> None: + """Register SpaceCore containers as Torch pytree nodes. + + Unlike the JAX adapter, this one has real translating to do: Torch's + registry expects children as a ``list`` and calls ``unflatten(children, + context)``, the reverse of the ``(aux, children)`` order + :class:`~spacecore.backend._container.PyTreeNode` defines (which follows + JAX). Absorbing that mismatch here is the point of a per-backend adapter + — no container class has to know either convention. + + Registering with ``torch.utils._pytree`` also covers + ``torch.utils._cxx_pytree`` (the optree-backed implementation selected by + ``PYTORCH_USE_CXX_PYTREE=1``): Torch mirrors registrations across the two, + so a single call serves both and registering twice would raise. + + ``torch.utils._pytree`` is private API — there is no public alias in + Torch 2.x — so this is a deliberate coupling to a Torch internal. It is + confined to this adapter, and a breaking change there degrades to a + warning from the registry rather than an import failure. + """ + import torch.utils._pytree as torch_pytree + + from .._container import registry + + def registrar(klass: type) -> None: + def flatten(obj: Any) -> tuple[list[Any], Any]: + children, aux = obj.tree_flatten() + return list(children), aux + + def unflatten(children: Any, aux: Any) -> Any: + return klass.tree_unflatten(aux, tuple(children)) + + torch_pytree.register_pytree_node(klass, flatten, unflatten) + + registry.register_backend("torch", registrar) + @staticmethod def _defined_kwargs(**kwargs: Any) -> dict[str, Any]: return {key: value for key, value in kwargs.items() if value is not None} @@ -462,7 +499,10 @@ def eigh( raise TypeError("eigh requires a dense array; sparse input is not supported.") kwargs = {} if backend_kwargs is None else dict(backend_kwargs) kwargs.update(self._defined_kwargs(out=out)) - return self.torch.linalg.eigh(x, UPLO=UPLO, **kwargs) + eigenvalues, eigenvectors = self.torch.linalg.eigh(x, UPLO=UPLO, **kwargs) + # Same ordering/gauge normalization as the neutral method: Torch's raw + # complex Hermitian eigenvectors come back as -1 times NumPy's and JAX's. + return self._order_eigenpairs(eigenvalues, eigenvectors) def norm( self, @@ -483,6 +523,14 @@ def norm( out=out, ) + def diagonal(self, x: DenseArray) -> DenseArray: + # torch.diagonal spells the axes dim1/dim2; same trailing-two-axes + # convention as the neutral method. + return self.torch.diagonal(x, dim1=-2, dim2=-1) + + def _take_along_axis(self, x: DenseArray, indices: DenseArray, axis: int) -> DenseArray: + return self.torch.take_along_dim(x, indices, dim=axis) + def solve( self, A: DenseArray, @@ -507,7 +555,8 @@ def svd( ) -> tuple[DenseArray, DenseArray, DenseArray]: kwargs = {} if backend_kwargs is None else dict(backend_kwargs) kwargs.update(self._defined_kwargs(driver=driver, out=out)) - return self.torch.linalg.svd(A, full_matrices=full_matrices, **kwargs) + U, s, Vh = self.torch.linalg.svd(A, full_matrices=full_matrices, **kwargs) + return self._order_svd(U, s, Vh) def cholesky( self, diff --git a/spacecore/_contextual/__init__.py b/spacecore/contextual/__init__.py similarity index 51% rename from spacecore/_contextual/__init__.py rename to spacecore/contextual/__init__.py index 63b824f..aa360f9 100644 --- a/spacecore/_contextual/__init__.py +++ b/spacecore/contextual/__init__.py @@ -1,14 +1,29 @@ +"""Backend/dtype context, context-bound objects, and ambient context state. + +This package is the public home of :class:`Context` (which backend ops and +dtype a value is expressed in) and :class:`ContextBound` (the mixin every +space, operator, and functional derives from). Validation policy lives on the +bound object as ``check_level``, not on the context; the module-level +``get_check_level`` / ``set_check_level`` / ``use_check_level`` helpers seed +and scope that default. +""" + +from ._context import Context from ._bound import ContextBound as ContextBound from ._state import ( enforce_convert_policy as enforce_convert_policy, + get_check_level as get_check_level, get_context as get_context, normalize_context as normalize_context, normalize_ops as normalize_ops, register_ops as register_ops, resolve_context_priority as resolve_context_priority, + set_check_level as set_check_level, set_context as set_context, + use_check_level as use_check_level, + use_context as use_context, ) -from ._policies import ( +from .._errors import ( ContextConflictError as ContextConflictError, ContextError as ContextError, ContextInferenceError as ContextInferenceError, @@ -16,16 +31,21 @@ ) __all__ = [ + "Context", "ContextBound", "ContextConflictError", "ContextError", "ContextInferenceError", "UnknownBackendError", "enforce_convert_policy", + "get_check_level", "get_context", "normalize_context", "normalize_ops", "register_ops", "resolve_context_priority", + "set_check_level", "set_context", + "use_check_level", + "use_context", ] diff --git a/spacecore/_contextual/_bound.py b/spacecore/contextual/_bound.py similarity index 63% rename from spacecore/_contextual/_bound.py rename to spacecore/contextual/_bound.py index b0c02dd..294b938 100644 --- a/spacecore/_contextual/_bound.py +++ b/spacecore/contextual/_bound.py @@ -1,20 +1,25 @@ from __future__ import annotations from abc import ABC -from typing import TYPE_CHECKING, Any, Self - -from .._check_policy import CheckLevel, check_level_at_least, level_to_enabled +from typing import Any, Self + +from ..backend import BackendFamily, BackendOps +from .._check_policy import ( + CheckLevel, + check_level_at_least, + level_to_enabled, + minimum_check_level, + normalize_check_level, +) from .._repr import format_dtype from ..types import DType -from ._state import enforce_convert_policy, normalize_context, resolve_context_priority - -if TYPE_CHECKING: - from ..backend import BackendFamily, BackendOps, Context - - -def _same_math_context(left: Context, right: Context) -> bool: - """Return whether contexts match for algebra, ignoring validation checks.""" - return left.ops == right.ops and left.dtype == right.dtype +from ._state import ( + enforce_convert_policy, + get_check_level, + normalize_context, + resolve_context_priority, +) +from ._context import Context class ContextBound(ABC): @@ -24,30 +29,49 @@ class ContextBound(ABC): Parameters ---------- ctx : Context, str, or None, optional - Context specification used to resolve backend operations, dtype, and - validation policy. + Context specification used to resolve backend operations and dtype. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. """ - def __init__(self, ctx: Context | str | None = None): + def __init__( + self, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ): ctx = normalize_context(ctx) self._ctx = ctx - - def _eq_backend_compatible(self, other: Any) -> bool: - """Tier-1 equality gate: same concrete type and same backend. - - "Backend compatibility" means the same backend ops family and the same - representation dtype (via :func:`_same_math_context`), deliberately - ignoring ``check_level`` (a validation policy, not a mathematical or - backend property). This is the mandatory first check in every ``__eq__``: - it guarantees any later ``ops.allclose`` runs only on same-backend, - same-dtype arrays, and makes cross-backend objects compare unequal. + self._check_level = ( + normalize_check_level(check_level) if check_level is not None else get_check_level() + ) + + def same_math(self, other: Any) -> bool: + """Tier-1 equality gate: same concrete type and same math context. + + Combines a runtime-type identity check with :meth:`Context.same_math` + (equal backend ops and dtype, ignoring ``check_level`` — a validation + policy, not a mathematical or backend property). This is the mandatory + first check in every ``__eq__``: the type gate guarantees ``other`` + exposes the same attributes this object's ``__eq__`` will read, and + ``same_math`` guarantees any later ``ops.allclose`` runs only on + same-backend, same-dtype arrays, making cross-backend objects compare + unequal. Callers that fail this gate should return ``NotImplemented`` so Python can try the reflected comparison and fall back to identity symmetrically. """ - return type(self) is type(other) and _same_math_context(self._ctx, other._ctx) - - def _bind_context(self, ctx: Any, *children: Any, sources: tuple[Any, ...] | None = None): + return type(self) is type(other) and self._ctx.same_math(other.ctx) + + def _bind_context( + self, + ctx: Any, + *children: Any, + sources: tuple[Any, ...] | None = None, + check_level: CheckLevel | bool | None = None, + ): """Resolve the priority context, store it, and re-bind children onto it. This factors the context-binding prologue shared by every contextual @@ -77,8 +101,6 @@ def _bind_context(self, ctx: Any, *children: Any, sources: tuple[Any, ...] | Non tuple ``children`` converted onto the resolved context, in order. """ - from ..backend import Context - if isinstance(ctx, Context): resolved = ctx else: @@ -86,6 +108,11 @@ def _bind_context(self, ctx: Any, *children: Any, sources: tuple[Any, ...] | Non ctx, *(children if sources is None else sources) ) self._ctx = resolved + if check_level is not None: + self._check_level = normalize_check_level(check_level) + else: + levels = [c._check_level for c in children if isinstance(c, ContextBound)] + self._check_level = minimum_check_level(tuple(levels)) if levels else get_check_level() return tuple(child.convert(resolved) for child in children) @property @@ -105,8 +132,15 @@ def ctx(self) -> Context: @property def check_level(self) -> CheckLevel: - """Return this object's runtime validation level.""" - return self.ctx.check_level + """Return this object's runtime validation level. + + The level is stored on the bound object (``_check_level``), seeded at + construction. The fallback covers any construction path that assigns + ``_ctx`` without going through ``__init__``/``_bind_context``. + """ + if hasattr(self, "_check_level"): + return self._check_level + return get_check_level() @property def _enable_checks(self) -> bool: @@ -166,8 +200,15 @@ def _convert(self, new_ctx: Context) -> Self: raise NotImplementedError() def convert(self, new_ctx: Context | BackendFamily | str | None = None) -> Self: - """Return this object represented in ``new_ctx``.""" + """Return this object represented in ``new_ctx``. + + The object's ``check_level`` is a property of the object, not of the + backend/dtype context, so it is preserved across conversion rather than + re-seeded from the ambient default. + """ _, new_ctx = enforce_convert_policy(self, new_ctx) if self.ctx == new_ctx: return self - return self._convert(new_ctx) + result = self._convert(new_ctx) + result._check_level = self._check_level + return result diff --git a/spacecore/backend/_context.py b/spacecore/contextual/_context.py similarity index 58% rename from spacecore/backend/_context.py rename to spacecore/contextual/_context.py index f1013c5..3c5e01b 100644 --- a/spacecore/backend/_context.py +++ b/spacecore/contextual/_context.py @@ -1,21 +1,24 @@ from dataclasses import dataclass from typing import Any -from .._check_policy import CheckLevel, level_to_enabled, normalize_check_level -from ._ops import BackendOps +from ..backend import BackendOps from ..types import DenseArray, SparseArray, DType, ArrayLike @dataclass(frozen=True, slots=True, init=False) class Context: """ - Select backend operations, representation dtype, and validation policy. + Select backend operations and representation dtype. - A context collects the backend operations object, default dtype, and runtime - validation policy used by spaces, linear operators, and context-bound - values. It is intentionally small: it does not own arrays, but it defines - how new arrays are created and how existing arrays are checked or converted - for a backend family. + A context collects the backend operations object and default dtype used by + spaces, linear operators, and context-bound values. It is intentionally + small: it does not own arrays, but it defines how new arrays are created and + how existing arrays are checked or converted for a backend family. + + Validation policy is deliberately *not* part of a context. ``check_level`` is + a property of the context-bound object (space, operator, functional) — see + :class:`spacecore.contextual.ContextBound` — because two objects on the same + backend and dtype may legitimately validate at different strictness. Parameters ---------- @@ -30,11 +33,6 @@ class Context: ``ops.sanitize_dtype`` during initialization. It does not independently define a mathematical scalar field; spaces expose that contract through :attr:`spacecore.space.Space.field`. - enable_checks : bool or None, optional - Deprecated compatibility alias. ``True`` maps to ``"standard"`` and - ``False`` maps to ``"none"``. Passing both policy arguments is an error. - check_level : {"none", "cheap", "standard", "strict"}, optional - Runtime validation policy. The default is ``"standard"``. Attributes ---------- @@ -48,7 +46,12 @@ class Context: ``Context`` is frozen and slot-based. Methods that convert values return new backend arrays or sparse objects; they do not mutate the context itself. - Equality compares backend family, dtype, and ``check_level``. + Equality (:meth:`__eq__`) and :meth:`same_math` both compare backend ops and + dtype. They coincide now that ``check_level`` has moved off the context — the + two are kept distinct because they answer different questions at their call + sites (object identity versus "may these be combined algebraically"), and + only :meth:`same_math` is part of the ``ContextBound`` equality contract. + :meth:`same_backend` remains strictly coarser, ignoring dtype. Examples -------- @@ -64,16 +67,8 @@ class Context: ops: BackendOps dtype: DType | None - check_level: CheckLevel - - def __init__( - self, - ops: BackendOps, - dtype: DType | None = None, - enable_checks: bool | None = None, - *, - check_level: CheckLevel | None = None, - ) -> None: + + def __init__(self, ops: BackendOps, dtype: DType | None = None) -> None: """ Validate and normalize the context after dataclass initialization. @@ -82,7 +77,7 @@ def __init__( TypeError If ``ops`` is not a :class:`BackendOps` instance. """ - from .._contextual._state import normalize_ops + from ._state import normalize_ops try: ops = normalize_ops(ops) @@ -90,20 +85,6 @@ def __init__( raise TypeError("Unknown ops type.") object.__setattr__(self, "ops", ops) object.__setattr__(self, "dtype", self.ops.sanitize_dtype(dtype)) - object.__setattr__( - self, - "check_level", - normalize_check_level( - check_level, - enable_checks=enable_checks, - warn_legacy=enable_checks is not None, - ), - ) - - @property - def enable_checks(self) -> bool: - """Deprecated Boolean view of :attr:`check_level`.""" - return level_to_enabled(self.check_level) def assert_dense(self, x: Any) -> DenseArray: """ @@ -229,6 +210,71 @@ def convert(self, x: Any) -> ArrayLike: else: raise NotImplementedError + def same_math(self, other: Any) -> bool: + """ + Return whether ``other`` shares this context's mathematical backend. + + Two contexts "share a math context" when they have equal backend + ``ops`` and ``dtype``. Validation policy is not consulted, being a + property of the bound object rather than of the context. This is a + *fixed* equivalence relation — currently coextensive with :meth:`__eq__` + — and the first-tier gate in every context-bound + ``__eq__`` (via :meth:`ContextBound.same_math`): it + guarantees any later elementwise ``ops.allclose`` runs only on + same-backend, same-dtype arrays. + + It is intentionally *not* policy-configurable. Compatibility decisions + that may legitimately vary — family-only matching, dtype promotion — + belong to the context-resolution policy + (``Contextual.are_compatible_contexts``), which may use ``same_math`` + as its strict floor. + + Parameters + ---------- + other: + Object to compare against. + + Returns + ------- + bool + ``True`` when ``other`` is a ``Context`` with equal backend + operations and dtype. Validation policy is not consulted: it lives + on the bound object, not the context. + """ + if isinstance(other, Context): + return self.ops == other.ops and self.dtype == other.dtype + return False + + def same_backend(self, other: Any) -> bool: + """ + Return whether ``other`` runs on the same backend as this context. + + Two contexts "share a backend" when they have equal backend ``ops`` + (i.e. the same backend family — :class:`BackendOps` equality compares + ``family``), ignoring ``dtype``. This is the loosest of the three context + relations and the coarsest point of the strictness chain ``__eq__`` ≡ + :meth:`same_math` (ops ∧ dtype) ⊃ ``same_backend`` (ops). + + It is the dtype-agnostic notion used by the context-resolution policy + (``Contextual.are_compatible_contexts``): operands on the same backend + are combinable, with any dtype difference reconciled by conversion onto + the resolved context. + + Parameters + ---------- + other: + Object to compare against. + + Returns + ------- + bool + ``True`` when ``other`` is a ``Context`` on the same backend, + regardless of ``dtype``. + """ + if isinstance(other, Context): + return self.ops == other.ops + return False + def __eq__(self, other: Any) -> bool: """ Return whether another object has the same execution context. @@ -242,18 +288,12 @@ def __eq__(self, other: Any) -> bool: ------- bool ``True`` when ``other`` is a ``Context`` with equal backend - operations, dtype, and ``check_level``. + operations and dtype. Validation policy (``check_level``) is not a + property of the context — it lives on the context-bound object. """ if isinstance(other, Context): - return ( - self.ops == other.ops - and self.dtype == other.dtype - and self.check_level == other.check_level - ) + return self.ops == other.ops and self.dtype == other.dtype return False def __repr__(self) -> str: - return ( - f"Context(ops={self.ops!r}, dtype={self.dtype!r}, " - f"check_level={self.check_level!r})" - ) + return f"Context(ops={self.ops!r}, dtype={self.dtype!r})" diff --git a/spacecore/contextual/_contextual.py b/spacecore/contextual/_contextual.py new file mode 100644 index 0000000..ee4da92 --- /dev/null +++ b/spacecore/contextual/_contextual.py @@ -0,0 +1,260 @@ +from __future__ import annotations + +from contextlib import contextmanager +from contextvars import ContextVar +from typing import Any, Iterable, Iterator, Tuple +from warnings import warn + +from .._check_policy import CheckLevel, normalize_check_level +from ..types import DType +from ..backend import BackendFamily, BackendOps, NumpyOps, OpsRegistry +from ..backend import ops_registry as _default_ops_registry +from ._context import Context +from .._errors import ContextInferenceError + + +# Scoped ambient overrides installed by ``Contextual.scoped_ctx`` / +# ``scoped_check_level`` (public entry points: ``use_context`` / +# ``use_check_level``). +# +# Module level by design. ``ContextVar`` objects are never garbage collected, so +# the contextvars documentation requires creating them at module scope rather +# than per call or per instance. Sharing one pair across ``Contextual`` instances +# is also the intended semantics: an override is ambient to the current thread / +# async task, not a property of whichever resolver observes it. +# +# ``None`` is the sentinel for "no override installed here" — readers fall back to +# the process-wide baseline held on the instance. +_ctx_override: ContextVar[Context | None] = ContextVar("spacecore_ctx", default=None) +_check_override: ContextVar[CheckLevel | None] = ContextVar( + "spacecore_check_level", default=None +) + + +class Contextual: + """Resolve which context and validation level are currently in effect. + + Holds **ambient policy only** — the process-wide baseline context and check + level, each read through a scoped :class:`~contextvars.ContextVar` override + before falling back to the baseline stored here. + + *Which backends exist* is a different kind of state and lives in + :class:`~spacecore.backend.OpsRegistry`: a registry is shared by every thread + and must never be scoped (a backend registered inside a ``with`` block that + vanished on exit would be a bug), whereas ambient policy is swappable by + design. This object consults the registry; it does not own it. + + Parameters + ---------- + ops_registry : OpsRegistry, optional + Backend registry to resolve names against. Defaults to the process-wide + instance; injectable so tests can resolve against their own. + """ + + _baseline_ctx: Context + _default_dtype: DType | None = None + # Baseline ambient validation level seeded onto newly created bound objects. + # check_level is a property of the bound object, not of the Context. + _baseline_check_level: CheckLevel = "standard" + + def __init__(self, ops_registry: OpsRegistry | None = None) -> None: + ops = NumpyOps() + self._baseline_ctx = Context(ops=ops, dtype=ops.sanitize_dtype(self._default_dtype)) + self._baseline_check_level = type(self)._baseline_check_level + self._ops_registry = ops_registry if ops_registry is not None else _default_ops_registry + + @property + def ops_registry(self) -> OpsRegistry: + """Backend registry this resolver consults.""" + return self._ops_registry + + def get_check_level(self) -> CheckLevel: + """Return the active ambient validation level for new bound objects. + + Resolution order: the scoped override installed by + :meth:`scoped_check_level`, then the process-wide baseline. + """ + level = _check_override.get() + return level if level is not None else self._baseline_check_level + + def set_check_level(self, level: CheckLevel | bool | None) -> None: + """Set the process-wide *baseline* validation level for new bound objects. + + Deliberately does not disturb scoped overrides: a ``scoped_check_level`` + block already in progress keeps winning until it exits. + """ + self._baseline_check_level = normalize_check_level(level) + + @contextmanager + def scoped_check_level(self, level: CheckLevel | bool | None) -> Iterator[CheckLevel]: + """Override the ambient validation level for the duration of the block. + + The override is visible only to the current thread / async task. The + :class:`~contextvars.Token` returned by ``set`` unwinds exactly this + override, so nesting and concurrent scopes cannot clobber one another. + """ + resolved = normalize_check_level(level) + token = _check_override.set(resolved) + try: + yield resolved + finally: + _check_override.reset(token) + + def normalize_context( + self, + ctx: Context | BackendFamily | str | None = None, + dtype: Any = None, + ) -> Context: + if ctx is None: + if dtype is not None: + warn("Provided context is None; dtype is ignored.", UserWarning) + return self.default_ctx + if isinstance(ctx, Context): + if dtype is not None: + warn("Provided concrete context; dtype is ignored.", UserWarning) + return Context(ops=ctx.ops, dtype=ctx.ops.sanitize_dtype(ctx.dtype)) + if isinstance(ctx, (str, BackendFamily)): + ops = self._ops_registry.get(ctx) + return self.ctx_from_ops(ops, dtype=dtype) + else: + raise TypeError(f"Expected Context, BackendFamily, str, or None, got {type(ctx)}.") + + def ctx_from_ops(self, ops: BackendOps, dtype: DType | None = None) -> Context: + return Context(ops=ops, dtype=ops.sanitize_dtype(dtype)) + + @property + def default_ctx(self) -> Context: + """Return the active default context. + + This is the scoped override if one is installed here, otherwise the + process-wide baseline. + + Every internal resolution path (:meth:`normalize_context` with ``None``, + :meth:`resolve_context_priority`, :meth:`infer_context`) reads through this + property, so a scoped override is honoured everywhere rather than only at + the public ``get_context`` boundary. + """ + ctx = _ctx_override.get() + return ctx if ctx is not None else self._baseline_ctx + + @default_ctx.setter + def default_ctx(self, ctx: Context | BackendFamily | str | None = None) -> None: + """Set the process-wide *baseline* context; scoped overrides still win.""" + self._baseline_ctx = self.normalize_context(ctx) + + @contextmanager + def scoped_ctx( + self, + ctx: Context | BackendFamily | str | None = None, + dtype: Any = None, + ) -> Iterator[Context]: + """Override the default context for the duration of the block. + + The override is visible only to the current thread / async task; see + :meth:`scoped_check_level` for the Token-unwind rationale. + """ + resolved = self.normalize_context(ctx, dtype=dtype) + token = _ctx_override.set(resolved) + try: + yield resolved + finally: + _ctx_override.reset(token) + + def infer_context(self, x: Any) -> Context | None: + """Infer context from `.ctx` first, then registered backend arrays. + + The reverse lookup — which registered backend claims this array — is the + registry's job (:meth:`~spacecore.backend.OpsRegistry.match`); turning the + match into a :class:`Context` is this object's. + """ + if isinstance(x, Context): + return x + + ctx = getattr(x, "ctx", None) + if isinstance(ctx, Context): + return ctx + + matched = self._ops_registry.match(x) + + if not matched: + return None + if len(matched) > 1: + raise ContextInferenceError( + f"Ambiguous backend inference for object of type {type(x)!r}: {matched!r}." + ) + + ops = matched[0] + try: + dtype = ops.get_dtype(x) + except Exception: + dtype = getattr(x, "dtype", self.default_ctx.dtype) + + return self.ctx_from_ops(ops, dtype) + + def infer_contexts(self, values: Iterable[Any]) -> Tuple[Context, ...]: + out: list[Context] = [] + for x in values: + ctx = self.infer_context(x) + if ctx is not None: + out.append(ctx) + return tuple(out) + + def are_compatible_contexts(self, *ctxs: Context) -> bool: + if len(ctxs) < 2: + return True + first = ctxs[0] + return all(ctx.same_backend(first) for ctx in ctxs[1:]) + + def are_compatible_values(self, *values: Any) -> bool: + return self.are_compatible_contexts(*self.infer_contexts(values)) + + def are_compatible_ops(self, *ops: BackendOps) -> bool: + if not ops: + return True + first = ops[0] + return all(op == first for op in ops) + + def enforce_convert_policy( + self, x: Any, to: Context | BackendFamily | str | None = None + ) -> Tuple[Any, Context]: + """Resolve the target context for ``x``.""" + self.infer_context(x) + ctx = self.normalize_context(to) + return x, ctx + + def resolve_context_priority( + self, + priority_ctx: Context | BackendFamily | str | None = None, + *other_ctx: object, + ) -> Context: + """Resolve explicit context first, then compatible inferred contexts.""" + if priority_ctx is not None: + return self.normalize_context(priority_ctx) + + inferred = self.infer_contexts(other_ctx) + if not inferred: + return self.default_ctx + + if not self.are_compatible_contexts(*inferred): + fams = tuple(ctx.ops.family for ctx in inferred) + raise ValueError(f"Incompatible inferred contexts: {fams!r}") + + first = inferred[0] + ops = type(first.ops)() + dtype = self._join_dtypes(ops, *(ctx.dtype for ctx in inferred)) + + return self.ctx_from_ops(ops=ops, dtype=dtype) + + def _join_dtypes(self, ops: BackendOps, *dtypes: DType | None) -> DType | None: + clean = [ops.sanitize_dtype(dt) for dt in dtypes if dt is not None] + if not clean: + return ops.sanitize_dtype(None) + + # Promote through the operands' OWN backend namespace. NumPy's + # ``result_type`` cannot interpret a torch/jax dtype, so joining the + # inferred contexts of a non-NumPy operator (for example a + # ``BlockDiagonalLinOp`` built with ``from_operators``, which infers a + # ``TreeSpace`` and joins its leaf dtypes) would otherwise raise + # ``TypeError: Cannot interpret 'torch.float64' as a data type``. + joined = ops.xp.result_type(*clean) + return ops.sanitize_dtype(joined) diff --git a/spacecore/contextual/_state.py b/spacecore/contextual/_state.py new file mode 100644 index 0000000..27b9ac9 --- /dev/null +++ b/spacecore/contextual/_state.py @@ -0,0 +1,263 @@ +from __future__ import annotations + +from contextlib import contextmanager +from typing import Any + +from ._context import Context +from ._contextual import Contextual +from .._check_policy import CheckLevel +from ..backend import BackendFamily, BackendOps +from ..backend import ops_registry as _ops_registry + + +_contextual: Contextual | None = None + + +def _state() -> Contextual: + """Return the process-wide contextual singleton.""" + global _contextual + if _contextual is None: + _contextual = Contextual() + return _contextual + + +def set_context( + ctx: Context | BackendFamily | str | None = None, + dtype: Any = None, +) -> None: + """ + Set the process-wide *baseline* SpaceCore context. + + Visible to every thread and async task. A :func:`use_context` override already + in progress keeps winning until its block exits. + + Parameters + ---------- + ctx : Context, BackendFamily, str, or None, optional + Context or backend specification. + dtype : Any, optional + Default dtype override. + """ + state = _state() + state.default_ctx = state.normalize_context(ctx, dtype=dtype) + + +def get_check_level() -> CheckLevel: + """ + Return the ambient default validation level applied to new bound objects. + + Returns + ------- + CheckLevel + The level that seeds a context-bound object constructed without an + explicit ``check_level``: the scoped override installed for this thread + or task if there is one, otherwise the process-wide baseline. + """ + return _state().get_check_level() + + +def set_check_level(level: CheckLevel | bool | None) -> None: + """ + Set the ambient default validation level applied to new bound objects. + + ``check_level`` is a property of the context-bound object (space, operator, + functional), not of the :class:`Context`. This sets the process-wide default + that seeds a new object when it is constructed without an explicit level. + + Parameters + ---------- + level : {"none", "cheap", "standard", "strict"} or None + New process-wide baseline. ``None`` restores the built-in default. + """ + _state().set_check_level(level) + + +@contextmanager +def use_check_level(level: CheckLevel | bool | None): + """ + Temporarily override the ambient default validation level within a ``with`` block. + + Scoped to the current thread / async task; see :func:`use_context` for the + scoping rules and the thread-pool caveat. + + Parameters + ---------- + level : {"none", "cheap", "standard", "strict"} or None + Level to install for the duration of the block. ``None`` leaves the + ambient default in place. + + Yields + ------ + CheckLevel + The resolved level in effect inside the block. + """ + with _state().scoped_check_level(level) as resolved: + yield resolved + + +@contextmanager +def use_context( + ctx: Context | BackendFamily | str | None = None, + dtype: Any = None, +): + """ + Temporarily override the default context within a ``with`` block. + + Yields the resolved :class:`Context` and unwinds on exit, so the override is + exception-safe and never leaks — the dependency-injection-friendly counterpart + to :func:`set_context`. + + Unlike :func:`set_context`, which moves the process-wide baseline seen by every + thread, this override is stored in a :class:`~contextvars.ContextVar` and is + visible only to the **current thread and async task**. Each ``asyncio`` task + receives its own copy of the ambient context at creation, so interleaved + coroutines cannot clobber one another's override. + + Parameters + ---------- + ctx : Context, BackendFamily, str, or None, optional + Context to install for the duration of the block, or a backend family + name (for example ``"numpy"`` or ``"jax"``) to resolve into one. ``None`` + keeps the active context and only applies ``dtype``. + dtype : dtype-like, optional + Representation dtype for the scoped context. ``None`` keeps the dtype of + the context being overridden. + + Yields + ------ + Context + The resolved context in effect inside the block. + + Notes + ----- + ``ContextVar`` values are deliberately *not* inherited by threads started with + :class:`threading.Thread`, so a scoped override does not silently leak into + workers you spawn:: + + with use_context("jax"): + list(executor.map(solve, problems)) # workers see the BASELINE + + To propagate it on purpose, snapshot the context and run the worker inside it:: + + from contextvars import copy_context + + with use_context("jax"): + snapshot = copy_context() + list(executor.map(lambda p: snapshot.run(solve, p), problems)) + + ``asyncio.create_task`` needs no such handling. + """ + with _state().scoped_ctx(ctx, dtype=dtype) as resolved: + yield resolved + + +def get_context() -> Context: + """ + Return the currently active default SpaceCore context. + + Resolves the :func:`use_context` override for this thread / task if one is + installed, otherwise the process-wide baseline set by :func:`set_context`. + + Returns + ------- + Context + Active default context. + """ + return _state().default_ctx + + +def resolve_context_priority( + priority_ctx: Context | BackendFamily | str | None = None, + *other_ctx: object, +) -> Context: + """ + Resolve the context assigned to a newly created object. + + Parameters + ---------- + priority_ctx : Context, BackendFamily, str, or None, optional + Explicit context that takes precedence when provided. + *other_ctx : object + Objects or contexts used as fallback context sources. + + Returns + ------- + Context + Resolved context. + """ + return _state().resolve_context_priority(priority_ctx, *other_ctx) + + +def register_ops(ops: type[BackendOps]) -> type[BackendOps]: + """ + Register a backend operations implementation. + + Writes to the process-wide backend registry + (:data:`spacecore.backend.ops_registry`), not to the ambient context state: + a registration is shared by every thread and is never scoped. + + Parameters + ---------- + ops : type of BackendOps + Backend operations class to register. + + Returns + ------- + type of BackendOps + Registered backend operations class. + + Raises + ------ + ContextConflictError + If the backend family is already registered. + """ + return _ops_registry.register(ops) + + +def normalize_context( + ctx: Context | BackendFamily | str | None = None, + dtype: Any = None, +) -> Context: + """ + Normalize a context specification through the process-wide state. + + Parameters + ---------- + ctx : Context, BackendFamily, str, or None, optional + Context or backend specification. + dtype : Any, optional + Default dtype override. + + Returns + ------- + Context + Normalized context. + """ + return _state().normalize_context(ctx, dtype=dtype) + + +def normalize_ops(ops: str | BackendFamily | BackendOps | type[BackendOps] | Context) -> BackendOps: + """ + Normalize backend operations through the process-wide state. + + Parameters + ---------- + ops : str, BackendFamily, BackendOps, type of BackendOps, or Context + Backend operations specification. + + Returns + ------- + BackendOps + Normalized backend operations singleton. + """ + if isinstance(ops, BackendOps): + return ops + return _ops_registry.get(ops) + + +def enforce_convert_policy( + x: Any, + to: Context | BackendFamily | str | None = None, +) -> tuple[Any, Context]: + """Resolve a conversion target context.""" + return _state().enforce_convert_policy(x, to) diff --git a/spacecore/functional/__init__.py b/spacecore/functional/__init__.py index dc31788..8dc8493 100644 --- a/spacecore/functional/__init__.py +++ b/spacecore/functional/__init__.py @@ -9,10 +9,14 @@ from ._base import Functional from ._algebra import ( + ConstantFunctional, + ProductFunctional, ScaledFunctional, ShiftedFunctional, SumFunctional, ZeroFunctional, + make_constant_functional, + make_functional_product, make_functional_sum, make_scaled_functional, make_shifted_functional, @@ -20,14 +24,17 @@ from ._composed import ComposedFunctional, make_functional_composed from ._linear import InnerProductFunctional, LinearFunctional, MatrixFreeLinearFunctional from ._quadratic import LinOpQuadraticForm, QuadraticForm +from ._realified import RealifiedFunctional, realify from .tools import ( + SpectralFunctional, + eigenvalue_space, + spectralize, HuberFunctional, KLDivergenceFunctional, L1NormFunctional, LpNormFunctional, NegativeEntropyFunctional, NuclearNormFunctional, - SpectralLpNormFunctional, SquaredL2NormFunctional, generalized_shrinkage, least_squares, @@ -38,6 +45,7 @@ __all__ = [ "ComposedFunctional", + "ConstantFunctional", "Functional", "HuberFunctional", "InnerProductFunctional", @@ -49,20 +57,27 @@ "MatrixFreeLinearFunctional", "NegativeEntropyFunctional", "NuclearNormFunctional", + "ProductFunctional", "QuadraticForm", + "RealifiedFunctional", "ScaledFunctional", "ShiftedFunctional", - "SpectralLpNormFunctional", + "SpectralFunctional", "SquaredL2NormFunctional", "SumFunctional", "ZeroFunctional", "generalized_shrinkage", "least_squares", + "make_constant_functional", "make_functional_composed", + "make_functional_product", "make_functional_sum", "make_scaled_functional", "make_shifted_functional", "project_nonneg", "prox_l1", "prox_l2sq", + "eigenvalue_space", + "realify", + "spectralize", ] diff --git a/spacecore/functional/_algebra.py b/spacecore/functional/_algebra.py index 02b425e..5619601 100644 --- a/spacecore/functional/_algebra.py +++ b/spacecore/functional/_algebra.py @@ -1,10 +1,19 @@ -"""Lazy functional algebra: scalar multiples and sums (mirrors ``linop/_algebra.py``). +"""Lazy functional algebra: scalar multiples, sums, and pointwise products. -`Functional` gains the additive/scalar algebra `LinOp` already has, so objectives -compose as ``a * F``, ``F + G``, ``F - G``, ``-F``. The operator overloads live on +Mirrors ``linop/_algebra.py`` for the parts they share. `Functional` gains the +additive/scalar algebra `LinOp` already has, so objectives compose as ``a * F``, +``F + G``, ``F - G``, ``-F``. The operator overloads live on :class:`~spacecore.functional.Functional` and delegate to the ``make_*`` factories here, which do local canonicalization (fold nested scalars, flatten nested sums). +Unlike the operator algebra this one is also **multiplicative**: functionals are +scalar-valued, so ``F * G`` is the pointwise product +(:class:`ProductFunctional`), with :class:`ConstantFunctional` as the embedding +of a plain scalar. Multiplying by a constant functional folds back to +:class:`ScaledFunctional` — scaling is the constant-factor case of a product. +Canonicalization stays *structural*: it reads node types, never values, so a +functional that merely happens to be constant is not recognized as one. + A functional's ``grad`` is a *metric (Riesz) gradient* -- an element of the domain ``X`` -- so the algebra combines child gradients through the domain's own vector ops (``X.add`` / ``X.scale``), never raw ``+`` / ``*`` (which would be wrong on a @@ -12,51 +21,43 @@ """ from __future__ import annotations -from numbers import Number from typing import Any from ._base import Functional from .._checks import checked_method -from ..backend import Context, jax_pytree_class - - -def is_scalar_like(value: Any) -> bool: - """Return whether ``value`` can be used as a scalar multiplier for a functional.""" - if isinstance(value, Number): - return True - shape = getattr(value, "shape", None) - if shape is not None: - return tuple(shape) == () - ndim = getattr(value, "ndim", None) - return ndim == 0 - - -def _scalar_eq(a: Any, b: Any) -> bool: - """Return whether two scalar-likes are equal, NaN-reflexive, as a real ``bool``.""" - if bool(a == b): - return True - try: - return bool(a != a) and bool(b != b) - except Exception: - return False - - -def _require_same_domain(terms: Any) -> None: +from .._check_policy import CheckLevel, minimum_check_level +from ..contextual import Context +from .._lazy_algebra import ( + finalize_sum, + flatten_sum, + fold_scaled, + is_scalar_like as is_scalar_like, # re-exported for functional/_base.py + scalar_eq, +) + + +def _require_same_domain(terms: Any, node: str = "SumFunctional") -> None: """Raise unless every functional in ``terms`` shares the first term's domain. Domain equality folds in the backend/dtype context, so this also rejects a same-shape space on a different backend or dtype. + + Parameters + ---------- + terms : sequence of Functional + Operands to compare. + node : str, optional + Node name used in the error message; sums and products share this check. """ domain = terms[0].domain for i, term in enumerate(terms[1:], start=1): if term.domain != domain: raise ValueError( - "All SumFunctional operands must have the same domain; operand 0 has " + f"All {node} operands must have the same domain; operand 0 has " f"domain {domain!r}, operand {i} has domain {term.domain!r}." ) -@jax_pytree_class class ScaledFunctional(Functional): """ Lazy scalar multiple ``scalar * functional``. @@ -67,18 +68,31 @@ class ScaledFunctional(Functional): Scalar coefficient. functional : Functional Functional to scale. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this functional. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Validation + policy is a property of the functional, not of the ``Context``. """ - def __init__(self, scalar: Any, functional: Functional) -> None: + def __init__( + self, + scalar: Any, + functional: Functional, + check_level: CheckLevel | bool | None = None, + ) -> None: if not isinstance(functional, Functional): raise TypeError(f"functional must be a Functional, got {type(functional).__name__}.") if not is_scalar_like(scalar): raise TypeError(f"scalar must be scalar-like, got {type(scalar).__name__}.") - super().__init__(functional.domain, functional.ctx) + # Default the policy from the OPERAND, not from its domain space, so the + # result inherits the least-strict operand independently of order. + if check_level is None: + check_level = functional.check_level + super().__init__(functional.domain, functional.ctx, check_level=check_level) self.scalar = scalar self.functional = functional.convert(self.ctx) - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any, *args: Any, **kwargs: Any) -> Any: """Return ``scalar * functional.value(x)``.""" return self._value_core(x, *args, **kwargs) @@ -106,9 +120,9 @@ def value_and_grad(self, x: Any, *args: Any, **kwargs: Any) -> tuple[Any, Any]: def __eq__(self, other: Any) -> bool: """Return whether another scaled functional has the same scalar and operand.""" - if not self._eq_backend_compatible(other): + if not self.same_math(other): return NotImplemented - return _scalar_eq(self.scalar, other.scalar) and self.functional == other.functional + return scalar_eq(self.scalar, other.scalar) and self.functional == other.functional def tree_flatten(self): """Flatten this functional for pytree registration (scalar is a traced child).""" @@ -149,29 +163,34 @@ def make_scaled_functional(scalar: Any, functional: Functional) -> Functional: raise TypeError(f"functional must be a Functional, got {type(functional).__name__}.") if not is_scalar_like(scalar): raise TypeError(f"scalar must be scalar-like, got {type(scalar).__name__}.") - if _scalar_eq(scalar, 0): - return ZeroFunctional(functional.domain, functional.ctx) - if _scalar_eq(scalar, 1): - return functional - if isinstance(functional, ZeroFunctional): - return functional - if isinstance(functional, ScaledFunctional): - return make_scaled_functional(scalar * functional.scalar, functional.functional) - return ScaledFunctional(scalar, functional) + + return fold_scaled( + scalar, + functional, + is_zero=lambda f: isinstance(f, ZeroFunctional), + unwrap_scaled=lambda f: ( + (f.scalar, f.functional) if isinstance(f, ScaledFunctional) else None + ), + make_zero=lambda: ZeroFunctional(functional.domain, functional.ctx), + make_scaled_node=ScaledFunctional, + ) -@jax_pytree_class class SumFunctional(Functional): """ - Lazy sum ``F_1 + ... + F_n`` of functionals on a common domain. + Lazy sum of finitely many functionals on a common domain. Parameters ---------- terms : sequence of Functional Nonempty sequence of functionals sharing one domain. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this functional. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Validation + policy is a property of the functional, not of the ``Context``. """ - def __init__(self, terms: Any) -> None: + def __init__(self, terms: Any, check_level: CheckLevel | bool | None = None) -> None: parts = tuple(terms) if not parts: raise ValueError( @@ -181,7 +200,9 @@ def __init__(self, terms: Any) -> None: if not isinstance(term, Functional): raise TypeError(f"operand {i} must be a Functional, got {type(term).__name__}.") _require_same_domain(parts) - super().__init__(parts[0].domain, parts[0].ctx) + if check_level is None: + check_level = minimum_check_level(tuple(term.check_level for term in parts)) + super().__init__(parts[0].domain, parts[0].ctx, check_level=check_level) self.terms = tuple(term.convert(self.ctx) for term in parts) @property @@ -189,7 +210,7 @@ def parts(self) -> tuple[Functional, ...]: """Return the summed terms in order.""" return self.terms - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any, *args: Any, **kwargs: Any) -> Any: """Return the sum of the term values at ``x``.""" return self._value_core(x, *args, **kwargs) @@ -224,7 +245,7 @@ def value_and_grad(self, x: Any, *args: Any, **kwargs: Any) -> tuple[Any, Any]: def __eq__(self, other: Any) -> bool: """Return whether another sum has the same ordered terms.""" - if not self._eq_backend_compatible(other): + if not self.same_math(other): return NotImplemented if len(self.terms) != len(other.terms): return False @@ -244,19 +265,6 @@ def _convert(self, new_ctx: Context) -> "SumFunctional": return SumFunctional(tuple(term.convert(new_ctx) for term in self.terms)) -def _flatten_functional_sum_terms(terms: Any) -> tuple[Functional, ...]: - """Flatten nested :class:`SumFunctional` nodes into a flat term tuple.""" - flat: list[Functional] = [] - for i, term in enumerate(terms): - if not isinstance(term, Functional): - raise TypeError(f"operand {i} must be a Functional, got {type(term).__name__}.") - if isinstance(term, SumFunctional): - flat.extend(term.terms) - else: - flat.append(term) - return tuple(flat) - - def make_functional_sum(terms: Any) -> Functional: """ Return a locally simplified lazy sum of functionals. @@ -280,19 +288,25 @@ def make_functional_sum(terms: Any) -> Functional: raise ValueError( "make_functional_sum requires a nonempty sequence of Functional operands." ) - flat = _flatten_functional_sum_terms(terms) + for i, term in enumerate(terms): + if not isinstance(term, Functional): + raise TypeError(f"operand {i} must be a Functional, got {type(term).__name__}.") + flat = flatten_sum( + terms, + is_sum=lambda t: isinstance(t, SumFunctional), + parts=lambda t: t.terms, + ) # Validate all terms' domains BEFORE dropping zeros, so a domain mismatch is # never swallowed by the single-survivor unwrap or the all-zero collapse. _require_same_domain(flat) - nonzero = tuple(term for term in flat if not isinstance(term, ZeroFunctional)) - if not nonzero: - return ZeroFunctional(flat[0].domain, flat[0].ctx) - if len(nonzero) == 1: - return nonzero[0] - return SumFunctional(nonzero) + return finalize_sum( + flat, + is_zero=lambda f: isinstance(f, ZeroFunctional), + make_zero=lambda: ZeroFunctional(flat[0].domain, flat[0].ctx), + make_sum_node=SumFunctional, + ) -@jax_pytree_class class ZeroFunctional(Functional): """ The zero functional: value ``0``, gradient the domain's zero element. @@ -306,12 +320,20 @@ class ZeroFunctional(Functional): Domain space. ctx : Context, str, or None, optional Backend context specification. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. """ - def __init__(self, dom: Any, ctx: Context | str | None = None) -> None: - super().__init__(dom, ctx) + def __init__( + self, + dom: Any, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: + super().__init__(dom, ctx, check_level=check_level) - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any, *args: Any, **kwargs: Any) -> Any: """Return the scalar zero.""" return self._value_core(x, *args, **kwargs) @@ -330,7 +352,7 @@ def value_and_grad(self, x: Any, *args: Any, **kwargs: Any) -> tuple[Any, Any]: def __eq__(self, other: Any) -> bool: """Return whether another zero functional has the same domain.""" - if not self._eq_backend_compatible(other): + if not self.same_math(other): return NotImplemented return self.domain == other.domain @@ -349,7 +371,121 @@ def _convert(self, new_ctx: Context) -> "ZeroFunctional": return ZeroFunctional(self.domain.convert(new_ctx), new_ctx) -@jax_pytree_class +class ConstantFunctional(Functional): + """ + The constant functional ``x -> constant``: fixed value, zero gradient. + + The embedding of a scalar into the functional algebra. It generalizes + :class:`ZeroFunctional` (which is the ``constant = 0`` case, kept separate + because it is the additive identity the canonicalizers recognize), and it is + what :func:`make_functional_product` folds against: multiplying by a constant + functional is exactly scaling by its value, so ``C * F`` collapses to a + :class:`ScaledFunctional` rather than building a product node. + + Parameters + ---------- + dom : Space + Domain space. + constant : scalar-like + Value returned at every point. + ctx : Context, str, or None, optional + Backend context specification. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this functional. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Validation + policy is a property of the functional, not of the ``Context``. + + Attributes + ---------- + constant : scalar-like + The stored value. + """ + + def __init__( + self, + dom: Any, + constant: Any, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: + if not is_scalar_like(constant): + raise TypeError(f"constant must be scalar-like, got {type(constant).__name__}.") + super().__init__(dom, ctx, check_level=check_level) + self.constant = constant + + @checked_method(in_space="domain", out_scalar=True) + def value(self, x: Any, *args: Any, **kwargs: Any) -> Any: + """Return the stored constant.""" + return self._value_core(x, *args, **kwargs) + + def _value_core(self, x: Any, *args: Any, **kwargs: Any) -> Any: + """Check-free constant in the domain dtype (mirrors ``ZeroFunctional``).""" + return self.ctx.asarray(self.constant) + + def grad(self, x: Any, *args: Any, **kwargs: Any) -> Any: + """Return the domain's zero element (a constant has zero derivative).""" + return self.domain.zeros() + + def value_and_grad(self, x: Any, *args: Any, **kwargs: Any) -> tuple[Any, Any]: + """Return ``(constant, X.zeros())``.""" + return self._value_core(x, *args, **kwargs), self.domain.zeros() + + def __eq__(self, other: Any) -> bool: + """Return whether another constant functional has the same domain and value.""" + if not self.same_math(other): + return NotImplemented + return self.domain == other.domain and scalar_eq(self.constant, other.constant) + + def tree_flatten(self): + """Flatten this functional for pytree registration (constant is a traced child).""" + return (self.constant,), (self.domain, self.ctx) + + @classmethod + def tree_unflatten(cls, aux, children): + """Rebuild this functional from pytree data.""" + dom, ctx = aux + (constant,) = children + return cls(dom, constant, ctx) + + def _convert(self, new_ctx: Context) -> "ConstantFunctional": + """Convert the constant functional to ``new_ctx``.""" + return ConstantFunctional(self.domain.convert(new_ctx), self.constant, new_ctx) + + +def make_constant_functional( + dom: Any, + constant: Any, + ctx: Context | str | None = None, +) -> Functional: + """ + Return a locally simplified constant functional on ``dom``. + + A zero constant collapses to :class:`ZeroFunctional`, so the additive + identity keeps exactly one representation and the sum/scale canonicalizers + continue to recognize it. + + Parameters + ---------- + dom : Space + Domain space. + constant : scalar-like + Value returned at every point. + ctx : Context, str, or None, optional + Backend context specification. + + Returns + ------- + Functional + :class:`ZeroFunctional` for a zero constant, else + :class:`ConstantFunctional`. + """ + if not is_scalar_like(constant): + raise TypeError(f"constant must be scalar-like, got {type(constant).__name__}.") + if scalar_eq(constant, 0): + return ZeroFunctional(dom, ctx) + return ConstantFunctional(dom, constant, ctx) + + class ShiftedFunctional(Functional): """ Affine shift ``functional + offset``: value shifted, gradient unchanged. @@ -360,18 +496,29 @@ class ShiftedFunctional(Functional): Functional to shift. offset : scalar-like Constant added to the value. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this functional. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Validation + policy is a property of the functional, not of the ``Context``. """ - def __init__(self, functional: Functional, offset: Any) -> None: + def __init__( + self, + functional: Functional, + offset: Any, + check_level: CheckLevel | bool | None = None, + ) -> None: if not isinstance(functional, Functional): raise TypeError(f"functional must be a Functional, got {type(functional).__name__}.") if not is_scalar_like(offset): raise TypeError(f"offset must be scalar-like, got {type(offset).__name__}.") - super().__init__(functional.domain, functional.ctx) + if check_level is None: + check_level = functional.check_level + super().__init__(functional.domain, functional.ctx, check_level=check_level) self.functional = functional.convert(self.ctx) self.offset = offset - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any, *args: Any, **kwargs: Any) -> Any: """Return ``functional.value(x) + offset``.""" return self._value_core(x, *args, **kwargs) @@ -391,9 +538,9 @@ def value_and_grad(self, x: Any, *args: Any, **kwargs: Any) -> tuple[Any, Any]: def __eq__(self, other: Any) -> bool: """Return whether another shifted functional has the same offset and operand.""" - if not self._eq_backend_compatible(other): + if not self.same_math(other): return NotImplemented - return _scalar_eq(self.offset, other.offset) and self.functional == other.functional + return scalar_eq(self.offset, other.offset) and self.functional == other.functional def tree_flatten(self): """Flatten this functional for pytree registration (offset is a traced child).""" @@ -433,8 +580,183 @@ def make_shifted_functional(functional: Functional, offset: Any) -> Functional: raise TypeError(f"functional must be a Functional, got {type(functional).__name__}.") if not is_scalar_like(offset): raise TypeError(f"offset must be scalar-like, got {type(offset).__name__}.") - if _scalar_eq(offset, 0): + if scalar_eq(offset, 0): return functional if isinstance(functional, ShiftedFunctional): return make_shifted_functional(functional.functional, functional.offset + offset) return ShiftedFunctional(functional, offset) + + +class ProductFunctional(Functional): + r""" + Lazy pointwise product ``(F * G)(x) = F(x) * G(x)`` on a shared domain. + + Deliberately **binary**: the gradient is the two-factor product rule, which + does not generalize to an n-ary node without the full Leibniz expansion, and + nesting ``(F*G)*H`` expresses the same thing with the same cost. + + The gradient conjugates each cofactor, mirroring + :meth:`ScaledFunctional.grad`. Writing :math:`D` for the derivative, + + .. math:: + + D(FG)(x)[h] = G(x)\, DF(x)[h] + F(x)\, DG(x)[h], + + and the Riesz gradient is the element pairing to that under the domain inner + product. Since the inner product conjugates its *first* argument, + :math:`\langle \overline{a} g, h\rangle = a \langle g, h\rangle`, so the + coefficients enter conjugated: + + .. math:: + + \nabla(FG)(x) = \overline{G(x)}\, \nabla F(x) + + \overline{F(x)}\, \nabla G(x). + + For real-valued factors — the usual case — the conjugations are identities. + The two terms are combined through the domain's own ``scale``/``add``, never + raw ``*``/``+``, because a domain element may be a pytree. + + Parameters + ---------- + left, right : Functional + Factors sharing one domain. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this functional. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Validation + policy is a property of the functional, not of the ``Context``. + + Attributes + ---------- + left, right : Functional + The two factors, converted into this node's context. + """ + + def __init__( + self, + left: Functional, + right: Functional, + check_level: CheckLevel | bool | None = None, + ) -> None: + for name, factor in (("left", left), ("right", right)): + if not isinstance(factor, Functional): + raise TypeError( + f"{name} must be a Functional, got {type(factor).__name__}." + ) + _require_same_domain((left, right), node="ProductFunctional") + # Default the policy from the OPERANDS, not from their domain space, so + # the result inherits the least-strict operand independently of order + # (mirrors SumFunctional). + if check_level is None: + check_level = minimum_check_level((left.check_level, right.check_level)) + super().__init__(left.domain, left.ctx, check_level=check_level) + self.left = left.convert(self.ctx) + self.right = right.convert(self.ctx) + + @property + def factors(self) -> tuple[Functional, Functional]: + """Return the two factors in order.""" + return (self.left, self.right) + + @checked_method(in_space="domain", out_scalar=True) + def value(self, x: Any, *args: Any, **kwargs: Any) -> Any: + """Return ``left.value(x) * right.value(x)``.""" + return self._value_core(x, *args, **kwargs) + + def _value_core(self, x: Any, *args: Any, **kwargs: Any) -> Any: + """Check-free pointwise product of the factor values.""" + return ( + self.left._value_core(x, *args, **kwargs) + * self.right._value_core(x, *args, **kwargs) + ) + + def _product_grad(self, lv: Any, lg: Any, rv: Any, rg: Any) -> Any: + """Combine factor values/gradients by the product rule, in domain ops.""" + domain = self.domain + conj = self.ops.conj + return domain.add(domain.scale(conj(rv), lg), domain.scale(conj(lv), rg)) + + def grad(self, x: Any, *args: Any, **kwargs: Any) -> Any: + """Return the product-rule Riesz gradient at ``x``. + + Both factors are evaluated through ``value_and_grad`` because the rule + needs each factor's *value* as well as its gradient; asking for the + gradients alone would evaluate the values a second time. + """ + lv, lg = self.left.value_and_grad(x, *args, **kwargs) + rv, rg = self.right.value_and_grad(x, *args, **kwargs) + return self._product_grad(lv, lg, rv, rg) + + def value_and_grad(self, x: Any, *args: Any, **kwargs: Any) -> tuple[Any, Any]: + """Return ``(value, grad)`` from one fused evaluation per factor.""" + lv, lg = self.left.value_and_grad(x, *args, **kwargs) + rv, rg = self.right.value_and_grad(x, *args, **kwargs) + return lv * rv, self._product_grad(lv, lg, rv, rg) + + def __eq__(self, other: Any) -> bool: + """Return whether another product has the same ordered factors. + + Ordered, like :class:`SumFunctional`: the operation commutes, but this is + structural equality of expression trees, not semantic equivalence. + """ + if not self.same_math(other): + return NotImplemented + return self.left == other.left and self.right == other.right + + def tree_flatten(self): + """Flatten this functional for pytree registration.""" + return (self.left, self.right), () + + @classmethod + def tree_unflatten(cls, aux, children): + """Rebuild this functional from pytree data.""" + left, right = children + return cls(left, right) + + def _convert(self, new_ctx: Context) -> "ProductFunctional": + """Convert both factors to ``new_ctx``.""" + return ProductFunctional( + self.left.convert(new_ctx), self.right.convert(new_ctx) + ) + + +def make_functional_product(left: Functional, right: Functional) -> Functional: + """ + Return a locally simplified pointwise product of two functionals. + + Only *structural* simplifications are attempted — the ones visible from the + node types, with no evaluation: + + * a :class:`ZeroFunctional` factor collapses the product to zero (matching + how :func:`make_scaled_functional` treats a zero scalar); + * a :class:`ConstantFunctional` factor becomes a + :class:`ScaledFunctional` on the other factor, since multiplying by a + constant *is* scaling. + + There is deliberately no attempt to recognize a functional that merely + *happens* to be constant or zero: that is a fact about values, not about the + expression, and is not decidable here. + + Parameters + ---------- + left, right : Functional + Factors sharing one domain. + + Returns + ------- + Functional + Simplified product. + """ + for name, factor in (("left", left), ("right", right)): + if not isinstance(factor, Functional): + raise TypeError(f"{name} must be a Functional, got {type(factor).__name__}.") + # Validate domains BEFORE any collapse, so a mismatch is never swallowed by + # the zero/constant shortcuts. + _require_same_domain((left, right), node="ProductFunctional") + + if isinstance(left, ZeroFunctional) or isinstance(right, ZeroFunctional): + return ZeroFunctional(left.domain, left.ctx) + if isinstance(left, ConstantFunctional): + return make_scaled_functional(left.constant, right) + if isinstance(right, ConstantFunctional): + return make_scaled_functional(right.constant, left) + return ProductFunctional(left, right) diff --git a/spacecore/functional/_base.py b/spacecore/functional/_base.py index 420226c..83995d0 100644 --- a/spacecore/functional/_base.py +++ b/spacecore/functional/_base.py @@ -2,7 +2,7 @@ from abc import abstractmethod from numbers import Number -from typing import TYPE_CHECKING, Any, Generic, Self, TypeVar +from typing import TYPE_CHECKING, Any, Generic, TypeVar from .._batching import _leading_batch_size, _warn_vmap_fallback_once @@ -14,9 +14,11 @@ _check_scalar_shape, ) from .._checks import checked_method +from ..backend import PyTreeNode from .._repr import describe_space, field_symbol -from .._contextual import ContextBound -from ..backend import Context +from ..contextual import ContextBound +from ..contextual import Context +from .._check_policy import CheckLevel from ..space import CoordinateSpace if TYPE_CHECKING: @@ -26,14 +28,14 @@ Domain = TypeVar("Domain", bound=CoordinateSpace) -class Functional(ContextBound, Generic[Domain]): +class Functional(PyTreeNode, ContextBound, Generic[Domain]): r""" Scalar-valued map on a space. ``Functional`` represents a map ``F : X -> K`` without assuming any storage model. It mirrors the minimal ``LinOp`` contract: the domain is converted - into the resolved context, value checks follow ``ctx.check_level``, and - batched evaluation is implemented by a backend ``vmap`` fallback. + into the resolved context, value checks follow this object's ``check_level``, + and batched evaluation is implemented by a backend ``vmap`` fallback. Parameters ---------- @@ -41,6 +43,10 @@ class Functional(ContextBound, Generic[Domain]): Domain space ``X``. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom``. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this functional. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Validation + policy is a property of the functional, not of the ``Context``. Attributes ---------- @@ -50,8 +56,13 @@ class Functional(ContextBound, Generic[Domain]): Resolved backend context. """ - def __init__(self, dom: Domain, ctx: Context | str | None = None) -> None: - (self.dom,) = self._bind_context(ctx, dom) + def __init__( + self, + dom: Domain, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: + (self.dom,) = self._bind_context(ctx, dom, check_level=check_level) @property def domain(self) -> Domain: @@ -184,23 +195,44 @@ def __neg__(self) -> "Functional": return make_scaled_functional(-1, self) - def __mul__(self, scalar: Any) -> "Functional": - """Return the lazy right scalar multiple ``self * scalar``.""" - from ._algebra import is_scalar_like, make_scaled_functional + def __mul__(self, other: Any) -> Any: + """Return ``self * other``. - if not is_scalar_like(scalar): - return NotImplemented - return make_scaled_functional(scalar, self) + Three cases, by operand type: a scalar gives a ``ScaledFunctional``, + another ``Functional`` the pointwise product, and a ``LinOp`` the + functional-weighted map ``x -> F(x) A x``. The last is **not** a + ``LinOp`` — it is non-linear — so it returns an + :class:`~spacecore.opfamily.OperatorFamily`; see that module. + """ + from ..linop import LinOp + from ._algebra import is_scalar_like, make_functional_product, make_scaled_functional - def __rmul__(self, scalar: Any) -> "Functional": - """Return the lazy left scalar multiple ``scalar * self``.""" - from ._algebra import is_scalar_like, make_scaled_functional + if isinstance(other, Functional): + return make_functional_product(self, other) + if isinstance(other, LinOp): + from ..opfamily import make_functional_scaled_operator - if not is_scalar_like(scalar): - return NotImplemented - return make_scaled_functional(scalar, self) + return make_functional_scaled_operator(self, other) + if is_scalar_like(other): + return make_scaled_functional(other, self) + return NotImplemented + + def __rmul__(self, other: Any) -> Any: + """Return ``other * self`` — see :meth:`__mul__` for the operand cases.""" + from ..linop import LinOp + from ._algebra import is_scalar_like, make_functional_product, make_scaled_functional + + if isinstance(other, Functional): + return make_functional_product(other, self) + if isinstance(other, LinOp): + from ..opfamily import make_functional_scaled_operator - @checked_method(in_space="domain", in_batched=True) + return make_functional_scaled_operator(self, other) + if is_scalar_like(other): + return make_scaled_functional(other, self) + return NotImplemented + + @checked_method(in_space="domain", in_batched=True, out_batched_scalar=True) def vvalue(self, xs: Any) -> Any: """Evaluate over a leading batch axis. Input must have shape ``(N,) + domain.shape``; use ``moveaxis`` for other layouts.""" _warn_vmap_fallback_once(self, "vvalue", _leading_batch_size(self.domain, xs)) @@ -234,13 +266,8 @@ def _short_repr(self) -> str: """Return a bounded ``ClassName(domain → field)`` form for nesting.""" return f"{type(self).__name__}({self._arrow()})" - @abstractmethod - def tree_flatten(self) -> tuple[tuple[Any, ...], Any]: - """Flatten this functional for pytree registration.""" - ... - - @classmethod - @abstractmethod - def tree_unflatten(cls, aux: Any, children: Any) -> Self: - """Rebuild this functional from pytree data.""" - ... + # tree_flatten / tree_unflatten are inherited from PyTreeNode, which owns the + # flatten contract for every SpaceCore container and auto-registers concrete + # subclasses with each backend's tree protocol. Every concrete Functional is a + # container, so the capability belongs on this base rather than being + # re-declared (and separately registered) per functional. diff --git a/spacecore/functional/_composed.py b/spacecore/functional/_composed.py index 6bc2983..3fb301c 100644 --- a/spacecore/functional/_composed.py +++ b/spacecore/functional/_composed.py @@ -6,7 +6,8 @@ from ._linear import InnerProductFunctional from ._quadratic import LinOpQuadraticForm from .._checks import checked_method -from ..backend import Context, jax_pytree_class +from .._check_policy import CheckLevel, minimum_check_level +from ..contextual import Context from ..kernels import core_kernels from ..linop import LinOp @@ -52,7 +53,6 @@ def make_functional_composed(F: Functional, A: LinOp) -> Functional: @core_kernels("composed-functional") -@jax_pytree_class class ComposedFunctional(Functional): """ Generic pull-back of a functional through a linear operator. @@ -65,15 +65,26 @@ class ComposedFunctional(Functional): Functional defined on ``A.codomain``. A : LinOp Linear operator whose codomain is ``F.domain``. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this functional. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Validation + policy is a property of the functional, not of the ``Context``. """ - def __init__(self, F: Functional, A: LinOp) -> None: + def __init__( + self, + F: Functional, + A: LinOp, + check_level: CheckLevel | bool | None = None, + ) -> None: _require_composable(F, A) - super().__init__(A.domain, A.ctx) + if check_level is None: + check_level = minimum_check_level((F.check_level, A.check_level)) + super().__init__(A.domain, A.ctx, check_level=check_level) self.F = F.convert(A.ctx) self.A = A - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any) -> Any: """ Evaluate ``F(A x)``. @@ -90,9 +101,61 @@ def value(self, x: Any) -> Any: """ return self._value_core(x) + @checked_method(in_space="domain", out_space="domain") + def grad(self, x: Any) -> Any: + r""" + Return the Riesz gradient of ``F o A`` by the chain rule. + + For :math:`G = F \circ A`, differentiating gives + :math:`DG(x)[h] = DF(Ax)[Ah] = \langle \nabla F(Ax), Ah\rangle_Y`. + Moving ``A`` across the pairing with the adjoint's defining identity + :math:`\langle Au, v\rangle_Y = \langle u, A^{\#}v\rangle_X` (and + conjugate symmetry) turns that into + :math:`\langle A^{\#}\nabla F(Ax), h\rangle_X`, so + + .. math:: + + \nabla (F \circ A)(x) = A^{\#}\, \nabla F(A x). + + ``LinOp.rapply`` **is** :math:`A^{\#}`, the metric adjoint (ADR-009), so + the geometry of both spaces is already accounted for; applying a Riesz + map on top would count it twice. Correspondingly there is no explicit + conjugation here — the adjoint identity absorbs it. + + Raises :class:`NotImplementedError` when the inner ``F`` has no gradient. + + Parameters + ---------- + x: + Element of ``A.domain``. + + Returns + ------- + Any + Riesz gradient in ``A.domain``. + """ + return self._grad_core(x) + + def value_and_grad(self, x: Any, *args: Any, **kwargs: Any) -> tuple[Any, Any]: + """Return ``(F(Ax), A^#(grad F(Ax)))`` from a single application of ``A``. + + The default base implementation would call ``value`` and ``grad`` + separately and therefore apply ``A`` twice; here the image ``A x`` is + computed once and shared, which is the point of the fused path for a + composition. + """ + y = self.A.apply(x) + value, gradient = self.F.value_and_grad(y, *args, **kwargs) + return value, self.A.rapply(gradient) + + @checked_method(in_space="domain", out_space="domain", in_batched=True, out_batched=True) + def vgrad(self, xs: Any) -> Any: + """Evaluate the chain-rule gradient over a leading batch axis.""" + return self._vgrad_core(xs) + def __eq__(self, other: Any) -> bool: """Return whether another composed functional has the same operands.""" - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented return self.F == other.F and self.A == other.A diff --git a/spacecore/functional/_linear.py b/spacecore/functional/_linear.py index 5c3f61c..68ef2fb 100644 --- a/spacecore/functional/_linear.py +++ b/spacecore/functional/_linear.py @@ -4,9 +4,10 @@ from typing import Any, Callable from ._base import Domain, Functional -from .._batching import _check_scalar_shape, _leading_batch_size +from .._batching import _check_scalar_shape from .._checks import checked_method -from ..backend import Context, jax_pytree_class +from .._check_policy import CheckLevel +from ..contextual import Context from ..kernels import core_kernels from ..space import Space, TreeElement, TreeSpace @@ -42,6 +43,9 @@ class LinearFunctional(Functional[Domain]): Domain space. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom``. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. """ @property @@ -72,7 +76,6 @@ def vgrad(self, xs: Any) -> Any: @core_kernels("inner-product-functional") -@jax_pytree_class class InnerProductFunctional(LinearFunctional[Domain]): r""" Linear functional represented by a domain element. @@ -88,6 +91,9 @@ class InnerProductFunctional(LinearFunctional[Domain]): Domain space. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom``. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Attributes ---------- @@ -100,8 +106,9 @@ def __init__( c: Any, dom: Domain, ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: - super().__init__(dom, ctx) + super().__init__(dom, ctx, check_level=check_level) self._c = _convert_space_element(self.domain, c) if self._checks_at_least("standard"): self.domain._check_member(self._c) @@ -111,22 +118,19 @@ def representer(self) -> Any: """Stored domain element ``c`` defining ``ell_c(x) = ``.""" return self._c - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any) -> Any: """Return ``domain.inner(representer, x)``.""" return self._value_core(x) - @checked_method(in_space="domain", in_batched=True) + @checked_method(in_space="domain", in_batched=True, out_batched_scalar=True) def vvalue(self, xs: Any) -> Any: """Evaluate ``domain.inner(representer, xs[i])`` without a Python loop.""" - values = self._vvalue_core(xs) - if self._checks_at_least("standard"): - _check_scalar_shape(values, (_leading_batch_size(self.domain, xs),)) - return values + return self._vvalue_core(xs) def __eq__(self, other: Any) -> bool: """Return whether another inner-product functional has the same representer.""" - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented if self.domain != other.domain: # Tier 2: domain before allclose return False @@ -155,7 +159,6 @@ def _convert(self, new_ctx: Context) -> InnerProductFunctional: @core_kernels("matrixfree-linear-functional") -@jax_pytree_class class MatrixFreeLinearFunctional(LinearFunctional[Domain]): """ Linear functional defined by user-supplied evaluation callables. @@ -176,6 +179,9 @@ class MatrixFreeLinearFunctional(LinearFunctional[Domain]): vvalue : callable or None, optional Optional callable with signature ``vvalue(xs: Any) -> Any`` for batched evaluation. If omitted, backend ``vmap`` fallback is used. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Returns ------- @@ -190,6 +196,7 @@ def __init__( dom: Domain, ctx: Context | str | None = None, vvalue: Callable[[Any], Any] | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: """ Initialize a matrix-free linear functional. @@ -218,7 +225,7 @@ def __init__( raise TypeError(f"value must be callable, got {type(value).__name__}.") if vvalue is not None and not callable(vvalue): raise TypeError(f"vvalue must be callable, got {type(vvalue).__name__}.") - super().__init__(dom, ctx) + super().__init__(dom, ctx, check_level=check_level) self.value_fn = value self.vvalue_fn = vvalue @@ -238,7 +245,7 @@ def representer(self) -> Any: """ raise NotImplementedError(f"{type(self).__name__} does not store a Riesz representer.") - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any) -> Any: """ Evaluate the scalar functional. @@ -253,11 +260,14 @@ def value(self, x: Any) -> Any: Any Scalar-like backend value returned by ``value_fn``. """ - y = self._value_core(x) - if self._checks_at_least("standard"): - _check_scalar_shape(y, ()) - return y + return self._value_core(x) + # NOT ``out_batched_scalar=True``: the decorator derives the expected batch + # shape as ``(_leading_batch_size(domain, xs),)`` — one leading axis, the + # contract ``Functional.vvalue`` documents. The check below is deliberately + # more permissive, stripping the domain shape so a user-supplied + # ``vvalue_fn`` may take *several* leading axes. Narrowing that to the + # documented contract is a real behavior change and is left alone here. @checked_method(in_space="domain", in_batched=True) def vvalue(self, xs: Any) -> Any: """ @@ -286,7 +296,7 @@ def vvalue(self, xs: Any) -> Any: def __eq__(self, other: Any) -> bool: """Return whether another matrix-free functional uses the same callables.""" - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented if self.domain != other.domain: # Tier 2: domain return False diff --git a/spacecore/functional/_quadratic.py b/spacecore/functional/_quadratic.py index 0ce6646..7696c40 100644 --- a/spacecore/functional/_quadratic.py +++ b/spacecore/functional/_quadratic.py @@ -10,8 +10,9 @@ ) from ._linear import LinearFunctional from .._checks import checked_method -from .._contextual import resolve_context_priority -from ..backend import Context, jax_pytree_class +from .._check_policy import CheckLevel +from ..contextual import resolve_context_priority +from ..contextual import Context from ..kernels import core_kernels from ..linop import LinOp @@ -26,6 +27,9 @@ class QuadraticForm(Functional[Domain]): Domain space. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom``. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. """ def hess_apply(self, x: Any) -> Any: @@ -44,7 +48,6 @@ def vgrad(self, xs: Any) -> Any: @core_kernels("linop-quadratic-form") -@jax_pytree_class class LinOpQuadraticForm(QuadraticForm[Domain]): r""" Represent a quadratic form backed by a linear operator. @@ -72,6 +75,9 @@ class LinOpQuadraticForm(QuadraticForm[Domain]): ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``Q`` and ``linear``. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Attributes ---------- @@ -89,6 +95,7 @@ def __init__( linear: LinearFunctional[Domain] | None = None, a: Any = 0, ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: if not isinstance(Q, LinOp): raise TypeError(f"Q must be a LinOp, got {type(Q).__name__}.") @@ -107,7 +114,7 @@ def __init__( if linear.domain != Q.domain: raise ValueError("linear.domain must match Q.domain.") - super().__init__(Q.domain, resolved_ctx) + super().__init__(Q.domain, resolved_ctx, check_level=check_level) self.Q = Q self.linear = linear self.a = self.ctx.asarray(a) @@ -120,7 +127,7 @@ def _check_hermitian_structure(Q: LinOp[Domain, Domain]) -> None: if result is False: raise ValueError("LinOpQuadraticForm requires Q to be Hermitian/self-adjoint.") - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any) -> Any: """Return ``1/2 * + linear(x) + a``.""" return self._value_core(x) @@ -144,13 +151,10 @@ def hess_apply(self, x: Any) -> Any: """Return the Hessian action ``Q x`` under the Hermitian assumption.""" return self.Q.apply(x) - @checked_method(in_space="domain", in_batched=True) + @checked_method(in_space="domain", in_batched=True, out_batched_scalar=True) def vvalue(self, xs: Any) -> Any: """Evaluate the quadratic objective over a leading batch axis.""" - values = self._vvalue_core(xs) - if self._checks_at_least("standard"): - _check_scalar_shape(values, (_leading_batch_size(self.domain, xs),)) - return values + return self._vvalue_core(xs) @checked_method(in_space="domain", out_space="domain", in_batched=True, out_batched=True) def vgrad(self, xs: Any) -> Any: @@ -159,7 +163,7 @@ def vgrad(self, xs: Any) -> Any: def __eq__(self, other: Any) -> bool: """Return whether another quadratic form has the same stored terms.""" - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented if self.Q != other.Q: # Tier 2/3: operator (own gate) return False diff --git a/spacecore/functional/_realified.py b/spacecore/functional/_realified.py new file mode 100644 index 0000000..883f86c --- /dev/null +++ b/spacecore/functional/_realified.py @@ -0,0 +1,240 @@ +r"""Real-coordinate view of a complex-domain functional. + +Optimizers and line searches that assume a real vector space cannot consume a +functional whose domain is complex — ``spacecore.optimize`` flattens to real +coordinates, and a complex iterate has no ordering, no real inner product to +descend in, and no meaningful ``result_type`` against a real step size. The usual +workaround is to hand-write a real wrapper per objective, re-deriving the +Cauchy-Riemann bookkeeping each time and getting the conjugate wrong. + +:class:`RealifiedFunctional` does it once. A real-valued ``F`` on a complex +coordinate space has differential ``dF = Re`` with metric gradient ``g``, +so writing ``v = a + i b`` gives + +.. math:: + + \partial F/\partial a = \operatorname{Re} g_v, \qquad + \partial F/\partial b = \operatorname{Im} g_v, + +where ``g_v = flatten(X.riesz(g))`` is the *coordinate* gradient. The real +gradient is therefore exactly the stacked pair — no Wirtinger calculus at the +call site, and the metric correction stays where ADR-010 put it. + +References +---------- +.. [Realification] Treating ``C^m`` as ``R^{2m}`` by ``v = a + i b`` is the + standard realification of a complex vector space, and ``Re<., .>_C`` is a + real inner product on ``R^{2m}`` inducing the same norm, so the two views are + isometric. + + This is **derived**, from the axioms of Conway, *A Course in Functional + Analysis*, 2nd ed., Springer, 1990, **Definition I.1.1** — stated uniformly + over ``F = R`` or ``C``, with the note that ``conj(alpha) = alpha`` when + ``F = R``. Each property of a real inner product needs a different axiom: + + * *real-bilinearity* — (a) and (b) restricted to **real** scalars, where the + conjugation in (b) is the identity; ``Re`` of a real-linear map is + real-linear; + * *symmetry* — (d), since ``Re conj(z) = Re z``, so + ``Re = Re``; + * *positive-definiteness* — (c) for `` >= 0`` **and** the extra + inner-product axiom (f), `` = 0 => v = 0``. Without these two, + ``Re<., .>`` would only be a symmetric real bilinear form; + * *same norm* — (c) and (d) make ```` real and nonnegative, so + ``Re = = ||v||^2``. + + Convention caveat: Conway's I.1.1 is linear in the **first** argument and + conjugate-linear in the second; SpaceCore's ``Space.inner`` conjugates the + **first**. The realification is unaffected — ``Re<., .>`` is symmetric either + way — but do not transcribe the axioms slot-for-slot into this codebase. + +.. [Wirtinger] For a **real-valued** ``F`` the Wirtinger derivatives satisfy + ``dF/d(conj v) = conj(dF/dv)``, so the real gradient is + ``2 dF/d(conj v)`` — the same stacked pair derived above. See Kreutz-Delgado, + "The complex gradient operator and the CR-calculus", arXiv:0906.4835, 2009, + §4 (the real-valued case and the steepest-ascent direction). The conjugate is + where hand-written wrappers usually go wrong, which is the reason this view + exists once rather than per objective. +""" +from __future__ import annotations + +from typing import Any, Self, Tuple + +from .._checks import checked_method +from ..contextual import Context +from ..space import CoordinateSpace, DenseCoordinateSpace +from ._base import Functional + + +class RealifiedFunctional(Functional): + r""" + View of a complex-domain functional over stacked real coordinates. + + For a base functional :math:`F` on a coordinate space :math:`X` with complex + flattened coordinates :math:`v \in \mathbb{C}^m`, this functional acts on + :math:`w = (\operatorname{Re} v, \operatorname{Im} v) \in \mathbb{R}^{2m}`. + + Because a real-valued :math:`F` has differential + :math:`dF = \operatorname{Re}\langle g, dv \rangle` with metric gradient + :math:`g`, the real coordinate gradient is exactly + :math:`(\operatorname{Re} g_v, \operatorname{Im} g_v)` where + :math:`g_v = \operatorname{flatten}(X.\operatorname{riesz}(g))`. + + Parameters + ---------- + base : Functional + Functional on a complex :class:`~spacecore.space.CoordinateSpace`. + + Raises + ------ + TypeError + If the base domain is not a coordinate space (there are no flattened + coordinates to split). + ValueError + If the base domain is already real — use the functional directly rather + than paying for a no-op view. + + Notes + ----- + The domain is classified by ``Space.field``, which is derived from the dtype. + That is exact for a genuinely complex coordinate space, but **over-reports** + for a space whose complex storage is constrained — most importantly + :class:`~spacecore.HermitianSpace`, where Hermitian matrices form a *real* + vector space of dimension :math:`n^2` yet ``field`` reads ``"complex"``. + + Realifying such a space is *safe but redundant*: ``unflatten`` symmetrizes, + so ``from_real`` always lands back on the manifold and no step can escape it, + but the view carries :math:`2n^2` real coordinates for :math:`n^2` real + dimensions. ``flatten . unflatten`` is then a projection rather than the + identity, so the realified problem has a rank-deficient Hessian with + :math:`n^2` null directions — harmless for first-order methods, degenerate + for Newton-type ones. Prefer optimizing such a space directly. + + Examples + -------- + >>> import numpy as np + >>> import spacecore as sc + >>> ctx = sc.Context(sc.NumpyOps(), dtype=np.complex128) + >>> X = sc.DenseCoordinateSpace((2,), ctx=ctx) + >>> F = sc.SquaredL2NormFunctional(X) + >>> R = sc.realify(F) + >>> R.domain.shape + (4,) + >>> float(R.value(R.to_real(ctx.asarray([3.0 + 4.0j, 0.0])))) + 12.5 + """ + + def __init__(self, base: Functional) -> None: + X = base.domain + if not isinstance(X, CoordinateSpace): + raise TypeError( + "RealifiedFunctional requires a CoordinateSpace domain; " + f"got {type(X).__name__}." + ) + if X.field == "real": + raise ValueError( + "The base functional already has a real domain; use it directly." + ) + ops = X.ops + # check_level is a property of the bound object, not of the Context, so + # it is threaded through the constructors rather than the context. + real_ctx = Context(ops, dtype=ops.real_dtype(X.dtype)) + dom = DenseCoordinateSpace( + (2 * X.size,), real_ctx, check_level=X.check_level + ) + super().__init__(dom, real_ctx, check_level=X.check_level) + self.base = base + self.complex_space = X + + def to_real(self, y: Any) -> Any: + """Flatten a base-domain element into stacked real coordinates.""" + ops = self.ops + v = self.complex_space.flatten(y) + return ops.concatenate([ops.real(v), ops.imag(v)]) + + def from_real(self, w: Any) -> Any: + """Rebuild a base-domain element from stacked real coordinates.""" + m = self.complex_space.size + v = w[:m] + 1j * w[m:] + return self.complex_space.unflatten(self.complex_space.ctx.asarray(v)) + + # Output-only: the input is validated by ``base.value`` against the *complex* + # domain after ``from_real``, so an ``in_space`` here would check ``w`` + # against the wrong space. + @checked_method(out_scalar=True) + def value(self, w: Any, *args: Any, **kwargs: Any) -> Any: + """Return ``base.value`` at the complex element ``w`` encodes.""" + return self.base.value(self.from_real(w), *args, **kwargs) + + def grad(self, w: Any, *args: Any, **kwargs: Any) -> Any: + """Return the stacked real gradient ``(Re g_v, Im g_v)``.""" + return self.value_and_grad(w, *args, **kwargs)[1] + + def value_and_grad(self, w: Any, *args: Any, **kwargs: Any) -> Tuple[Any, Any]: + """Return ``(value, stacked real gradient)`` from one fused base evaluation. + + ``riesz`` maps the base's metric gradient back to coordinates before the + real/imaginary split; skipping it would silently return the wrong vector + on any non-Euclidean geometry. + """ + X, ops = self.complex_space, self.ops + val, g = self.base.value_and_grad(self.from_real(w), *args, **kwargs) + gv = X.flatten(X.riesz(g)) + return val, ops.concatenate([ops.real(gv), ops.imag(gv)]) + + def __eq__(self, other: Any) -> bool: + """Return whether another realified functional wraps an equal base.""" + if not self.same_math(other): + return NotImplemented + return self.base == other.base + + def _convert(self, new_ctx: Context) -> "RealifiedFunctional": + """Convert the wrapped functional to ``new_ctx`` and re-derive the view. + + ``new_ctx`` describes the *real view*, so the base must be converted to + the matching **complex** dtype — passing the real context straight + through would make the base's domain real and the view impossible to + rebuild. The pairing comes from ``ops.complex_dtype`` because it is not + portable to derive: NumPy promotes ``float32`` against a Python complex + to ``complex128``, JAX and Torch to ``complex64``. + """ + ops = new_ctx.ops + base_ctx = Context(ops, dtype=ops.complex_dtype(new_ctx.dtype)) + return RealifiedFunctional(self.base.convert(base_ctx)) + + def tree_flatten(self) -> tuple[tuple[Any, ...], Any]: + """Flatten this functional for pytree registration.""" + return (self.base,), () + + @classmethod + def tree_unflatten(cls, aux: Any, children: Any) -> Self: + """Rebuild this functional from pytree data.""" + (base,) = children + return cls(base) + + +def realify(F: Functional) -> Functional: + """ + Return ``F`` unchanged on a real domain, realified on a complex one. + + The idempotent entry point: calling it on an already-real functional is a + no-op rather than an error, so a caller preparing an objective for a + real-only optimizer need not branch on the field. + + Parameters + ---------- + F : Functional + Functional on a real or complex coordinate space. + + Returns + ------- + Functional + ``F`` itself when its domain is real, else a :class:`RealifiedFunctional` + over stacked real coordinates. + """ + if F.domain.field == "real": + return F + return RealifiedFunctional(F) + + +__all__ = ["RealifiedFunctional", "realify"] diff --git a/spacecore/functional/tools/__init__.py b/spacecore/functional/tools/__init__.py index ebf36e7..aeda2ac 100644 --- a/spacecore/functional/tools/__init__.py +++ b/spacecore/functional/tools/__init__.py @@ -10,7 +10,12 @@ from ._least_squares import least_squares from ._norms import L1NormFunctional, LpNormFunctional, SquaredL2NormFunctional from ._proximal import generalized_shrinkage, project_nonneg, prox_l1, prox_l2sq -from ._spectral import NuclearNormFunctional, SpectralLpNormFunctional +from ._spectral import ( + NuclearNormFunctional, + SpectralFunctional, + eigenvalue_space, + spectralize, +) __all__ = [ "HuberFunctional", @@ -19,11 +24,13 @@ "LpNormFunctional", "NegativeEntropyFunctional", "NuclearNormFunctional", - "SpectralLpNormFunctional", + "SpectralFunctional", "SquaredL2NormFunctional", "generalized_shrinkage", "least_squares", "project_nonneg", "prox_l1", "prox_l2sq", + "eigenvalue_space", + "spectralize", ] diff --git a/spacecore/functional/tools/_coordinate.py b/spacecore/functional/tools/_coordinate.py index d3ad000..82a20ca 100644 --- a/spacecore/functional/tools/_coordinate.py +++ b/spacecore/functional/tools/_coordinate.py @@ -7,6 +7,21 @@ ``domain.riesz_inverse``. Centralizing the correction is the single defense against the [ADR-019](019_everyday_toolbox.md) trap of pairing a Euclidean gradient with a non-Euclidean metric. + +The correction is the Riesz representation theorem: the differential ``DF(x)`` +is a bounded linear functional on ``X``, so it is represented by a *unique* +``g`` with ``DF(x)[h] = _X`` for all ``h`` — and that ``g``, not the array +of partials, is what ``grad`` returns. The two coincide only when ``X`` is +Euclidean. + +References +---------- +.. [Riesz] A. V. Balakrishnan, *Applied Functional Analysis*, 2nd ed., + Springer, 1981, §1.7 (Riesz representation theorem) — existence and + uniqueness of the representer in a Hilbert space. See also Brezis, + *Functional Analysis, Sobolev Spaces and PDE*, Springer, 2011, Thm. 5.5. +.. [ADR010] The library contract: ``grad`` is the metric (Riesz) gradient, an + element of ``X``; ``X.riesz`` maps it back to coordinates. """ from __future__ import annotations diff --git a/spacecore/functional/tools/_entropy.py b/spacecore/functional/tools/_entropy.py index 558af38..04dd6d5 100644 --- a/spacecore/functional/tools/_entropy.py +++ b/spacecore/functional/tools/_entropy.py @@ -1,16 +1,36 @@ -"""Entropy objectives: negative entropy and KL divergence (ADR-019).""" +"""Entropy objectives: negative entropy and KL divergence (ADR-019). + +References +---------- +.. [Beck] A. Beck, *First-Order Methods in Optimization*, MOS-SIAM, 2017. + **§4.4.10** (a worked section, p. 97, not a numbered example) gives the + conjugate of the negative entropy over the unit simplex — the log-sum-exp + function of §4.4.11; Example 5.27 gives its ``1``-strong convexity over the + simplex w.r.t. the ``l_1`` norm — the property that makes it the standard + mirror-descent kernel (Example 9.10, entropic mirror descent / the + multiplicative-weights update). +.. [Lifted] Composed with a Jordan spectrum via + :func:`~spacecore.spectralize`, the negative entropy becomes the **von Neumann + entropy** ``S(rho) = -tr(rho log rho)``; see Aubrun & Szarek, *Alice and Bob + Meet Banach*, AMS, 2017, **§1.3.3**, eq. (1.36) for the definition and + **Proposition 1.19(i)** for concavity of ``S`` — equivalently convexity of the + negative entropy — whose proof uses concavity of ``x -> -x log x`` together + with Klein's lemma. The majorization machinery it sits on is §1.3.1, + **Proposition 1.12** (``x < y`` iff ``x`` is a convex combination of + coordinate permutations of ``y`` iff ``y = Bx`` for a bistochastic ``B``). +""" from __future__ import annotations from typing import Any, cast from .._base import Domain from .._linear import _convert_space_element -from ...backend import Context, jax_pytree_class +from ...contextual import Context +from ..._check_policy import CheckLevel from ..._checks import checked_method from ._coordinate import _CoordinateFunctional -@jax_pytree_class class NegativeEntropyFunctional(_CoordinateFunctional[Domain]): r""" Negative (Shannon) entropy ``F(x) = sum_i x_i log x_i``. @@ -25,6 +45,9 @@ class NegativeEntropyFunctional(_CoordinateFunctional[Domain]): Domain space ``X``. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom``. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Examples -------- @@ -39,10 +62,15 @@ class NegativeEntropyFunctional(_CoordinateFunctional[Domain]): array([1., 1.]) """ - def __init__(self, dom: Domain, ctx: Context | str | None = None) -> None: - super().__init__(dom, ctx) + def __init__( + self, + dom: Domain, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: + super().__init__(dom, ctx, check_level=check_level) - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any) -> Any: """Return ``sum_i x_i log x_i`` with ``0 log 0 = 0``.""" o = self.ops @@ -69,7 +97,6 @@ def _convert(self, new_ctx: Context) -> "NegativeEntropyFunctional": return NegativeEntropyFunctional(self.domain.convert(new_ctx), new_ctx) -@jax_pytree_class class KLDivergenceFunctional(_CoordinateFunctional[Domain]): r""" Kullback--Leibler divergence to a fixed positive ``target``. @@ -91,6 +118,9 @@ class KLDivergenceFunctional(_CoordinateFunctional[Domain]): bare element does not carry its space. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom``. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Examples -------- @@ -110,8 +140,9 @@ def __init__( target: Any, dom: Domain, ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: - super().__init__(dom, ctx) + super().__init__(dom, ctx, check_level=check_level) self._target = _convert_space_element(self.domain, target) if self._checks_at_least("standard"): self.domain._check_member(self._target) @@ -121,7 +152,7 @@ def target(self) -> Any: """Stored reference element ``t``.""" return self._target - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any) -> Any: """Return ``sum_i x_i log(x_i / t_i)`` with ``0 log 0 = 0``.""" o = self.ops diff --git a/spacecore/functional/tools/_huber.py b/spacecore/functional/tools/_huber.py index 2025d0a..211e96b 100644 --- a/spacecore/functional/tools/_huber.py +++ b/spacecore/functional/tools/_huber.py @@ -5,12 +5,12 @@ from typing import Any from .._base import Domain -from ...backend import Context, jax_pytree_class +from ...contextual import Context +from ..._check_policy import CheckLevel from ..._checks import checked_method from ._coordinate import _CoordinateFunctional -@jax_pytree_class class HuberFunctional(_CoordinateFunctional[Domain]): r""" Separable Huber loss ``F(x) = sum_i h_delta(x_i)``. @@ -18,7 +18,15 @@ class HuberFunctional(_CoordinateFunctional[Domain]): The per-coordinate loss is quadratic near the origin and linear in the tails: ``h_delta(r) = 1/2 r^2`` for ``|r| <= delta`` and ``delta (|r| - delta/2)`` otherwise. It is everywhere differentiable, with - gradient ``r`` in the quadratic region and ``delta sign(r)`` in the tails. + gradient ``r`` in the quadratic region and ``delta sign(r)`` in the tails + (value and gradient agree at ``|r| = delta``, so ``h_delta`` is ``C^1``). + + **Convention.** This is Huber's original robust-statistics scaling [Huber1964]_, + which is ``delta`` times the Moreau envelope used in the optimization + literature: with ``M_mu`` the envelope of the absolute value from + [Beck]_ Example 6.54, ``h_delta = delta * M_delta``. The distinction matters + when transcribing a prox or a smoothing constant from either source — the two + differ by exactly one factor of ``delta``. Parameters ---------- @@ -28,6 +36,19 @@ class HuberFunctional(_CoordinateFunctional[Domain]): Transition threshold; must be finite and ``> 0``. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom``. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. + + References + ---------- + .. [Huber1964] P. J. Huber, "Robust estimation of a location parameter", + *Ann. Math. Statist.* 35(1):73-101, 1964, §2 — the original + quadratic-near-zero / linear-in-the-tails loss. + .. [Beck] A. Beck, *First-Order Methods in Optimization*, MOS-SIAM, 2017, + Example 6.54 (the Huber function as the Moreau envelope of the norm), + Example 6.62 (its smoothness: gradient ``1/mu``-Lipschitz) and + Example 6.66 (prox of the Huber function). Examples -------- @@ -40,14 +61,20 @@ class HuberFunctional(_CoordinateFunctional[Domain]): 2.625 """ - def __init__(self, dom: Domain, delta: Any, ctx: Context | str | None = None) -> None: - super().__init__(dom, ctx) + def __init__( + self, + dom: Domain, + delta: Any, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: + super().__init__(dom, ctx, check_level=check_level) delta = float(delta) if not math.isfinite(delta) or delta <= 0.0: raise ValueError(f"HuberFunctional requires a finite delta > 0, got {delta}.") self.delta = delta - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any) -> Any: """Return ``sum_i h_delta(x_i)``.""" o = self.ops diff --git a/spacecore/functional/tools/_norms.py b/spacecore/functional/tools/_norms.py index 57288bf..9fde0d0 100644 --- a/spacecore/functional/tools/_norms.py +++ b/spacecore/functional/tools/_norms.py @@ -5,12 +5,12 @@ from typing import Any, cast from .._base import Domain -from ...backend import Context, jax_pytree_class +from ...contextual import Context +from ..._check_policy import CheckLevel from ..._checks import checked_method from ._coordinate import _CoordinateFunctional, _inner_core, lp_coordinate_grad, lp_value -@jax_pytree_class class SquaredL2NormFunctional(_CoordinateFunctional[Domain]): r""" Half the squared space norm ``F(x) = 1/2 ||x||_X^2 = 1/2 _X``. @@ -27,6 +27,9 @@ class SquaredL2NormFunctional(_CoordinateFunctional[Domain]): Domain space ``X``. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom``. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Examples -------- @@ -41,10 +44,15 @@ class SquaredL2NormFunctional(_CoordinateFunctional[Domain]): array([3., 4.]) """ - def __init__(self, dom: Domain, ctx: Context | str | None = None) -> None: - super().__init__(dom, ctx) + def __init__( + self, + dom: Domain, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: + super().__init__(dom, ctx, check_level=check_level) - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any) -> Any: """Return ``1/2 _X`` as a real scalar.""" return 0.5 * self.ops.real(_inner_core(self.domain, x, x)) @@ -73,7 +81,6 @@ def _convert(self, new_ctx: Context) -> "SquaredL2NormFunctional": return SquaredL2NormFunctional(self.domain.convert(new_ctx), new_ctx) -@jax_pytree_class class LpNormFunctional(_CoordinateFunctional[Domain]): r""" Coordinate ``p``-norm ``F(x) = (sum_i |x_i|^p)^{1/p}`` for ``p >= 1``. @@ -86,6 +93,9 @@ class LpNormFunctional(_CoordinateFunctional[Domain]): Norm order; must be finite and ``>= 1``. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom``. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Notes ----- @@ -105,14 +115,20 @@ class LpNormFunctional(_CoordinateFunctional[Domain]): 6.0 """ - def __init__(self, dom: Domain, p: Any, ctx: Context | str | None = None) -> None: - super().__init__(dom, ctx) + def __init__( + self, + dom: Domain, + p: Any, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: + super().__init__(dom, ctx, check_level=check_level) p = float(p) if not math.isfinite(p) or p < 1.0: raise ValueError(f"LpNormFunctional requires a finite p >= 1, got {p}.") self.p = p - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any) -> Any: """Return ``(sum_i |x_i|^p)^{1/p}``.""" return lp_value(self.ops, x, self.p) @@ -137,7 +153,9 @@ def _convert(self, new_ctx: Context) -> "LpNormFunctional": def L1NormFunctional( - dom: Domain, ctx: Context | str | None = None + dom: Domain, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> "LpNormFunctional[Domain]": r""" Coordinate 1-norm ``||x||_1`` -- a thin wrapper for ``LpNormFunctional(X, 1)``. @@ -148,10 +166,13 @@ def L1NormFunctional( Domain space ``X``. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom``. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for the returned functional. When omitted, the + ambient default (see :func:`spacecore.get_check_level`) is used. Returns ------- LpNormFunctional The ``p = 1`` instance of :class:`LpNormFunctional`. """ - return LpNormFunctional(dom, 1.0, ctx) + return LpNormFunctional(dom, 1.0, ctx, check_level=check_level) diff --git a/spacecore/functional/tools/_proximal.py b/spacecore/functional/tools/_proximal.py index 3878ccf..e04f1eb 100644 --- a/spacecore/functional/tools/_proximal.py +++ b/spacecore/functional/tools/_proximal.py @@ -20,6 +20,26 @@ metrics; on a non-diagonal metric the subproblem does not separate, so the primitive **raises** rather than returning a wrong (separable) answer (ADR-019 / ADR-020 diagonal-metric rule). + +References +---------- +.. [Beck] A. Beck, *First-Order Methods in Optimization*, MOS-SIAM, 2017. + Definition 6.1 (the proximal mapping ``prox_f(v) = argmin_u f(u) + 1/2||u-v||^2``) + and Theorem 6.3 (existence and uniqueness for closed proper convex ``f``). + §6.2.4 (``g_2``) gives the scalar soft-threshold + ``T_lambda(x) = [|x| - lambda]_+ sgn(x)``, which is ``prox`` of + ``lambda|.|``; separability (Theorem 6.6, in §6.3) lifts it coordinatewise, + and Example 6.8 states the resulting ``l_1``-norm prox directly. The ``prox`` + of ``t/2 ||.||^2`` is the linear shrinkage ``v/(1+t)``: **§6.2.3** (Convex + Quadratic, p. 132) — a worked section that carries no numbered result, so it + is cited as a section. (Not "Example 6.9", which is the negative sum of + logs.) +.. [BC] H. H. Bauschke and P. L. Combettes, *Convex Analysis and Monotone + Operator Theory in Hilbert Spaces*, 2nd ed., Springer, 2017, Ch. 24 + (proximity operators): §24.2 basic properties, §24.4 functions on the real + line, §24.5 proximal thresholding — the Hilbert-space statement, which is the + setting SpaceCore actually works in (the threshold below is metric-aware for + exactly this reason). """ from __future__ import annotations @@ -144,13 +164,25 @@ def _require_step(t: float, fname: str) -> float: return float(t) -def prox_l1(v: Any, t: Any, X: Any) -> Any: +def prox_l1(v: Any, t: Any, X: Any, *, nonneg: bool = False) -> Any: r""" Proximal operator of ``t * ||.||_1`` in the space metric (soft-threshold). Returns ``argmin_x 1/2 ||x - v||^2_X + t ||x||_1``, i.e. the metric-aware soft-threshold with per-coordinate width ``t / w_i`` on a diagonal metric. + With ``nonneg=True`` the minimization is additionally constrained to + ``x >= 0``, giving ``argmin_{x >= 0} 1/2 ||x - v||^2_X + t ||x||_1``. On the + nonnegative orthant ``||x||_1 = sum_i x_i`` is *linear*, so the subproblem + ``1/2 w_i (x_i - v_i)^2 + t x_i`` has the one-sided solution + ``max(v_i - t / w_i, 0)`` — derived without the sign/absolute-value form of + the unconstrained soft-threshold. The two nevertheless agree numerically: + clipping the symmetric threshold at zero gives the same result for every + ``t >= 0`` (for ``v_i <= 0`` both are ``0``; for ``v_i > 0`` the symmetric + form already reduces to ``max(v_i - t / w_i, 0)``). The equality is pinned by + a test rather than relied on, since it does not hold for proximal operators + in general. + Parameters ---------- v : array-like @@ -159,6 +191,9 @@ def prox_l1(v: Any, t: Any, X: Any) -> Any: Nonnegative threshold / step size. X : Space Ambient inner-product space (Euclidean or diagonal metric). + nonneg : bool, optional + Constrain the solution to ``x >= 0`` (real spaces only). Default + ``False``. With ``t = 0`` this degenerates to :func:`project_nonneg`. Returns ------- @@ -168,10 +203,10 @@ def prox_l1(v: Any, t: Any, X: Any) -> Any: t = _require_step(t, "prox_l1") v = X.ctx.asarray(v) _require_member_shape(X, v, "v", "prox_l1") - return generalized_shrinkage(X, c=X.zeros(), x0=v, eps=0.5, lam=t) + return generalized_shrinkage(X, c=X.zeros(), x0=v, eps=0.5, lam=t, nonneg=nonneg) -def prox_l2sq(v: Any, t: Any, X: Any) -> Any: +def prox_l2sq(v: Any, t: Any, X: Any, *, nonneg: bool = False) -> Any: r""" Proximal operator of ``t * (1/2) ||.||_X^2`` (linear shrinkage ``v / (1 + t)``). @@ -181,6 +216,11 @@ def prox_l2sq(v: Any, t: Any, X: Any) -> Any: minimizer of ``<-v, x>_X + ((1 + t)/2) ||x||^2_X``, i.e. ``generalized_shrinkage`` with ``c = -v``, ``x0 = 0`` and ``eps = (1 + t)/2``. + With ``nonneg=True`` the minimization is constrained to ``x >= 0``. The + objective is separable and strictly convex, so the constrained minimizer is + the unconstrained one clipped coordinatewise: + ``max(v_i / (1 + t), 0) = max(v_i, 0) / (1 + t)`` since ``1 + t > 0``. + Parameters ---------- v : array-like @@ -189,6 +229,9 @@ def prox_l2sq(v: Any, t: Any, X: Any) -> Any: Nonnegative step size. X : Space Ambient inner-product space (Euclidean or diagonal metric). + nonneg : bool, optional + Constrain the solution to ``x >= 0`` (real spaces only). Default + ``False``. With ``t = 0`` this degenerates to :func:`project_nonneg`. Returns ------- @@ -198,7 +241,9 @@ def prox_l2sq(v: Any, t: Any, X: Any) -> Any: t = _require_step(t, "prox_l2sq") v = X.ctx.asarray(v) _require_member_shape(X, v, "v", "prox_l2sq") - return generalized_shrinkage(X, c=-v, x0=X.zeros(), eps=0.5 * (1.0 + t), lam=0.0) + return generalized_shrinkage( + X, c=-v, x0=X.zeros(), eps=0.5 * (1.0 + t), lam=0.0, nonneg=nonneg + ) def project_nonneg(v: Any, X: Any) -> Any: diff --git a/spacecore/functional/tools/_spectral.py b/spacecore/functional/tools/_spectral.py index 8be3fba..c148a2d 100644 --- a/spacecore/functional/tools/_spectral.py +++ b/spacecore/functional/tools/_spectral.py @@ -1,54 +1,127 @@ -"""Spectral (Schatten) ``p``-norm functional over a Jordan-algebra spectrum (ADR-019). - -Where :class:`~spacecore.functional.LpNormFunctional` is the ``p``-norm of the -*coordinates*, :class:`SpectralLpNormFunctional` is the ``p``-norm of the -*spectrum*: for a Hermitian element ``X = U diag(lambda) U^*`` it is -``(sum_i |lambda_i|^p)^{1/p}`` -- the Schatten-``p`` norm (nuclear norm at -``p = 1``, Frobenius at ``p = 2``). - -It is a *spectral function* ``F(X) = f(lambda(X))`` with ``f`` the symmetric -coordinate ``p``-norm. Its gradient is the spectral function gradient -``U diag(grad f(lambda)) U^*`` (Lewis), which is exactly -``from_spectrum(grad f(lambda), frame)`` on the [ADR-012](012_jordan_spectrum.md) -Jordan spectral API. Building it on ``spectrum`` / ``spectral_decompose`` / -``from_spectrum`` (rather than reaching for backend ``eigh``) keeps it correct on -every Jordan space: on an elementwise Jordan space the spectrum *is* the -coordinates, so it coincides with ``LpNormFunctional``. +"""Spectral lift of coordinate functionals over a Jordan-algebra spectrum (ADR-019). + +A *spectral function* is ``F(X) = f(lambda(X))`` for a symmetric ``f`` of the +eigenvalues. By Lewis's theorem its gradient is ``U diag(grad f(lambda)) U^*``, +which is exactly ``from_spectrum(grad f(lambda), frame)`` on the +[ADR-012](012_jordan_spectrum.md) Jordan spectral API. Building on +``spectrum`` / ``spectral_decompose`` / ``from_spectrum`` rather than reaching +for a backend ``eigh`` keeps it correct on every Jordan space; on an elementwise +Jordan space the spectrum *is* the coordinates, so a lifted functional coincides +with its coordinate original. + +:class:`SpectralFunctional` implements that lift once, for **any** coordinate +functional. There is deliberately no per-formula spectral class: a Schatten +``p``-norm is ``spectralize(X, lambda s: LpNormFunctional(s, p))``, the von +Neumann entropy is ``spectralize(X, NegativeEntropyFunctional)``, and so on. The +alternative — one hand-written class per formula — duplicates the value, the +gradient and the validation of its coordinate twin, which is how the retired +``SpectralLpNormFunctional`` came to re-check ``p >= 1`` that +``LpNormFunctional`` already enforced. + +References +---------- +.. [Lewis] A. S. Lewis, Theorem 2.3.2 in H. Wolkowicz, R. Saigal, L. Vandenberghe + (eds.), *Handbook of Semidefinite Programming*, Kluwer, 2000, §2.3.2 + "Smoothness of eigenvalues", pp. 18-19: for a *permutation-invariant* ``f``, + ``F(X) = f(lambda(X))`` is (Frechet) differentiable at ``X`` **if and only if** + ``f`` is differentiable at ``lambda(X)``, and then + ``DF(X) = U^T Diag(f'(lambda(X))) U`` for any orthogonal ``U`` diagonalizing + ``X``. Permutation-invariance is the hypothesis, not a convenience — it is why + :class:`SpectralFunctional` states symmetry as a caller contract. +.. [Beck7] A. Beck, *First-Order Methods in Optimization*, MOS-SIAM, 2017, + Ch. 7: Definition 7.11 (spectral functions over ``S^n``), Definition 7.12 + (symmetric spectral functions), Theorem 7.9 (symmetric conjugate theorem) and + §7.2.2 (the proximal operator of a symmetric spectral function) — the modern + optimization treatment, including the conjugate and prox of the lift. """ from __future__ import annotations -import math from typing import Any, cast -from .._base import Domain -from ...backend import Context, jax_pytree_class +from .._base import Domain, Functional +from ...contextual import Context +from ..._check_policy import CheckLevel from ...space import JordanAlgebraSpace from ..._checks import checked_method -from ._coordinate import _CoordinateFunctional, lp_coordinate_grad, lp_value +from ._coordinate import _CoordinateFunctional +from ._norms import LpNormFunctional + + +def eigenvalue_space(dom: Any, check_level: CheckLevel | bool | None = None) -> Any: + r""" + Return the real coordinate space the spectrum of ``dom`` lives in. + + The spectrum of a Jordan-algebra element is a real vector of length equal to + the algebra's rank, regardless of whether the element itself is stored with + complex entries (a Hermitian matrix has real eigenvalues). Rank is read off + the zero element rather than declared, so this works for any Jordan space + without a ``rank`` attribute. + + Parameters + ---------- + dom : JordanAlgebraSpace + Domain whose spectrum is to be modelled. + check_level : {"none", "cheap", "standard", "strict"}, optional + Validation policy for the returned space. Defaults to ``dom``'s. + + Returns + ------- + DenseCoordinateSpace + Euclidean real space of shape ``(rank,)``. + """ + from ...space import DenseCoordinateSpace + if not isinstance(dom, JordanAlgebraSpace): + raise TypeError( + "eigenvalue_space requires a Jordan-algebra domain with a spectral " + f"decomposition; got {type(dom).__name__}." + ) + ops = dom.ops + rank = int(dom.spectrum(dom.zeros()).shape[0]) + real_ctx = Context(ops, dtype=ops.real_dtype(dom.dtype)) + level = dom.check_level if check_level is None else check_level + return DenseCoordinateSpace((rank,), real_ctx, check_level=level) -@jax_pytree_class -class SpectralLpNormFunctional(_CoordinateFunctional[Domain]): + +class SpectralFunctional(_CoordinateFunctional[Domain]): r""" - Schatten ``p``-norm ``F(X) = (sum_i |lambda_i(X)|^p)^{1/p}`` for ``p >= 1``. + Lift a **symmetric** coordinate functional onto a Jordan spectrum. + + Given ``f`` on :math:`\mathbb{R}^r` and a Jordan space of rank :math:`r`, + this is the spectral function :math:`F(X) = f(\lambda(X))`. By Lewis's + theorem its gradient is :math:`U \operatorname{diag}(\nabla f(\lambda)) U^*`, + which is exactly ``from_spectrum(grad f(lambda), frame)`` on the ADR-012 + spectral API. - ``lambda(X)`` is the Jordan-algebraic spectrum of ``X`` (eigenvalues for a - Hermitian matrix). The gradient is the spectral function gradient: with - ``X = U diag(lambda) U^*``, it is ``U diag(g) U^*`` where ``g`` is the - coordinate ``p``-norm gradient of ``lambda``, reconstructed through the - space's ``from_spectrum``. ``p = 1`` is the nuclear / trace norm (see - :func:`NuclearNormFunctional`); ``p = 2`` is the Frobenius norm. + Any coordinate functional gains a spectral counterpart by wrapping rather + than by being re-implemented: ``LpNormFunctional`` lifts to the Schatten + ``p``-norm (nuclear at ``p = 1``, Frobenius at ``p = 2``), + ``NegativeEntropyFunctional`` to the **von Neumann entropy**, + ``SquaredL2NormFunctional`` to the squared Frobenius norm, and + ``HuberFunctional`` to its spectral analogue. + + .. warning:: + + ``f`` **must be symmetric** (permutation-invariant). This is not a + technicality about gradients: eigenvalues have no canonical order, so for + a non-symmetric ``f`` the composition :math:`f(\lambda(X))` is not even a + well-defined function of :math:`X` — its value would depend on the + backend's sort convention. Symmetry cannot be checked programmatically, + so it is a contract the caller must honour. + :class:`~spacecore.KLDivergenceFunctional` is the notable member of the + toolbox that does **not** qualify, being weighted by a fixed target. Parameters ---------- dom : JordanAlgebraSpace - Domain space with a spectral decomposition (e.g. - :class:`~spacecore.HermitianSpace`). A space without a Jordan spectrum - raises ``TypeError``. - p : float - Norm order; must be finite and ``>= 1``. + Domain with a spectral decomposition (e.g. :class:`~spacecore.HermitianSpace`). + base : Functional + Symmetric functional on the Euclidean real space of shape ``(rank,)``; + see :func:`eigenvalue_space`. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom``. + check_level : {"none", "cheap", "standard", "strict"}, optional + Validation policy for this functional. Examples -------- @@ -56,60 +129,139 @@ class SpectralLpNormFunctional(_CoordinateFunctional[Domain]): >>> import spacecore as sc >>> ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) >>> X = sc.HermitianSpace(2, ctx=ctx) - >>> A = ctx.asarray([[2.0, 0.0], [0.0, -3.0]]) - >>> f = sc.SpectralLpNormFunctional(X, 1) # nuclear norm |2| + |-3| - >>> float(f.value(A)) - 5.0 + >>> S = sc.eigenvalue_space(X) + >>> vn = sc.SpectralFunctional(X, sc.NegativeEntropyFunctional(S)) + >>> rho = ctx.asarray([[0.5, 0.0], [0.0, 0.5]]) # maximally mixed state + >>> float(vn.value(rho)) # -log 2, i.e. sum p log p + -0.6931471805599453 """ - def __init__(self, dom: Domain, p: Any, ctx: Context | str | None = None) -> None: - super().__init__(dom, ctx) - if not isinstance(self.domain, JordanAlgebraSpace): + def __init__( + self, + dom: Domain, + base: Functional, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: + super().__init__(dom, ctx, check_level=check_level) + domain = cast(Any, self.domain) + if not isinstance(domain, JordanAlgebraSpace): raise TypeError( - "SpectralLpNormFunctional requires a Jordan-algebra domain with a " + "SpectralFunctional requires a Jordan-algebra domain with a " f"spectral decomposition (e.g. HermitianSpace); got " - f"{type(self.domain).__name__}." + f"{type(domain).__name__}." + ) + if not isinstance(base, Functional): + raise TypeError(f"base must be a Functional, got {type(base).__name__}.") + expected = eigenvalue_space(domain) + if base.domain.shape != expected.shape: + raise ValueError( + f"base domain shape {base.domain.shape} does not match the " + f"spectrum of {type(domain).__name__} (rank {expected.shape[0]}); " + f"build it on eigenvalue_space(dom)." + ) + if not getattr(base.domain, "is_euclidean", True): + # Lewis needs the EUCLIDEAN gradient of f. On a Euclidean eigenvalue + # space base.grad already is that; under any other metric it would be + # G^-1 df, silently scaling the reconstructed spectral gradient. + raise ValueError( + "SpectralFunctional requires a Euclidean eigenvalue space so that " + "base.grad is the plain coordinate gradient Lewis's formula needs." ) - p = float(p) - if not math.isfinite(p) or p < 1.0: - raise ValueError(f"SpectralLpNormFunctional requires a finite p >= 1, got {p}.") - self.p = p + self.base = base - @checked_method(in_space="domain") + @checked_method(in_space="domain", out_scalar=True) def value(self, x: Any) -> Any: - """Return the Schatten-``p`` norm ``(sum_i |lambda_i|^p)^{1/p}``.""" - spectrum = cast(Any, self.domain).spectrum(x) - return lp_value(self.ops, spectrum, self.p) + """Return ``f(lambda(x))``.""" + return self.base.value(cast(Any, self.domain).spectrum(x)) def _coordinate_grad(self, x: Any) -> Any: - """Spectral gradient ``U diag(grad f(lambda)) U^*`` via ``from_spectrum``.""" + """Return ``U diag(grad f(lambda)) U^*`` via ``from_spectrum`` (Lewis).""" domain = cast(Any, self.domain) eigvals, frame = domain.spectral_decompose(x) - spectral_grad = lp_coordinate_grad(self.ops, eigvals, self.p) - return domain.from_spectrum(spectral_grad, frame) + return domain.from_spectrum(self.base.grad(eigvals), frame) + + def __eq__(self, other: Any) -> bool: + """Return whether another spectral functional has the same domain and base.""" + if not self.same_math(other): + return NotImplemented + return self.domain == other.domain and self.base == other.base def tree_flatten(self): """Flatten this functional for pytree registration.""" - return (), (self.domain, self.p, self.ctx) + return (self.base,), (self.domain, self.ctx) @classmethod def tree_unflatten(cls, aux, children): """Rebuild this functional from pytree data.""" - domain, p, ctx = aux - return cls(domain, p, ctx) + domain, ctx = aux + (base,) = children + return cls(domain, base, ctx) + + def _convert(self, new_ctx: Context) -> "SpectralFunctional": + """Convert this functional and its base to ``new_ctx``.""" + return SpectralFunctional( + self.domain.convert(new_ctx), self.base.convert(new_ctx), new_ctx + ) + - def _convert(self, new_ctx: Context) -> "SpectralLpNormFunctional": - """Convert this functional to ``new_ctx``.""" - return SpectralLpNormFunctional(self.domain.convert(new_ctx), self.p, new_ctx) +def spectralize( + dom: Domain, + make_base: Any, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, +) -> "SpectralFunctional[Domain]": + r""" + Build the spectral counterpart of a coordinate functional. + + Constructs the eigenvalue space for ``dom``, hands it to ``make_base``, and + wraps the result — so a caller never has to derive the rank or build the + intermediate space by hand. + + Parameters + ---------- + dom : JordanAlgebraSpace + Domain with a spectral decomposition. + make_base : callable + ``make_base(eigenvalue_space) -> Functional``; the functional must be + symmetric (see :class:`SpectralFunctional`). + ctx : Context, str, or None, optional + Backend context specification. + check_level : {"none", "cheap", "standard", "strict"}, optional + Validation policy. + + Returns + ------- + SpectralFunctional + ``f(lambda(.))`` on ``dom``. + + Examples + -------- + >>> import numpy as np + >>> import spacecore as sc + >>> ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + >>> X = sc.HermitianSpace(2, ctx=ctx) + >>> vn = sc.spectralize(X, sc.NegativeEntropyFunctional) # von Neumann entropy + >>> float(vn.value(ctx.asarray([[0.5, 0.0], [0.0, 0.5]]))) + -0.6931471805599453 + """ + base = make_base(eigenvalue_space(dom, check_level=check_level)) + return SpectralFunctional(dom, base, ctx, check_level=check_level) def NuclearNormFunctional( - dom: Domain, ctx: Context | str | None = None -) -> "SpectralLpNormFunctional[Domain]": + dom: Domain, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, +) -> "SpectralFunctional[Domain]": r""" - Nuclear (trace) norm, a thin wrapper for ``SpectralLpNormFunctional(X, 1)``. + Nuclear (trace) norm ``sum_i |lambda_i(X)|`` — the Schatten-1 norm. - Computes ``sum_i |lambda_i(X)|``, the Schatten-1 norm of the Jordan spectrum. + Kept as a named constructor because the nuclear norm is a concept in its own + right (the convex envelope of rank, and the workhorse of low-rank recovery), + not because the formula needs its own class: it is exactly + ``spectralize(dom, lambda s: LpNormFunctional(s, 1.0))``, and ``p >= 1`` is + validated once, by :class:`~spacecore.LpNormFunctional`. Parameters ---------- @@ -117,10 +269,24 @@ def NuclearNormFunctional( Domain space with a spectral decomposition. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom``. + check_level : {"none", "cheap", "standard", "strict"}, optional + Validation policy for the constructed functional. Returns ------- - SpectralLpNormFunctional - The ``p = 1`` instance of :class:`SpectralLpNormFunctional`. + SpectralFunctional + The Schatten-1 lift of the coordinate 1-norm. + + Examples + -------- + >>> import numpy as np + >>> import spacecore as sc + >>> ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + >>> X = sc.HermitianSpace(2, ctx=ctx) + >>> f = sc.NuclearNormFunctional(X) + >>> float(f.value(ctx.asarray([[2.0, 0.0], [0.0, -3.0]]))) + 5.0 """ - return SpectralLpNormFunctional(dom, 1.0, ctx) + return spectralize( + dom, lambda s: LpNormFunctional(s, 1.0), ctx, check_level=check_level + ) diff --git a/spacecore/kernels/core/algebra.py b/spacecore/kernels/core/algebra.py index 25c1434..94a5d7f 100644 --- a/spacecore/kernels/core/algebra.py +++ b/spacecore/kernels/core/algebra.py @@ -22,16 +22,11 @@ from ._rules import CoreKernelSet, register_core_kernels from ..specs._dispatch import dispatch, should_consult_dispatch -# --------------------------------------------------------------------------- -# Shared rules / helpers -# --------------------------------------------------------------------------- -def conjugate_scalar(value: Any) -> Any: - """Return the scalar conjugate when the value supports conjugation.""" - if hasattr(value, "conjugate"): - return value.conjugate() - if hasattr(value, "conj"): - return value.conj() - return value +# Re-exported: the scalar predicates/helpers live together in +# :mod:`spacecore._lazy_algebra` (stdlib-only, so importable from anywhere), +# next to ``is_scalar_like`` / ``scalar_eq`` which share their duck-typing +# assumptions. Kept importable from here for the kernel call sites. +from ..._lazy_algebra import conjugate_scalar as conjugate_scalar # noqa: F401 def compose_chain(op: Any) -> tuple: @@ -88,9 +83,9 @@ def _composed_chain_apply(chain: Any, x: Any) -> Any: def composed_apply_core(op: Any, x: Any) -> Any: chain = op._apply_chain - if should_consult_dispatch(op.ctx): + if should_consult_dispatch(op): return dispatch( - _COMPOSED_APPLY_KEY, chain, x, generic=_composed_chain_apply, ctx=op.ctx + _COMPOSED_APPLY_KEY, chain, x, generic=_composed_chain_apply, ctx=op ) return _composed_chain_apply(chain, x) diff --git a/spacecore/kernels/core/functional.py b/spacecore/kernels/core/functional.py index b003819..81d44ec 100644 --- a/spacecore/kernels/core/functional.py +++ b/spacecore/kernels/core/functional.py @@ -127,6 +127,17 @@ def composed_functional_value_core(op: Any, x: Any) -> Any: return op.F._value_core(op.A._apply_core(x)) +def composed_functional_grad_core(op: Any, x: Any) -> Any: + # Chain rule in Riesz form: grad(F o A)(x) = A^#(grad F(A x)). + # ``rapply`` IS the metric adjoint A^# (ADR-009), so no extra Riesz map is + # applied here; inserting one would count the geometry twice. + return op.A._rapply_core(op.F._grad_core(op.A._apply_core(x))) + + +def composed_functional_vgrad_core(op: Any, xs: Any) -> Any: + return op.A._rvapply_core(op.F._vgrad_core(op.A._vapply_core(xs))) + + register_core_kernels(CoreKernelSet( "functional-linear", grad=linear_grad_core, vgrad=linear_vgrad_core, notes="Constant Riesz gradient of a linear functional.", @@ -148,6 +159,8 @@ def composed_functional_value_core(op: Any, x: Any) -> Any: notes="Reaches Q and linear-term cores without re-validating intermediates.", )) register_core_kernels(CoreKernelSet( - "composed-functional", value=composed_functional_value_core, - notes="Pull-back F(Ax) via the operand cores.", + "composed-functional", + value=composed_functional_value_core, grad=composed_functional_grad_core, + vgrad=composed_functional_vgrad_core, + notes="Pull-back F(Ax) and its chain rule A^#(grad F(Ax)) via the operand cores.", )) diff --git a/spacecore/kernels/specs/_dispatch.py b/spacecore/kernels/specs/_dispatch.py index 229bc81..640a084 100644 --- a/spacecore/kernels/specs/_dispatch.py +++ b/spacecore/kernels/specs/_dispatch.py @@ -122,7 +122,13 @@ def _is_strict(ctx: Any) -> bool: def effective_mode(ctx: Any = None) -> DispatchMode: - """Resolve the dispatch mode for an operand context. + """Resolve the dispatch mode for an operand. + + ``ctx`` is the context-*bound* operand (space, operator, or functional), not + a :class:`Context`. Only two attributes are read from it: ``check_level`` + for the strict rule below, and ``ops`` for the memory gate. Validation + policy lives on the bound object, so passing a bare ``Context`` silently + disables the strict rule. ``check_level="strict"`` (ADR-014) implies ``verify``, overriding both the context override and the global default — the strictest policy always runs @@ -243,8 +249,11 @@ def dispatch( dispatch is ``off``, when no eligible spec applies/fits, and (in ``verify``) as the value the optimized result is checked against. ctx : optional - The operand execution context. Supplies the check level (for the - ``strict`` → ``verify`` rule) and the backend (for the memory gate). + The context-*bound* operand (space, operator, or functional) — not a + :class:`Context`. Supplies ``check_level`` (for the ``strict`` → + ``verify`` rule) and ``ops`` (for the memory gate). A bare ``Context`` + carries no ``check_level``, so passing one silently disables the strict + rule. Returns ------- diff --git a/spacecore/linalg/_power.py b/spacecore/linalg/_power.py index af5ca11..711d90f 100644 --- a/spacecore/linalg/_power.py +++ b/spacecore/linalg/_power.py @@ -3,7 +3,7 @@ from collections.abc import Callable from typing import Any, NamedTuple, cast -from ..backend import Context +from ..contextual import Context from ..functional import QuadraticForm from ..linop import LinOp from ..space import Space diff --git a/spacecore/linop/_algebra.py b/spacecore/linop/_algebra.py index 307ee86..ce8f33c 100644 --- a/spacecore/linop/_algebra.py +++ b/spacecore/linop/_algebra.py @@ -1,16 +1,15 @@ from __future__ import annotations +import warnings from math import prod -from numbers import Number from typing import Any, Callable, Sequence, cast from ._base import LinOp, Domain, Codomain from ._metric import _requires_euclidean_or_riesz, metric_rapply, metric_rvapply +from .._check_policy import CheckLevel, minimum_check_level from .._checks import checked_method -from .._contextual import resolve_context_priority -from .._contextual._bound import _same_math_context from .._repr import summarize_value -from ..backend import Context, jax_pytree_class +from ..contextual import Context from ..kernels import core_kernels from ..kernels.core.algebra import ( batched_zeros as _batched_zeros, @@ -20,40 +19,32 @@ ) -def is_scalar_like(value: Any) -> bool: - """Return whether ``value`` can be used as a scalar multiplier for a ``LinOp``.""" - if isinstance(value, Number): - return True - shape = getattr(value, "shape", None) - if shape is not None: - return tuple(shape) == () - ndim = getattr(value, "ndim", None) - return ndim == 0 - +from .._lazy_algebra import ( + finalize_sum, + flatten_sum, + fold_scaled, + is_scalar_like as is_scalar_like, # re-exported for linop/_base.py + scalar_eq, +) -def _scalar_eq(a: Any, b: Any) -> bool: - """Return whether two scalar-likes are equal, NaN-reflexive. - Mirrors the ``equal_nan=True`` used for array values: two matching NaN - scalars compare equal so a NaN-scaled operator equals itself. Always returns - a real Python ``bool`` (a 0-d backend-array ``==`` would otherwise yield - ``np.bool_``, which leaks through the ``and`` combinator of any container). - """ - if bool(a == b): - return True - try: - # ``x != x`` is True only for NaN (including a complex value with a NaN - # component), so this branch matches NaN against NaN. - return bool(a != a) and bool(b != b) - except Exception: - return False +#: Sentinel for "the caller did not mention ``euclidean_adjoint`` at all". +#: Passing ``False`` explicitly is an assertion that ``rapply`` already carries the +#: geometry, so it silences the advisory; omitting it means the question was never +#: considered, which is the only case worth warning about. +_ADJOINT_UNSPECIFIED = object() def _require_same_context(ops: Sequence[LinOp]) -> Context: - """Return the common context for algebra operands or raise.""" + """Return the common mathematical context for algebra operands or raise. + + Operands must share a mathematical context (:meth:`Context.same_math`). The + surviving ``check_level`` is a property of the resulting bound object and is + combined (via the minimum) when the container binds its operands, not here. + """ ctx = ops[0].ctx for i, op in enumerate(ops[1:], start=1): - if not _same_math_context(ops[0].ctx, op.ctx): + if not ctx.same_math(op.ctx): raise ValueError( "All LinOp operands in an algebraic expression must have the same ctx; " f"operand 0 has ctx {ctx!r}, operand {i} has ctx {op.ctx!r}." @@ -69,7 +60,7 @@ def _same_space_for_algebra(left: Any, right: Any) -> bool: return False if tuple(left.shape) != tuple(right.shape): return False - if not _same_math_context(left.ctx, right.ctx): + if not left.ctx.same_math(right.ctx): return False try: return left.convert(right.ctx) == right @@ -84,36 +75,6 @@ def _require_linop(op: Any, name: str) -> LinOp: return op -def _scalar_equal(value: Any, target: Any) -> bool: - """Return whether two scalar-like values compare equal.""" - try: - return bool(value == target) - except Exception: - return False - - -def _is_zero_scalar(value: Any) -> bool: - """Return whether ``value`` is scalar-like zero.""" - return _scalar_equal(value, 0) - - -def _is_one_scalar(value: Any) -> bool: - """Return whether ``value`` is scalar-like one.""" - return _scalar_equal(value, 1) - - -def _flatten_sum_terms(ops: Sequence[LinOp]) -> tuple[LinOp, ...]: - """Flatten nested lazy sums into a tuple of terms.""" - terms: list[LinOp] = [] - for i, op in enumerate(ops): - op = _require_linop(op, f"ops[{i}]") - if isinstance(op, SumLinOp): - terms.extend(_flatten_sum_terms(op.parts)) - else: - terms.append(op) - return tuple(terms) - - def make_sum(ops: Sequence[LinOp]) -> LinOp: """ Return a locally simplified lazy sum of linear operators. @@ -137,7 +98,11 @@ def make_sum(ops: Sequence[LinOp]) -> LinOp: if not ops: raise ValueError("make_sum requires a nonempty sequence of LinOp operands.") - terms = _flatten_sum_terms(ops) + terms = flatten_sum( + tuple(_require_linop(op, f"ops[{i}]") for i, op in enumerate(ops)), + is_sum=lambda t: isinstance(t, SumLinOp), + parts=lambda t: t.parts, + ) ctx = _require_same_context(terms) domain = terms[0].domain codomain = terms[0].codomain @@ -151,12 +116,12 @@ def make_sum(ops: Sequence[LinOp]) -> LinOp: f"operand {i} maps {op.domain!r} -> {op.codomain!r}." ) - nonzero_terms = tuple(op for op in terms if not isinstance(op, ZeroLinOp)) - if not nonzero_terms: - return ZeroLinOp(domain, codomain, ctx) - if len(nonzero_terms) == 1: - return nonzero_terms[0] - return SumLinOp(nonzero_terms) + return finalize_sum( + terms, + is_zero=lambda t: isinstance(t, ZeroLinOp), + make_zero=lambda: ZeroLinOp(domain, codomain, ctx), + make_sum_node=SumLinOp, + ) def make_scaled(scalar: Any, op: LinOp) -> LinOp: @@ -185,15 +150,14 @@ def make_scaled(scalar: Any, op: LinOp) -> LinOp: if not is_scalar_like(scalar): raise TypeError(f"scalar must be scalar-like, got {type(scalar).__name__}.") - if _is_zero_scalar(scalar): - return ZeroLinOp(op.domain, op.codomain, op.ctx) - if _is_one_scalar(scalar): - return op - if isinstance(op, ZeroLinOp): - return op - if isinstance(op, ScaledLinOp): - return make_scaled(scalar * op.scalar, op.op) - return ScaledLinOp(scalar, op) + return fold_scaled( + scalar, + op, + is_zero=lambda o: isinstance(o, ZeroLinOp), + unwrap_scaled=lambda o: (o.scalar, o.op) if isinstance(o, ScaledLinOp) else None, + make_zero=lambda: ZeroLinOp(op.domain, op.codomain, op.ctx), + make_scaled_node=ScaledLinOp, + ) def make_composed(left: LinOp, right: LinOp) -> LinOp: @@ -240,7 +204,6 @@ def make_composed(left: LinOp, right: LinOp) -> LinOp: @core_kernels("scaled") -@jax_pytree_class class ScaledLinOp(LinOp[Domain, Codomain]): r""" Lazy scalar multiple of a linear operator. @@ -261,6 +224,11 @@ class ScaledLinOp(LinOp[Domain, Codomain]): Scalar multiplier. op : LinOp Operator being scaled. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. Attributes ---------- @@ -270,11 +238,21 @@ class ScaledLinOp(LinOp[Domain, Codomain]): Stored operand. """ - def __init__(self, scalar: Any, op: LinOp[Domain, Codomain]) -> None: + def __init__( + self, + scalar: Any, + op: LinOp[Domain, Codomain], + check_level: CheckLevel | bool | None = None, + ) -> None: op = _require_linop(op, "op") if not is_scalar_like(scalar): raise TypeError(f"scalar must be scalar-like, got {type(scalar).__name__}.") - super().__init__(op.domain, op.codomain, op.ctx) + # Default the policy from the OPERAND, not from its domain/codomain + # spaces: an algebra node inherits the least-strict operand so the + # result does not depend on operand order. + if check_level is None: + check_level = op.check_level + super().__init__(op.domain, op.codomain, op.ctx, check_level=check_level) self.scalar = scalar self.op = op @@ -340,10 +318,10 @@ def is_hermitian(self) -> bool | None: def __eq__(self, other: Any) -> bool: """Return whether another scaled operator has the same scalar and operand.""" - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented # NaN-reflexive, returns a real Python bool (no np.bool_ leak). - if not _scalar_eq(self.scalar, other.scalar): # Tier 3: scalar value + if not scalar_eq(self.scalar, other.scalar): # Tier 3: scalar value return False return self.op == other.op # operand (own gate) @@ -368,7 +346,6 @@ def _convert(self, new_ctx: Context) -> ScaledLinOp: @core_kernels("sum") -@jax_pytree_class class SumLinOp(LinOp[Domain, Codomain]): r""" Lazy finite sum of linear operators with common spaces. @@ -387,6 +364,11 @@ class SumLinOp(LinOp[Domain, Codomain]): ops : sequence of LinOp Nonempty sequence of operators with common context, domain, and codomain. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. Attributes ---------- @@ -394,7 +376,11 @@ class SumLinOp(LinOp[Domain, Codomain]): Stored operands in the lazy sum. """ - def __init__(self, ops: Sequence[LinOp[Domain, Codomain]]) -> None: + def __init__( + self, + ops: Sequence[LinOp[Domain, Codomain]], + check_level: CheckLevel | bool | None = None, + ) -> None: if not ops: raise ValueError("SumLinOp requires a nonempty sequence of LinOp operands.") parts = tuple(_require_linop(op, f"ops[{i}]") for i, op in enumerate(ops)) @@ -410,7 +396,9 @@ def __init__(self, ops: Sequence[LinOp[Domain, Codomain]]) -> None: f"operand 0 maps {domain!r} -> {codomain!r}, " f"operand {i} maps {op.domain!r} -> {op.codomain!r}." ) - super().__init__(domain, codomain, ctx) + if check_level is None: + check_level = minimum_check_level(tuple(op.check_level for op in parts)) + super().__init__(domain, codomain, ctx, check_level=check_level) self.ops_tuple = parts @property @@ -489,7 +477,7 @@ def is_hermitian(self) -> bool | None: def __eq__(self, other: Any) -> bool: """Return whether another sum has the same operands, in order.""" - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented if len(self.ops_tuple) != len(other.ops_tuple): # Tier 2: operand count before zip return False @@ -519,7 +507,6 @@ def _convert(self, new_ctx: Context) -> SumLinOp: @core_kernels("composed") -@jax_pytree_class class ComposedLinOp(LinOp[Domain, Codomain]): r""" Lazy composition of two linear operators. @@ -539,6 +526,11 @@ class ComposedLinOp(LinOp[Domain, Codomain]): Operator applied second. right : LinOp Operator applied first. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. Attributes ---------- @@ -548,7 +540,12 @@ class ComposedLinOp(LinOp[Domain, Codomain]): Right operand. """ - def __init__(self, left: LinOp, right: LinOp) -> None: + def __init__( + self, + left: LinOp, + right: LinOp, + check_level: CheckLevel | bool | None = None, + ) -> None: left = _require_linop(left, "left") right = _require_linop(right, "right") _require_same_context((left, right)) @@ -557,7 +554,9 @@ def __init__(self, left: LinOp, right: LinOp) -> None: "ComposedLinOp requires right.codomain == left.domain; " f"got {right.codomain!r} and {left.domain!r}." ) - super().__init__(right.domain, left.codomain, left.ctx) + if check_level is None: + check_level = minimum_check_level((left.check_level, right.check_level)) + super().__init__(right.domain, left.codomain, left.ctx, check_level=check_level) self.left = left self.right = right # Fuse the (possibly nested) composition into one flat chain of leaf @@ -632,7 +631,7 @@ def is_hermitian(self) -> bool | None: def __eq__(self, other: Any) -> bool: """Return whether another composition has the same operands, in order.""" - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented return self.left == other.left and self.right == other.right @@ -657,7 +656,6 @@ def _convert(self, new_ctx: Context) -> ComposedLinOp: @core_kernels("zero") -@jax_pytree_class class ZeroLinOp(LinOp[Domain, Codomain]): r""" Lazy zero map between two spaces. @@ -678,6 +676,11 @@ class ZeroLinOp(LinOp[Domain, Codomain]): Codomain space. ctx : Context, str, or None, optional Backend context specification. Default is resolved from the spaces. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. """ def __init__( @@ -685,8 +688,9 @@ def __init__( dom: Domain, cod: Codomain, ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: - super().__init__(dom, cod, ctx) + super().__init__(dom, cod, ctx, check_level=check_level) @checked_method(in_space="domain", out_space="codomain") def apply(self, x: Any) -> Any: @@ -731,7 +735,7 @@ def is_hermitian(self) -> bool: def __eq__(self, other: Any) -> bool: """Return whether another zero map has the same spaces.""" - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented return self.domain == other.domain and self.codomain == other.codomain # Tier 2 @@ -753,7 +757,6 @@ def _convert(self, new_ctx: Context) -> ZeroLinOp: @core_kernels("identity") -@jax_pytree_class class IdentityLinOp(LinOp[Domain, Domain]): r""" Lazy identity map on a space. @@ -771,10 +774,20 @@ class IdentityLinOp(LinOp[Domain, Domain]): Domain and codomain space. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``space``. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. """ - def __init__(self, space: Domain, ctx: Context | str | None = None) -> None: - super().__init__(space, space, ctx) + def __init__( + self, + space: Domain, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: + super().__init__(space, space, ctx, check_level=check_level) @checked_method(in_space="domain", out_space="codomain") def apply(self, x: Any) -> Any: @@ -821,7 +834,7 @@ def is_hermitian(self) -> bool: def __eq__(self, other: Any) -> bool: """Return whether another identity map has the same space.""" - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented return self.domain == other.domain # Tier 2 (square: cod == dom) @@ -848,7 +861,6 @@ def _convert(self, new_ctx: Context) -> IdentityLinOp: @core_kernels("matrixfree") -@jax_pytree_class class MatrixFreeLinOp(LinOp[Domain, Codomain]): """ Linear operator defined by user-supplied forward and reverse callables. @@ -895,6 +907,27 @@ class MatrixFreeLinOp(LinOp[Domain, Codomain]): Optional callable with signature ``rvapply(ys: Any) -> Any`` for batched adjoint application. If omitted, backend ``vmap`` fallback is used. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. + euclidean_adjoint : bool, optional + How ``rapply`` is interpreted. Omitted or ``False`` (the default): the + callable **is** the metric adjoint and is stored verbatim -- the design + choice that lets a caller supply a hand-derived adjoint no wrapper could + produce. ``True``: the callable is the *Euclidean coordinate* adjoint + ``A^dagger``, wrapped once into ``R_X^-1 A^dagger R_Y``; this is exactly + what :meth:`from_coordinate_adjoint` does, and it requires usable Riesz + maps on any non-Euclidean space. + + Because verbatim storage is trusted and unverifiable, *omitting* the + argument on a non-Euclidean geometry emits a ``UserWarning``: there a + coordinate adjoint is silently wrong and nothing downstream detects it. + Passing ``False`` **explicitly** asserts that the supplied adjoint + already carries the geometry and silences the advisory -- the + distinction drawn is between "never considered" and "considered and + declared". Returns ------- @@ -918,6 +951,9 @@ def __init__( ctx: Context | str | None = None, vapply: Callable[[Any], Any] | None = None, rvapply: Callable[[Any], Any] | None = None, + check_level: CheckLevel | bool | None = None, + *, + euclidean_adjoint: bool | Any = _ADJOINT_UNSPECIFIED, ) -> None: """ Initialize a matrix-free linear operator. @@ -943,6 +979,9 @@ def __init__( rvapply: Optional callable for batched adjoint application over ``cod`` batches. + check_level: + Optional runtime validation policy for this operator. When omitted, + the ambient default is used. Returns ------- @@ -958,7 +997,39 @@ def __init__( raise TypeError(f"vapply must be callable, got {type(vapply).__name__}.") if rvapply is not None and not callable(rvapply): raise TypeError(f"rvapply must be callable, got {type(rvapply).__name__}.") - super().__init__(dom, cod, ctx) + super().__init__(dom, cod, ctx, check_level=check_level) + + non_euclidean = not (self.domain.is_euclidean and self.codomain.is_euclidean) + if euclidean_adjoint is True: + # The caller supplied the *coordinate* adjoint A^dagger; wrap it into + # the metric adjoint R_X^-1 A^dagger R_Y once, here. + try: + _requires_euclidean_or_riesz( + self.domain, + self.codomain, + "MatrixFreeLinOp(euclidean_adjoint=True) / " + "MatrixFreeLinOp.from_coordinate_adjoint", + ) + except TypeError as exc: + raise ValueError(str(exc)) from exc + rapply, rvapply = self._wrap_coordinate_adjoint(rapply, rvapply) + elif non_euclidean and euclidean_adjoint is _ADJOINT_UNSPECIFIED: + # The callable is stored verbatim and *trusted* as the metric adjoint. + # On a non-Euclidean geometry a coordinate adjoint is simply wrong here + # and nothing downstream detects it, so say so once, at construction. + # Silent on Euclidean spaces (the two coincide) and silent when the + # caller passed ``euclidean_adjoint=False`` explicitly, which asserts + # that the supplied adjoint already carries the geometry. + warnings.warn( + "MatrixFreeLinOp stores `rapply` verbatim as the metric adjoint, but " + f"{'domain' if not self.domain.is_euclidean else 'codomain'} geometry is " + "non-Euclidean, where the metric adjoint R_X^-1 A^dagger R_Y differs from " + "the coordinate adjoint. Pass euclidean_adjoint=True to have the coordinate " + "adjoint wrapped for you, or confirm `rapply` already carries the geometry.", + UserWarning, + stacklevel=2, + ) + self.apply_fn = apply self.rapply_fn = rapply self.vapply_fn = vapply @@ -966,6 +1037,29 @@ def __init__( if self._checks_at_least("strict"): self._check_adjoint_consistency() + def _wrap_coordinate_adjoint( + self, + coordinate_rapply: Callable[[Any], Any], + coordinate_rvapply: Callable[[Any], Any] | None, + ) -> tuple[Callable[[Any], Any], Callable[[Any], Any] | None]: + """Return ``(rapply, rvapply)`` wrapping coordinate adjoints in Riesz maps.""" + dom, cod, ops = self.domain, self.codomain, self.ctx.ops + + def wrapped_rapply(y: Any) -> Any: + return metric_rapply(dom, cod, coordinate_rapply, y) + + wrapped_rvapply: Callable[[Any], Any] | None = None + if coordinate_rvapply is not None: + + def _wrapped_rvapply(ys: Any) -> Any: + return metric_rvapply( + dom, cod, coordinate_rapply, coordinate_rvapply, ys, + opname="MatrixFreeLinOp", ops=ops, + ) + + wrapped_rvapply = _wrapped_rvapply + return wrapped_rapply, wrapped_rvapply + def _check_adjoint_consistency(self) -> None: """Probe the declared adjoint identity on deterministic space elements.""" if not all(hasattr(space, "inner") for space in (self.domain, self.codomain)): @@ -1064,34 +1158,19 @@ def from_coordinate_adjoint( f"coordinate_rvapply must be callable, got {type(coordinate_rvapply).__name__}." ) - resolved_ctx = resolve_context_priority(ctx, dom, cod) - dom = dom.convert(resolved_ctx) - cod = cod.convert(resolved_ctx) - try: - _requires_euclidean_or_riesz(dom, cod, "MatrixFreeLinOp.from_coordinate_adjoint") - except TypeError as exc: - raise ValueError(str(exc)) from exc - - def wrapped_rapply(y: Any) -> Any: - return metric_rapply(dom, cod, coordinate_rapply, y) - - wrapped_rvapply: Callable[[Any], Any] | None = None - if coordinate_rvapply is not None: - - def _wrapped_rvapply(ys: Any) -> Any: - return metric_rvapply( - dom, - cod, - coordinate_rapply, - coordinate_rvapply, - ys, - opname="MatrixFreeLinOp.from_coordinate_adjoint", - ops=resolved_ctx.ops, - ) - - wrapped_rvapply = _wrapped_rvapply - - return cls(apply, wrapped_rapply, dom, cod, resolved_ctx, vapply, wrapped_rvapply) + # The wrapping now lives in __init__ behind euclidean_adjoint=True, so the + # two entry points cannot drift apart, and this path raises no warning: + # passing the coordinate adjoint is exactly what it is documented to take. + return cls( + apply, + coordinate_rapply, + dom, + cod, + ctx, + vapply, + coordinate_rvapply, + euclidean_adjoint=True, + ) @checked_method(in_space="domain", out_space="codomain") def apply(self, x: Any) -> Any: @@ -1183,7 +1262,7 @@ def fuse(self, *, materialize: bool = False) -> LinOp: return DenseLinOp(self.to_dense(), self.domain, self.codomain, self.ctx) def __eq__(self, other: Any) -> bool: - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented # Tier 2: spaces + callable identity. Extensional equality of callables # is undecidable, so 'is' is the only sound comparison. @@ -1237,6 +1316,11 @@ def _convert(self, new_ctx: Context) -> MatrixFreeLinOp: Operator with converted spaces and the same user-supplied callables. """ + # ``self.rapply_fn`` is already the metric adjoint by construction -- either + # supplied that way, or wrapped once at __init__ under + # ``euclidean_adjoint=True``. Declaring ``False`` here is therefore both + # correct (re-wrapping would apply the Riesz maps twice) and quiet: the + # construction advisory belongs to the original call, not to every convert. return MatrixFreeLinOp( self.apply_fn, self.rapply_fn, @@ -1245,11 +1329,11 @@ def _convert(self, new_ctx: Context) -> MatrixFreeLinOp: new_ctx, self.vapply_fn, self.rvapply_fn, + euclidean_adjoint=False, ) @core_kernels("adjoint") -@jax_pytree_class class _AdjointViewLinOp(LinOp[Codomain, Domain]): """ Hermitian-adjoint view of a linear operator. @@ -1260,11 +1344,27 @@ class _AdjointViewLinOp(LinOp[Codomain, Domain]): The forward action is ``apply(y) = A.rapply(y)`` for ``y in A.codomain``. The reverse action is ``rapply(x) = A.apply(x)`` for ``x in A.domain``. + + Parameters + ---------- + op : LinOp + Operator whose adjoint is viewed. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. """ - def __init__(self, op: LinOp[Domain, Codomain]) -> None: + def __init__( + self, + op: LinOp[Domain, Codomain], + check_level: CheckLevel | bool | None = None, + ) -> None: op = _require_linop(op, "op") - super().__init__(op.codomain, op.domain, op.ctx) + if check_level is None: + check_level = op.check_level + super().__init__(op.codomain, op.domain, op.ctx, check_level=check_level) self.op = op @checked_method(in_space="domain", out_space="codomain") @@ -1302,7 +1402,7 @@ def H(self) -> LinOp[Domain, Codomain]: return self.op def __eq__(self, other: Any) -> bool: - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented return self.op == other.op diff --git a/spacecore/linop/_base.py b/spacecore/linop/_base.py index 3f76f47..5fa30cc 100644 --- a/spacecore/linop/_base.py +++ b/spacecore/linop/_base.py @@ -5,20 +5,22 @@ from functools import cached_property from math import prod from numbers import Number -from typing import Any, Generic, Self, TypeVar +from typing import Any, Generic, TypeVar from .._batching import _leading_batch_size, _warn_vmap_fallback_once +from .._check_policy import CheckLevel from .._checks import checked_method +from ..backend import PyTreeNode from .._repr import describe_space from ..space import CoordinateSpace -from ..backend import Context -from .._contextual import ContextBound +from ..contextual import Context +from ..contextual import ContextBound Domain = TypeVar("Domain", bound=CoordinateSpace) Codomain = TypeVar("Codomain", bound=CoordinateSpace) -class LinOp(ContextBound, Generic[Domain, Codomain]): +class LinOp(PyTreeNode, ContextBound, Generic[Domain, Codomain]): r""" Represent a linear map between two spaces. @@ -30,6 +32,12 @@ class LinOp(ContextBound, Generic[Domain, Codomain]): :math:`x \in X` and :math:`y \in Y`. For complex operators this is the conjugate adjoint. + This is the **Hilbert-space adjoint** [Conway]_, defined by the pairing of + each space's own inner product — *not* the coordinate transpose, and not the + Banach (dual-space) adjoint. On a non-Euclidean geometry the two differ: see + :func:`~spacecore.linop._metric.metric_rapply` for the + :math:`R_X^{-1} A^\dagger R_Y` formula that realizes it. + Parameters ---------- dom : Space @@ -39,6 +47,11 @@ class LinOp(ContextBound, Generic[Domain, Codomain]): ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom`` and ``cod``. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. Attributes ---------- @@ -49,6 +62,19 @@ class LinOp(ContextBound, Generic[Domain, Codomain]): ctx : Context Resolved backend context. + References + ---------- + .. [Conway] J. B. Conway, *A Course in Functional Analysis*, 2nd ed., + Springer, 1990, II.2.4: for ``A`` in ``B(H,K)`` there is a **unique** + ``A*`` in ``B(K,H)`` with `` = ``; existence rests on the + Riesz representation theorem (I.3.4), which is why a metric-aware adjoint + needs the Riesz maps. Theorem II.2.6 gives the algebra this class and + ``linop/_algebra.py`` implement: + ``(aA + B)* = conj(a) A* + B*``, ``(AB)* = B* A*`` (**order reverses**), + ``A** = A``. II.2.2 is the uniqueness statement behind "the" adjoint. + Conway II.3 (``def-adjoint-banach``) is the *different*, dual-space + adjoint — not what SpaceCore means by ``rapply``. + Examples -------- Use a concrete dense operator as a :class:`LinOp`. @@ -62,8 +88,14 @@ class LinOp(ContextBound, Generic[Domain, Codomain]): array([3., 8.]) """ - def __init__(self, dom: Domain, cod: Codomain, ctx: Context | str | None = None): - self.dom, self.cod = self._bind_context(ctx, dom, cod) + def __init__( + self, + dom: Domain, + cod: Codomain, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ): + self.dom, self.cod = self._bind_context(ctx, dom, cod, check_level=check_level) @property def domain(self) -> Domain: @@ -97,7 +129,16 @@ def apply(self, x: Any) -> Any: @abstractmethod def rapply(self, y: Any) -> Any: - """Apply the adjoint map to an element of ``self.codomain``.""" + r"""Apply the adjoint map to an element of ``self.codomain``. + + This is the **Hilbert-space adjoint** ``A^#``, defined by + ``_Y = _X`` (Dunford-Schwartz, *Linear Operators I*, + Definition VI.2.9) -- *not* the Banach adjoint of Definition VI.2.1, + which maps ``Y^* -> X^*`` and is the coordinate transpose. The two + coincide only when both spaces are Euclidean. See + :func:`spacecore.linop._metric.metric_rapply` for the derivation and + the full reference block. + """ def _apply_core(self, x: Any) -> Any: """Apply without adding validation beyond the concrete implementation.""" @@ -120,7 +161,11 @@ def __call__(self, x: Any) -> Any: return self.apply(x) def adjoint_apply(self, y: Any) -> Any: - """Apply the adjoint of this linear operator to ``y``.""" + """Apply the adjoint of this linear operator to ``y``. + + Spelled-out alias for :meth:`rapply`; the same Definition VI.2.9 + adjoint, not a second notion. + """ return self.rapply(y) def is_hermitian(self) -> bool | None: @@ -358,13 +403,8 @@ def _short_repr(self) -> str: """ return f"{type(self).__name__}({self._arrow()})" - @abstractmethod - def tree_flatten(self) -> tuple[tuple[Any, ...], Any]: - """Flatten this operator for backend pytree registration.""" - ... - - @classmethod - @abstractmethod - def tree_unflatten(cls, aux: Any, children: Any) -> Self: - """Rebuild this operator from backend pytree data.""" - ... + # tree_flatten / tree_unflatten are inherited from PyTreeNode, which owns the + # flatten contract for every SpaceCore container and auto-registers concrete + # subclasses with each backend's tree protocol. Every concrete LinOp is a + # container, so the capability belongs on this base rather than being + # re-declared (and separately registered) per operator. diff --git a/spacecore/linop/_dense.py b/spacecore/linop/_dense.py index cf725ff..519e4b2 100644 --- a/spacecore/linop/_dense.py +++ b/spacecore/linop/_dense.py @@ -1,10 +1,13 @@ from __future__ import annotations +import warnings + from functools import cached_property from math import prod from typing import Any, cast from ._base import Codomain, Domain, LinOp +from .._check_policy import CheckLevel from .._checks import checked_method from ._metric import _metric_is_hermitian_by_basis, _requires_euclidean_or_riesz from ..space import ( @@ -14,14 +17,13 @@ WeightedInnerProduct, ) from ..types import DenseArray -from ..backend import jax_pytree_class, Context -from .._contextual import resolve_context_priority +from ..contextual import Context +from ..contextual import resolve_context_priority from ..kernels import core_kernels from ..kernels.core.dense import _DenseMode @core_kernels("dense") -@jax_pytree_class class DenseLinOp(LinOp[Domain, Codomain]): r""" Represent a dense coordinate tensor-backed linear operator. @@ -45,9 +47,26 @@ class DenseLinOp(LinOp[Domain, Codomain]): Domain space. cod : Space or None, optional Codomain space. If omitted, it is inferred from the leading axes of - ``A``. + ``A`` — as a **Euclidean** :class:`DenseCoordinateSpace`. Only the + *shape* is inferred, never the geometry. + + .. warning:: + + On a non-Euclidean domain this is almost certainly not the operator + you meant. ``DenseLinOp(M, X)`` with square ``M`` and weighted ``X`` + builds ``X -> Y_euclidean``, **not** an endomorphism of ``X``: the + adjoint is then ``R_X^{-1} M^T`` rather than ``R_X^{-1} M^T R_X``, and + ``is_hermitian()`` reports ``False`` for a matrix that is self-adjoint + with respect to ``X``. Pass ``cod=X`` explicitly. Note the asymmetry + with :class:`~spacecore.DiagonalLinOp` and + :class:`~spacecore.IdentityLinOp`, which do reuse the given space. ctx : Context, str, or None, optional Backend context specification. Default is resolved from the spaces. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. Attributes ---------- @@ -71,6 +90,7 @@ def __init__( dom: Domain, cod: Codomain | None = None, ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: ctx = resolve_context_priority(ctx, dom, cod) ctx.assert_dense(A) # Check if A is ndarray of ctx @@ -78,10 +98,26 @@ def __init__( if cod is None: cod_shape_len = len(A.shape) - len(dom.shape) cod = cast(Codomain, DenseCoordinateSpace(tuple(A.shape[:cod_shape_len]), ctx)) + if not dom.is_euclidean: + # Only the shape is inferable from ``A``; the geometry is not. On a + # non-Euclidean domain the Euclidean default is very often not what + # the caller meant -- ``DenseLinOp(M, X)`` with square ``M`` reads as + # an endomorphism of ``X`` but is not one, which changes the adjoint + # and makes ``is_hermitian`` report False for a self-adjoint matrix. + # Silent when the domain is Euclidean, where the default is right. + warnings.warn( + "DenseLinOp inferred a EUCLIDEAN codomain from the shape of `A`, but " + "the domain geometry is non-Euclidean. The operator therefore maps " + "X -> (Euclidean space), not X -> X, so its adjoint is " + "R_X^-1 A^dagger rather than R_X^-1 A^dagger R_X. Pass `cod=` " + "explicitly to state the codomain geometry you intend.", + UserWarning, + stacklevel=2, + ) _requires_euclidean_or_riesz(dom, cod, "DenseLinOp") - super(DenseLinOp, self).__init__(dom, cod, ctx) + super(DenseLinOp, self).__init__(dom, cod, ctx, check_level=check_level) expected = tuple(self.cod.shape) + tuple(self.dom.shape) if tuple(A.shape) != expected: @@ -217,7 +253,7 @@ def is_hermitian(self) -> bool | None: def __eq__(self, other: Any) -> bool: """Return whether another dense operator has the same spaces and values.""" - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented if self.dom != other.dom or self.cod != other.cod: # Tier 2: spaces before allclose return False diff --git a/spacecore/linop/_diagonal.py b/spacecore/linop/_diagonal.py index 50fc882..d8063c6 100644 --- a/spacecore/linop/_diagonal.py +++ b/spacecore/linop/_diagonal.py @@ -6,8 +6,9 @@ from ._base import LinOp from ._metric import _metric_is_hermitian_by_basis, _requires_euclidean_or_riesz +from .._check_policy import CheckLevel from .._checks import checked_method -from ..backend import Context, jax_pytree_class +from ..contextual import Context from ..space import ( CoordinateSpace, DenseCoordinateSpace, @@ -16,13 +17,12 @@ WeightedInnerProduct, ) from ..types import DenseArray -from .._contextual import resolve_context_priority +from ..contextual import resolve_context_priority from ..kernels import core_kernels from ..kernels.core.diagonal import _DiagonalMode @core_kernels("diagonal") -@jax_pytree_class class DiagonalLinOp(LinOp[CoordinateSpace, CoordinateSpace]): r""" Represent a coordinatewise diagonal linear operator. @@ -41,6 +41,11 @@ class DiagonalLinOp(LinOp[CoordinateSpace, CoordinateSpace]): ``diagonal.shape``. ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``space``. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. Attributes ---------- @@ -63,13 +68,14 @@ def __init__( diagonal: DenseArray, space: CoordinateSpace | None = None, ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: ctx = resolve_context_priority(ctx, space) ctx.assert_dense(diagonal) if space is None: space = DenseCoordinateSpace(tuple(diagonal.shape), ctx) _requires_euclidean_or_riesz(space, space, "DiagonalLinOp") - super().__init__(space, space, ctx) + super().__init__(space, space, ctx, check_level=check_level) expected = tuple(self.domain.shape) if tuple(diagonal.shape) != expected: raise TypeError( @@ -149,7 +155,7 @@ def is_hermitian(self) -> bool | None: def __eq__(self, other: Any) -> bool: """Return whether another diagonal operator has the same space and values.""" - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented if self.domain != other.domain: # Tier 2: space before allclose return False diff --git a/spacecore/linop/_metric.py b/spacecore/linop/_metric.py index 2cf7de8..2fc8450 100644 --- a/spacecore/linop/_metric.py +++ b/spacecore/linop/_metric.py @@ -82,7 +82,19 @@ def _warn_metric_batch_fallback(opname: str, error: Exception) -> None: def metric_rapply(domain, codomain, euclidean_rapply, y): - """Apply the metric adjoint ``R_X^{-1} A^dagger R_Y`` to one element.""" + r"""Apply the metric adjoint ``R_X^{-1} A^dagger R_Y`` to one element. + + Derivation. Write ``R_X``, ``R_Y`` for the Riesz maps sending an element to + the coordinate functional that represents it, and ``A^dagger`` for the plain + coordinate (conjugate-)transpose. The defining identity + ``_Y = _X`` expands to ``(Ax)^dagger R_Y y = x^dagger R_X A^# y`` + for all ``x``, hence ``A^dagger R_Y = R_X A^#`` and + ``A^# = R_X^{-1} A^dagger R_Y``. The Euclidean short-circuit below is the case + ``R_X = R_Y = I``, where the metric adjoint collapses to the transpose. + + Getting this wrong is silent: a coordinate transpose satisfies the identity on + every Euclidean space, so only a non-Euclidean test can detect it. + """ if domain.is_euclidean and codomain.is_euclidean: return euclidean_rapply(y) return domain.riesz_inverse(euclidean_rapply(codomain.riesz(y))) diff --git a/spacecore/linop/_sparse.py b/spacecore/linop/_sparse.py index 62cdbe9..c12b7e5 100644 --- a/spacecore/linop/_sparse.py +++ b/spacecore/linop/_sparse.py @@ -6,6 +6,7 @@ from ._base import LinOp from ._metric import _metric_is_hermitian_by_basis, _requires_euclidean_or_riesz +from .._check_policy import CheckLevel from .._checks import checked_method from ..space import ( CoordinateSpace, @@ -15,8 +16,8 @@ WeightedInnerProduct, ) from ..types import DenseArray, SparseArray -from ..backend import jax_pytree_class, Context -from .._contextual import resolve_context_priority +from ..contextual import Context +from ..contextual import resolve_context_priority from ..kernels import core_kernels from ..kernels.core.sparse import _SparseMode @@ -29,7 +30,6 @@ @core_kernels("sparse") -@jax_pytree_class class SparseLinOp(LinOp[CoordinateSpace, CoordinateSpace]): r""" Represent a sparse coordinate matrix-backed linear operator. @@ -56,6 +56,11 @@ class SparseLinOp(LinOp[CoordinateSpace, CoordinateSpace]): Codomain vector space, or a subclass of ``VectorSpace``. ctx : Context, str, or None, optional Backend context specification. Default is resolved from the spaces. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. Attributes ---------- @@ -82,6 +87,7 @@ def __init__( dom: CoordinateSpace, cod: CoordinateSpace, ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: ctx = resolve_context_priority(ctx, dom, cod) ctx.assert_sparse(A) # Check if A is sparse array of ctx @@ -90,7 +96,7 @@ def __init__( _requires_euclidean_or_riesz(dom, cod, "SparseLinOp") - super(SparseLinOp, self).__init__(dom, cod, ctx) + super(SparseLinOp, self).__init__(dom, cod, ctx, check_level=check_level) expected = (prod(self.cod.shape), prod(self.dom.shape)) if tuple(A.shape) != expected: @@ -252,7 +258,7 @@ def is_hermitian(self) -> bool | None: return None def __eq__(self, other: Any) -> bool: - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented if self.dom != other.dom or self.cod != other.cod: # Tier 2: spaces before allclose return False diff --git a/spacecore/linop/tree/_base.py b/spacecore/linop/tree/_base.py index 53c7b63..3da43f7 100644 --- a/spacecore/linop/tree/_base.py +++ b/spacecore/linop/tree/_base.py @@ -4,10 +4,10 @@ from typing import Tuple, Sequence, Any, Self from .._base import LinOp, Domain, Codomain -from ...backend import jax_pytree_class, Context +from ...contextual import Context +from ..._check_policy import CheckLevel -@jax_pytree_class class TreeLinOp(LinOp[Domain, Codomain]): """ Define a base class for operators assembled from component operators. @@ -23,17 +23,26 @@ class TreeLinOp(LinOp[Domain, Codomain]): ctx : Context, str, or None, optional Backend context specification. Default is resolved from ``dom`` and ``cod``. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this operator. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Validation + policy is a property of the operator, not of the ``Context``. """ parts: Tuple[LinOp, ...] def __init__( - self, dom: Domain, cod: Codomain, parts: Sequence[LinOp], ctx: Context | str | None = None + self, + dom: Domain, + cod: Codomain, + parts: Sequence[LinOp], + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: if not parts: raise ValueError("Parts must be non-empty.") - super().__init__(dom, cod, ctx) + super().__init__(dom, cod, ctx, check_level=check_level) self.parts = tuple(op.convert(self.ctx) for op in parts) self._num_parts = len(self.parts) @@ -68,7 +77,7 @@ def _endpoint_treedefs(self) -> tuple[Any, Any]: def __eq__(self, other: Any) -> bool: """Return whether another tree operator has the same layout.""" - if not self._eq_backend_compatible(other): # Tier 1: backend + if not self.same_math(other): # Tier 1: backend return NotImplemented if self.dom != other.dom or self.cod != other.cod: # Tier 2a: endpoint spaces return False diff --git a/spacecore/linop/tree/_block.py b/spacecore/linop/tree/_block.py index cbe07a7..a765054 100644 --- a/spacecore/linop/tree/_block.py +++ b/spacecore/linop/tree/_block.py @@ -9,8 +9,8 @@ from .._algebra import _same_space_for_algebra from .._base import LinOp from ..._checks import checked_method -from ..._contextual._bound import _same_math_context -from ...backend import Context, jax_pytree_class +from ...contextual import Context +from ..._check_policy import CheckLevel from ...kernels import CachedStackParts, dispatch, should_consult_dispatch from ...space import TreeSpace @@ -61,7 +61,7 @@ def _validate_blocks(blocks: Sequence[Any], owner: str) -> tuple[LinOp, ...]: first = validated[0] for index, block in enumerate(validated[1:], start=1): - if not _same_math_context(first.ctx, block.ctx): + if not first.ctx.same_math(block.ctx): raise ValueError( f"All {owner} blocks must have the same mathematical context; " f"block 0 has {first.ctx!r}, block {index} has {block.ctx!r}." @@ -85,7 +85,6 @@ def _sum_values(space: Any, values: Sequence[Any], *, batched: bool) -> Any: return result -@jax_pytree_class class BlockDiagonalLinOp(TreeLinOp[TreeSpace, TreeSpace]): r""" Represent independent blocks over a finite direct-product tree. @@ -108,6 +107,10 @@ class BlockDiagonalLinOp(TreeLinOp[TreeSpace, TreeSpace]): Block operators for the legacy form; inferred from ``blocks`` otherwise. ctx : Context, str, or None, optional Backend context specification. Default is resolved from the blocks. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this operator. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Validation + policy is a property of the operator, not of the ``Context``. Notes ----- @@ -122,6 +125,7 @@ def __init__( cod: TreeSpace | None = None, parts: Sequence[LinOp] | None = None, ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: if isinstance(blocks, TreeSpace): dom = blocks @@ -141,7 +145,7 @@ def __init__( dom = TreeSpace(treedef, tuple(block.domain for block in block_parts), ctx=ctx) cod = TreeSpace(treedef, tuple(block.codomain for block in block_parts), ctx=ctx) - super().__init__(dom, cod, block_parts, ctx) + super().__init__(dom, cod, block_parts, ctx, check_level=check_level) # ADR-022: carry the per-accessor stacked-block-matrix memo on the parts # so the uniform-dense batched fold (block_batched) stacks once and reuses # across applies. Built lazily on first optimized use, NumPy-only; dropped @@ -168,13 +172,13 @@ def apply(self, x: Any) -> Any: def _apply_unchecked(self, x: Any) -> Any: x_parts = self.dom._components(x) - if should_consult_dispatch(self.ctx): + if should_consult_dispatch(self): y_parts = dispatch( _BLOCK_DIAGONAL_APPLY_KEY, self.parts, x_parts, generic=_block_diagonal_apply, - ctx=self.ctx, + ctx=self, ) else: y_parts = _block_diagonal_apply(self.parts, x_parts) @@ -187,13 +191,13 @@ def rapply(self, y: Any) -> Any: def _rapply_unchecked(self, y: Any) -> Any: y_parts = self.cod._components(y) - if should_consult_dispatch(self.ctx): + if should_consult_dispatch(self): x_parts = dispatch( _BLOCK_DIAGONAL_RAPPLY_KEY, self.parts, y_parts, generic=_block_diagonal_rapply, - ctx=self.ctx, + ctx=self, ) else: x_parts = _block_diagonal_rapply(self.parts, y_parts) @@ -208,13 +212,13 @@ def vapply(self, x: Any) -> Any: def _vapply_unchecked(self, x: Any) -> Any: x_parts = self.dom._components(x) - if should_consult_dispatch(self.ctx): + if should_consult_dispatch(self): y_parts = dispatch( _BLOCK_DIAGONAL_VAPPLY_KEY, self.parts, x_parts, generic=_block_diagonal_vapply, - ctx=self.ctx, + ctx=self, ) else: y_parts = _block_diagonal_vapply(self.parts, x_parts) @@ -229,13 +233,13 @@ def rvapply(self, y: Any) -> Any: def _rvapply_unchecked(self, y: Any) -> Any: y_parts = self.cod._components(y) - if should_consult_dispatch(self.ctx): + if should_consult_dispatch(self): x_parts = dispatch( _BLOCK_DIAGONAL_RVAPPLY_KEY, self.parts, y_parts, generic=_block_diagonal_rvapply, - ctx=self.ctx, + ctx=self, ) else: x_parts = _block_diagonal_rvapply(self.parts, y_parts) @@ -280,7 +284,6 @@ def _convert(self, new_ctx: Context) -> BlockDiagonalLinOp: ) -@jax_pytree_class class BlockMatrixLinOp(TreeLinOp[TreeSpace, TreeSpace]): r""" Represent a rectangular matrix of blocks over direct products. @@ -296,9 +299,17 @@ class BlockMatrixLinOp(TreeLinOp[TreeSpace, TreeSpace]): Nonempty rectangular block matrix. Blocks in one row must have compatible codomains, and blocks in one column must have compatible domains. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this operator. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Validation + policy is a property of the operator, not of the ``Context``. """ - def __init__(self, block_rows: Sequence[Sequence[LinOp]]) -> None: + def __init__( + self, + block_rows: Sequence[Sequence[LinOp]], + check_level: CheckLevel | bool | None = None, + ) -> None: if not isinstance(block_rows, Sequence) or isinstance(block_rows, (str, bytes)): raise TypeError("BlockMatrixLinOp block_rows must be a sequence of rows.") rows = tuple(block_rows) @@ -354,7 +365,7 @@ def __init__(self, block_rows: Sequence[Sequence[LinOp]]) -> None: cod = TreeSpace.from_leaf_spaces(tuple(row[0].codomain for row in normalized_rows), ctx) self._row_count = len(normalized_rows) self._column_count = column_count - super().__init__(dom, cod, flat_blocks, ctx) + super().__init__(dom, cod, flat_blocks, ctx, check_level=check_level) self.block_rows = tuple( self.parts[index * column_count : (index + 1) * column_count] for index in range(self._row_count) diff --git a/spacecore/linop/tree/_from_single.py b/spacecore/linop/tree/_from_single.py index b0b4e8d..062efbf 100644 --- a/spacecore/linop/tree/_from_single.py +++ b/spacecore/linop/tree/_from_single.py @@ -7,7 +7,8 @@ from ..._checks import checked_method from ...kernels import CachedStackParts, dispatch, should_consult_dispatch from ...space import DenseCoordinateSpace, DenseVectorSpace, ElementwiseJordanSpace, TreeSpace -from ...backend import jax_pytree_class, Context +from ...contextual import Context +from ..._check_policy import CheckLevel # ADR-016 dispatch call site: a StackedLinOp applies one shared input through # every component. The per-component loop is the ``generic`` fallback; the @@ -22,7 +23,6 @@ def _stacked_apply(parts: Any, x: Any) -> tuple[Any, ...]: return tuple(p._apply_core(x) for p in parts) -@jax_pytree_class class StackedLinOp(TreeLinOp[Domain, TreeSpace]): r""" Represent operators from one domain as a tree-valued map. @@ -41,6 +41,10 @@ class StackedLinOp(TreeLinOp[Domain, TreeSpace]): Operators from ``dom`` to each component of ``cod``. ctx : Context, str, or None, optional Backend context specification. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this operator. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Validation + policy is a property of the operator, not of the ``Context``. """ def __init__( @@ -49,8 +53,9 @@ def __init__( cod: TreeSpace, parts: Sequence[LinOp], ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: - super().__init__(dom, cod, parts, ctx) + super().__init__(dom, cod, parts, ctx, check_level=check_level) # ADR-022: memoize the stacked component matrices for the stacked.apply # broadcast fold (built once on first optimized use, NumPy-only). self.parts = CachedStackParts(self.parts) @@ -98,13 +103,13 @@ def apply(self, x: Any) -> Any: def _apply_unchecked(self, x: Any) -> Any: """Apply component operators without checks and rebuild codomain representation.""" - if should_consult_dispatch(self.ctx): + if should_consult_dispatch(self): y_parts = dispatch( _STACKED_APPLY_KEY, self.parts, x, generic=_stacked_apply, - ctx=self.ctx, + ctx=self, ) elif self._num_parts == 2: y_parts = (self._apply_parts[0](x), self._apply_parts[1](x)) diff --git a/spacecore/linop/tree/_to_single.py b/spacecore/linop/tree/_to_single.py index 76e71d0..7bcceda 100644 --- a/spacecore/linop/tree/_to_single.py +++ b/spacecore/linop/tree/_to_single.py @@ -7,7 +7,8 @@ from ..._checks import checked_method from ...kernels import CachedStackParts, dispatch, should_consult_dispatch from ...space import DenseCoordinateSpace, DenseVectorSpace, ElementwiseJordanSpace, TreeSpace -from ...backend import jax_pytree_class, Context +from ...contextual import Context +from ..._check_policy import CheckLevel # ADR-016 dispatch call site: a SumToSingleLinOp applies one shared input through # every component adjoint. The per-component loop is the ``generic`` fallback; @@ -22,7 +23,6 @@ def _sum_to_single_rapply(parts: Any, y: Any) -> tuple[Any, ...]: return tuple(p._rapply_core(y) for p in parts) -@jax_pytree_class class SumToSingleLinOp(TreeLinOp[TreeSpace, Codomain]): r""" Represent a sum of leaf operators from a tree domain. @@ -41,6 +41,10 @@ class SumToSingleLinOp(TreeLinOp[TreeSpace, Codomain]): Operators from each product component to ``cod``. ctx : Context, str, or None, optional Backend context specification. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this operator. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Validation + policy is a property of the operator, not of the ``Context``. """ def __init__( @@ -49,8 +53,9 @@ def __init__( cod: Codomain, parts: Sequence[LinOp], ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: - super().__init__(dom, cod, parts, ctx) + super().__init__(dom, cod, parts, ctx, check_level=check_level) # ADR-022: memoize the stacked adjoint matrices for the sum_to_single.rapply # broadcast fold (built once on first optimized use, NumPy-only). self.parts = CachedStackParts(self.parts) @@ -131,13 +136,13 @@ def rapply(self, y: Any) -> Any: def _rapply_unchecked(self, y: Any) -> Any: """Apply component adjoints without checks and rebuild domain representation.""" - if should_consult_dispatch(self.ctx): + if should_consult_dispatch(self): x_parts = dispatch( _SUM_TO_SINGLE_RAPPLY_KEY, self.parts, y, generic=_sum_to_single_rapply, - ctx=self.ctx, + ctx=self, ) elif self._num_parts == 2: x_parts = (self._rapply_parts[0](y), self._rapply_parts[1](y)) diff --git a/spacecore/opfamily.py b/spacecore/opfamily.py new file mode 100644 index 0000000..4383cab --- /dev/null +++ b/spacecore/opfamily.py @@ -0,0 +1,329 @@ +r"""Point-indexed families of linear operators, and the functional-weighted map. + +A :class:`~spacecore.linop.LinOp` is a *single* linear map. Several constructions +that look like operators are really a **family** of them, one per point of the +domain: the object is linear once a point is fixed, but the point-to-operator +assignment is not. + +The motivating case is the functional-weighted map + +.. math:: + + m(x) = F(x) \, A x, + +with ``F`` a scalar-valued :class:`~spacecore.functional.Functional` and ``A`` a +linear operator. This is **not** linear in ``x`` — both the scale ``F(x)`` and the +direction ``A x`` move with ``x`` — so it cannot be a ``LinOp`` without breaking +the contract that ``rapply``, ``compose_chain``, and the kernel fusion rules rely +on. It *is* linear the moment ``x`` is frozen, and that is exactly what +:meth:`OperatorFamily.at` returns. + +Two different linear operators are attached to each point, and conflating them is +the easy mistake: + +* :meth:`OperatorFamily.at` — the **frozen member** :math:`A_x`. For the + functional-weighted family that is ``F(x) · A``, so ``at(x).apply(x) == m(x)``. +* :meth:`OperatorFamily.linearize_at` — the **derivative** :math:`Dm(x)`, the + linear map ``h -> Dm(x)[h]`` used by Newton/Gauss-Newton. These coincide only + when the family is constant (i.e. ``F`` is constant). + +This module deliberately sits outside both ``linop`` and ``functional``: it +depends on both, and neither depends on it. +""" +from __future__ import annotations + +from abc import abstractmethod +from typing import Any, Generic, TypeVar + +from ._checks import checked_method +from ._check_policy import CheckLevel, minimum_check_level +from .backend import PyTreeNode +from .contextual import Context, ContextBound +from .functional import Functional +from .linop import LinOp, MatrixFreeLinOp +from .linop._algebra import make_scaled +from .space import Space + +Domain = TypeVar("Domain", bound=Space) +Codomain = TypeVar("Codomain", bound=Space) + + +class OperatorFamily(PyTreeNode, ContextBound, Generic[Domain, Codomain]): + r""" + A point-indexed family of linear operators :math:`x \mapsto A_x \in L(X, Y)`. + + Not a :class:`~spacecore.linop.LinOp`, and deliberately not a subclass of one: + the map :meth:`apply` is non-linear in general. A ``LinOp`` is the special + case of a *constant* family, and every member of a family is a genuine + ``LinOp`` obtained with :meth:`at`. + + Subclasses implement :meth:`at`. Everything else is derived from it, though + subclasses may override :meth:`apply` with a fused evaluation that avoids + building the intermediate node. + + Parameters + ---------- + dom : Space + Domain space ``X`` — both the index set of the family and the domain of + each member. + cod : Space + Codomain space ``Y`` of each member. + ctx : Context, str, or None, optional + Backend context specification. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. + """ + + def __init__( + self, + dom: Domain, + cod: Codomain, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: + self.dom, self.cod = self._bind_context(ctx, dom, cod, check_level=check_level) + + @property + def domain(self) -> Domain: + """Domain space of every member of this family.""" + return self.dom + + @property + def codomain(self) -> Codomain: + """Codomain space of every member of this family.""" + return self.cod + + @abstractmethod + def at(self, x: Any) -> LinOp: + """Return the family member at ``x`` — a genuine :class:`LinOp`. + + Freezing the index point is what recovers linearity, so the result + supports the whole operator algebra: ``@``, ``+``, ``.H``, ``fuse()``. + """ + + def apply(self, x: Any) -> Any: + """Evaluate the map ``x -> A_x x`` at ``x``. + + Note the double role of ``x``: it selects the member *and* is the vector + the member is applied to. That coupling is precisely what makes this + non-linear. To apply one member to a *different* vector, use + ``family.at(x).apply(h)``. + """ + return self.at(x).apply(x) + + def __call__(self, x: Any) -> Any: + """Evaluate the map at ``x``.""" + return self.apply(x) + + def linearize_at(self, x: Any) -> LinOp: + r"""Return the derivative :math:`Dm(x)` as a :class:`LinOp`. + + The Fréchet derivative of ``m(x) = A_x x`` at a fixed ``x``, i.e. the + linear map ``h -> Dm(x)[h]``. This is **not** :meth:`at`: ``at(x)`` + ignores how the family varies, while ``linearize_at(x)`` accounts for it. + They agree exactly when the family is constant. + + Override in subclasses that can differentiate; the base raises + :class:`NotImplementedError`. + """ + raise NotImplementedError( + f"{type(self).__name__} does not implement linearize_at()." + ) + + def _arrow(self) -> str: + """Return the ``domain ⇝ codomain`` descriptor (``⇝``: not linear).""" + from ._repr import describe_space + + return f"{describe_space(self.dom)} ⇝ {describe_space(self.cod)}" + + def _repr_body(self) -> str: + return self._arrow() + + def _short_repr(self) -> str: + """Return a bounded ``ClassName(domain ⇝ codomain)`` form for nesting.""" + return f"{type(self).__name__}({self._arrow()})" + + +class FunctionalScaledOperator(OperatorFamily[Domain, Codomain]): + r""" + The functional-weighted map :math:`m(x) = F(x)\, A x`. + + ``F`` and ``A`` must share a domain. The family member at ``x`` is the + ordinary scalar multiple ``F(x) · A``, built through + :func:`~spacecore.linop.make_scaled`, so it folds and canonicalizes like any + other scaled operator. + + Parameters + ---------- + functional : Functional + Scalar weight ``F``, defined on ``op.domain``. + op : LinOp + Linear operator ``A``. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy. When omitted, the least-strict of the two + operands is used, matching the functional algebra's combinator rule. + + Attributes + ---------- + functional : Functional + The scalar weight. + op : LinOp + The scaled operator. + """ + + def __init__( + self, + functional: Functional, + op: LinOp, + check_level: CheckLevel | bool | None = None, + ) -> None: + if not isinstance(functional, Functional): + raise TypeError( + f"functional must be a Functional, got {type(functional).__name__}." + ) + if not isinstance(op, LinOp): + raise TypeError(f"op must be a LinOp, got {type(op).__name__}.") + if functional.domain != op.domain: + raise ValueError( + "FunctionalScaledOperator requires functional.domain == op.domain; " + f"got {functional.domain!r} and {op.domain!r}." + ) + if check_level is None: + check_level = minimum_check_level((functional.check_level, op.check_level)) + super().__init__(op.domain, op.codomain, op.ctx, check_level=check_level) + self.functional = functional.convert(self.ctx) + self.op = op.convert(self.ctx) + + @checked_method(in_space="domain") + def at(self, x: Any) -> LinOp: + """Return the frozen member ``F(x) · A``.""" + return make_scaled(self.functional.value(x), self.op) + + @checked_method(in_space="domain", out_space="codomain") + def apply(self, x: Any) -> Any: + """Return ``F(x) · (A x)`` without building the intermediate node.""" + return self.codomain.scale(self.functional.value(x), self.op.apply(x)) + + @checked_method(in_space="domain") + def linearize_at(self, x: Any) -> LinOp: + r"""Return the derivative of ``m(x) = F(x) A x`` at ``x``. + + By the product rule, for a direction ``h`` + + .. math:: + + Dm(x)[h] = DF(x)[h] \; A x \;+\; F(x)\, A h + = \langle \nabla F(x), h\rangle \, Ax + F(x)\, A h, + + using the library's Riesz convention `` = DF(x)[h]``. The + first term is a **rank-one** operator (its output is always a multiple of + the fixed vector ``Ax``), the second the frozen member — so the + derivative and :meth:`at` differ by exactly that rank-one correction, and + coincide when ``F`` is constant. + + The adjoint follows from + :math:`\langle Dm(x)h, w\rangle_Y = \langle h, Dm(x)^{\#}w\rangle_X`: + + .. math:: + + Dm(x)^{\#}[w] = \langle Ax, w\rangle_Y \, \nabla F(x) + + \overline{F(x)}\, A^{\#} w . + + Both coefficients are inner products in the *codomain* geometry, not + coordinate dot products, so this is the metric adjoint (ADR-009). + """ + X, Y = self.domain, self.codomain + value, gradient = self.functional.value_and_grad(x) + ax = self.op.apply(x) + conj_value = self.ops.conj(value) + + def _apply(h: Any) -> Any: + return Y.add( + Y.scale(X.inner(gradient, h), ax), + Y.scale(value, self.op.apply(h)), + ) + + def _rapply(w: Any) -> Any: + return X.add( + X.scale(Y.inner(ax, w), gradient), + X.scale(conj_value, self.op.rapply(w)), + ) + + # ``_rapply`` is already the *metric* adjoint: it pairs with ``Y.inner`` + # and defers to ``self.op.rapply`` (itself metric-aware). Stating that + # explicitly documents the choice and silences the construction advisory. + return MatrixFreeLinOp( + _apply, _rapply, X, Y, self.ctx, euclidean_adjoint=False + ) + + def __eq__(self, other: Any) -> bool: + """Return whether another family has the same weight and operator.""" + if not self.same_math(other): + return NotImplemented + return self.functional == other.functional and self.op == other.op + + def tree_flatten(self): + """Flatten this family for pytree registration.""" + return (self.functional, self.op), () + + @classmethod + def tree_unflatten(cls, aux, children): + """Rebuild this family from pytree data.""" + functional, op = children + return cls(functional, op) + + def _convert(self, new_ctx: Context) -> "FunctionalScaledOperator": + """Convert both operands to ``new_ctx``.""" + return FunctionalScaledOperator( + self.functional.convert(new_ctx), self.op.convert(new_ctx) + ) + + def _repr_body(self) -> str: + return f"{self.functional._short_repr()} · {self.op._short_repr()}, {self._arrow()}" + + +def make_functional_scaled_operator(functional: Functional, op: LinOp) -> Any: + """ + Return ``F · A`` as a :class:`LinOp` when possible, else as a family. + + A :class:`~spacecore.functional.ConstantFunctional` weight is a constant + family, so the result collapses to an ordinary + :class:`~spacecore.linop.ScaledLinOp` — the linear case is not forced through + the non-linear type. Everything else builds a + :class:`FunctionalScaledOperator`. + + As in the functional algebra, the collapse is **structural**: a functional + that merely happens to be constant is not recognized, because that is a fact + about values, not about the expression. + + Parameters + ---------- + functional : Functional + Scalar weight ``F``. + op : LinOp + Linear operator ``A``. + + Returns + ------- + LinOp or FunctionalScaledOperator + ``ScaledLinOp`` for a constant weight, otherwise the family. + """ + from .functional import ConstantFunctional, ZeroFunctional + + if not isinstance(functional, Functional): + raise TypeError( + f"functional must be a Functional, got {type(functional).__name__}." + ) + if not isinstance(op, LinOp): + raise TypeError(f"op must be a LinOp, got {type(op).__name__}.") + if functional.domain != op.domain: + raise ValueError( + "F · A requires functional.domain == op.domain; " + f"got {functional.domain!r} and {op.domain!r}." + ) + if isinstance(functional, ZeroFunctional): + return make_scaled(0, op) + if isinstance(functional, ConstantFunctional): + return make_scaled(functional.constant, op) + return FunctionalScaledOperator(functional, op) diff --git a/spacecore/optimize/_optax.py b/spacecore/optimize/_optax.py index 036f007..642bd3e 100644 --- a/spacecore/optimize/_optax.py +++ b/spacecore/optimize/_optax.py @@ -287,8 +287,8 @@ def minimize_optax( import optax import spacecore as sc - ctx = sc.Context(sc.JaxOps(), dtype=np.float32, check_level="none") - X = sc.DenseCoordinateSpace((2,), ctx) + ctx = sc.Context(sc.JaxOps(), dtype=np.float32) + X = sc.DenseCoordinateSpace((2,), ctx, check_level="none") Q = sc.DenseLinOp(ctx.asarray([[3.0, 0.0], [0.0, 1.0]]), X, X, ctx) linear = sc.InnerProductFunctional(ctx.asarray([-3.0, -2.0]), X) F = sc.LinOpQuadraticForm(Q, linear) diff --git a/spacecore/space/base/_coordinate.py b/spacecore/space/base/_coordinate.py index 30e6b8a..9f13d22 100644 --- a/spacecore/space/base/_coordinate.py +++ b/spacecore/space/base/_coordinate.py @@ -4,7 +4,8 @@ from math import prod from typing import Any, Tuple -from ...backend import Context +from ..._check_policy import CheckLevel +from ...contextual import Context from ..._repr import shape_descriptor from ...types import DenseArray from ._vector import VectorSpace @@ -20,12 +21,22 @@ class CoordinateSpace(VectorSpace): Canonical coordinate shape for one element of the space. ctx : Context, str, or None, optional Context specification used for coordinate arrays. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. """ shape: Tuple[int, ...] - def __init__(self, shape: Tuple[int, ...], ctx: Context | str | None = None) -> None: - super().__init__(ctx) + def __init__( + self, + shape: Tuple[int, ...], + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: + super().__init__(ctx, check_level=check_level) self.shape = tuple(shape) def _eq_algebra(self, other: Any) -> bool: @@ -47,11 +58,19 @@ def size(self) -> int: @abstractmethod def flatten(self, x: Any) -> DenseArray: - """Return a dense one-dimensional coordinate vector.""" + """Return a dense one-dimensional coordinate vector, in **C order**. + + C order (row-major, last axis varying fastest) is SpaceCore's declared + vec convention on every backend -- see the flattening-order note in + :mod:`spacecore.backend._ops`. It is what makes ``to_matrix`` and + ``flatten`` compose: ``to_matrix() @ flatten(x) == flatten(A.apply(x))`` + holds only for a consistent choice, and the Kronecker identities in + ``kernels/`` assume this one. + """ @abstractmethod def unflatten(self, v: DenseArray) -> Any: - """Inverse of flatten.""" + """Inverse of flatten, in the same C order.""" def flatten_batch(self, xs: Any) -> DenseArray: """Flatten a leading-axis batch of space elements to shape ``(N, size)``.""" diff --git a/spacecore/space/base/_inner_product.py b/spacecore/space/base/_inner_product.py index 5d58477..d442938 100644 --- a/spacecore/space/base/_inner_product.py +++ b/spacecore/space/base/_inner_product.py @@ -151,6 +151,9 @@ class InnerProductSpace(VectorSpace): ---------- ctx : Context, str, or None, optional Context specification used for elements and validation checks. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. """ geometry: InnerProduct diff --git a/spacecore/space/base/_jordan.py b/spacecore/space/base/_jordan.py index 3783167..41a183e 100644 --- a/spacecore/space/base/_jordan.py +++ b/spacecore/space/base/_jordan.py @@ -15,6 +15,9 @@ class JordanAlgebraSpace(VectorSpace): ---------- ctx : Context, str, or None, optional Context specification used for elements and validation checks. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. """ @abstractmethod @@ -59,4 +62,7 @@ class EuclideanJordanAlgebraSpace(JordanAlgebraSpace, InnerProductSpace): ---------- ctx : Context, str, or None, optional Context specification used for elements and validation checks. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. """ diff --git a/spacecore/space/base/_space.py b/spacecore/space/base/_space.py index 417cfc2..d65e8b0 100644 --- a/spacecore/space/base/_space.py +++ b/spacecore/space/base/_space.py @@ -3,9 +3,10 @@ from typing import Any, ClassVar, Literal from ..._check_policy import CheckLevel, check_level_at_least, normalize_check_level -from ..._contextual import ContextBound +from ..._lazy_algebra import is_recognizably_nonreal +from ...contextual import ContextBound from ..._repr import field_symbol -from ...backend import Context +from ...contextual import Context from ..checks import SpaceCheck, SpaceValidationError, _run_checks @@ -17,12 +18,21 @@ class Space(ContextBound): ---------- ctx : Context, str, or None, optional Context specification used for elements and validation checks. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. """ checks: ClassVar[tuple[SpaceCheck, ...]] = () - def __init__(self, ctx: Context | str | None = None) -> None: - super().__init__(ctx) + def __init__( + self, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: + super().__init__(ctx, check_level=check_level) # Lazy caches. Populated on first call to member_checks() / # _checks_for_level() and reused for every subsequent membership # validation. Spaces are immutable after construction (`convert` @@ -31,19 +41,65 @@ def __init__(self, ctx: Context | str | None = None) -> None: self._cached_member_checks: tuple[SpaceCheck, ...] | None = None self._cached_checks_by_level: dict[CheckLevel, tuple[SpaceCheck, ...]] = {} + #: Scalar field this space is closed under, when that is *narrower* than the + #: field implied by the dtype. ``None`` (the default) means "same as + #: :attr:`field`" — the common case, where entries and scalars agree. + #: Subclasses override it with a literal; see :attr:`scalar_field`. + declared_scalar_field: ClassVar[Literal["real", "complex"] | None] = None + @property def field(self) -> Literal["real", "complex"]: - """Return the mathematical scalar field derived from the context dtype. + """Return the field the *entries* live in, derived from the context dtype. ``Context.dtype`` controls array representation. This property records - only whether the space is over the real or complex scalar field. + only whether those entries are real or complex. + + This is a fact about representation, and it is **not** always the field + of scalars the space is a vector space over — see :attr:`scalar_field`. + Equality and repr use this one: a real-dtype and a complex-dtype space + of the same shape are different spaces regardless of what they are + closed under. """ return "complex" if self.ops.is_complex_dtype(self.dtype) else "real" + @property + def scalar_field(self) -> Literal["real", "complex"]: + """Return the field of scalars this space is closed under. + + Defaults to :attr:`field` — entries and scalars agree for coordinate + spaces. A subclass whose structure survives only real scaling declares + otherwise via :attr:`declared_scalar_field`: the complex Hermitian + matrices have complex entries but form a **real** vector space, because + ``i·H`` is anti-Hermitian for Hermitian ``H``. + + Declared, never inferred. The dtype cannot imply this — it is a property + of the structure the space enforces, not of how elements are stored. + """ + return self.declared_scalar_field or self.field + + def _check_scalar(self, a: Any) -> None: + """Raise if ``a`` is not an admissible multiplier for this space.""" + if self.scalar_field == "real" and is_recognizably_nonreal(a): + raise SpaceValidationError( + f"{self._space_descriptor()} is a vector space over ℝ, so scaling by " + f"the non-real scalar {a!r} leaves the space. Multiply the underlying " + "array directly if you intend to exit this space." + ) + + def check_scalar(self, a: Any) -> None: + """Raise if ``a`` is not an admissible multiplier, unless checks are off. + + Only rejects a multiplier that is *provably* inadmissible, so a traced + scalar under ``jax.jit`` always passes (see + :func:`spacecore._lazy_algebra.is_recognizably_nonreal`). + """ + if self.check_level != "none": + self._check_scalar(a) + def __eq__(self, other: Any) -> bool: # Tier 1: backend compatibility (type + ops family + dtype, ignoring # check_level). Tier 2/3: per-subclass algebraic comparison. - if not self._eq_backend_compatible(other): + if not self.same_math(other): return NotImplemented return self._eq_algebra(other) diff --git a/spacecore/space/base/_star.py b/spacecore/space/base/_star.py index 9824959..01bf349 100644 --- a/spacecore/space/base/_star.py +++ b/spacecore/space/base/_star.py @@ -14,6 +14,9 @@ class StarSpace(Space): ---------- ctx : Context, str, or None, optional Context specification used for elements and validation checks. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. """ @abstractmethod diff --git a/spacecore/space/base/_vector.py b/spacecore/space/base/_vector.py index 75ece31..a7e8fb6 100644 --- a/spacecore/space/base/_vector.py +++ b/spacecore/space/base/_vector.py @@ -14,6 +14,9 @@ class VectorSpace(Space): ---------- ctx : Context, str, or None, optional Context specification used for elements and validation checks. + check_level : {{"none", "cheap", "standard", "strict"}}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. """ @abstractmethod diff --git a/spacecore/space/concrete/_dense_coordinate.py b/spacecore/space/concrete/_dense_coordinate.py index 499e40c..f1341e3 100644 --- a/spacecore/space/concrete/_dense_coordinate.py +++ b/spacecore/space/concrete/_dense_coordinate.py @@ -5,8 +5,9 @@ from ..base import CoordinateSpace, EuclideanInnerProduct, InnerProduct, InnerProductSpace from ..checks import BackendCheck, DTypeCheck, FieldCheck, ShapeCheck, SpaceCheck +from ..._check_policy import CheckLevel from ..._checks import checked_method -from ...backend import Context +from ...contextual import Context from ...types import DenseArray @@ -23,6 +24,11 @@ class DenseCoordinateSpace(CoordinateSpace, InnerProductSpace): geometry : InnerProduct or None, optional Inner-product geometry. If omitted, Euclidean coordinate geometry is used. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. """ def __init__( @@ -30,8 +36,9 @@ def __init__( shape: Tuple[int, ...], ctx: Context | str | None = None, geometry: InnerProduct | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: - super().__init__(tuple(shape), ctx) + super().__init__(tuple(shape), ctx, check_level=check_level) self.geometry: InnerProduct = geometry if geometry is not None else EuclideanInnerProduct() self.geometry.validate_for(self) self._size = prod(self.shape) diff --git a/spacecore/space/concrete/_dense_vector.py b/spacecore/space/concrete/_dense_vector.py index 8596d6a..4f76207 100644 --- a/spacecore/space/concrete/_dense_vector.py +++ b/spacecore/space/concrete/_dense_vector.py @@ -9,10 +9,11 @@ JordanAlgebraSpace, StarSpace, ) -from ..._check_policy import require_mutually_exclusive +from ..._check_policy import CheckLevel, require_mutually_exclusive from ..._checks import checked_method -from ..._contextual import normalize_context -from ...backend import Context, jax_pytree_class +from ...contextual import normalize_context +from ...contextual import Context +from ...backend import PyTreeNode from ...types import DenseArray from ._dense_coordinate import DenseCoordinateSpace @@ -75,6 +76,11 @@ class DenseVectorSpace(DenseCoordinateSpace, StarSpace): geometry : InnerProduct or None, optional Inner-product geometry. If omitted, Euclidean coordinate geometry is used. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. """ def __init__( @@ -82,11 +88,12 @@ def __init__( shape: Tuple[int, ...], ctx: Context | str | None = None, geometry: InnerProduct | None = None, + check_level: CheckLevel | bool | None = None, ) -> None: shape = tuple(shape) if len(shape) != 1: raise ValueError(f"DenseVectorSpace requires one-dimensional shape, got {shape}.") - super().__init__(shape, ctx, geometry=geometry) + super().__init__(shape, ctx, geometry=geometry, check_level=check_level) @checked_method(in_space="self") def star(self, x: DenseArray) -> DenseArray: @@ -97,8 +104,7 @@ def _convert(self, new_ctx: Context) -> DenseVectorSpace: return DenseVectorSpace(self.shape, new_ctx, geometry=self.geometry.convert(new_ctx)) -@jax_pytree_class -class ElementwiseJordanSpace(JordanAlgebraSpace, DenseCoordinateSpace, StarSpace): +class ElementwiseJordanSpace(PyTreeNode, JordanAlgebraSpace, DenseCoordinateSpace, StarSpace): """ Elementwise Jordan algebra for real or complex dense coordinates. @@ -111,6 +117,11 @@ class ElementwiseJordanSpace(JordanAlgebraSpace, DenseCoordinateSpace, StarSpace geometry : InnerProduct or None, optional Inner-product geometry. If omitted, Euclidean coordinate geometry is used. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. inner_product : InnerProduct or None, optional Alias for ``geometry``. """ @@ -120,6 +131,7 @@ def __new__( shape: Tuple[int, ...], ctx: Context | str | None = None, geometry: InnerProduct | None = None, + check_level: CheckLevel | bool | None = None, *, inner_product: InnerProduct | None = None, ): @@ -135,11 +147,12 @@ def __init__( shape: Tuple[int, ...], ctx: Context | str | None = None, geometry: InnerProduct | None = None, + check_level: CheckLevel | bool | None = None, *, inner_product: InnerProduct | None = None, ) -> None: geometry = _resolve_elementwise_geometry(geometry, inner_product) - DenseCoordinateSpace.__init__(self, shape, ctx, geometry=geometry) + DenseCoordinateSpace.__init__(self, shape, ctx, geometry=geometry, check_level=check_level) @checked_method(in_space="self") def star(self, x: DenseArray) -> DenseArray: @@ -211,7 +224,6 @@ def tree_unflatten(cls, aux, children): return cls(shape, ctx, geometry=geometry) -@jax_pytree_class class EuclideanElementwiseJordanSpace(ElementwiseJordanSpace, EuclideanJordanAlgebraSpace): """ Real elementwise Euclidean Jordan algebra. @@ -225,6 +237,11 @@ class EuclideanElementwiseJordanSpace(ElementwiseJordanSpace, EuclideanJordanAlg geometry : InnerProduct or None, optional Inner-product geometry. This class is selected only for real contexts with Euclidean coordinate geometry. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. inner_product : InnerProduct or None, optional Alias for ``geometry``. """ @@ -234,10 +251,13 @@ def __init__( shape: Tuple[int, ...], ctx: Context | str | None = None, geometry: InnerProduct | None = None, + check_level: CheckLevel | bool | None = None, *, inner_product: InnerProduct | None = None, ) -> None: resolved_ctx = normalize_context(ctx) geometry = _resolve_elementwise_geometry(geometry, inner_product) - DenseCoordinateSpace.__init__(self, shape, resolved_ctx, geometry=geometry) + DenseCoordinateSpace.__init__( + self, shape, resolved_ctx, geometry=geometry, check_level=check_level + ) _validate_euclidean_elementwise_jordan(self, self.geometry) diff --git a/spacecore/space/concrete/_hermitian.py b/spacecore/space/concrete/_hermitian.py index 83a5902..170f558 100644 --- a/spacecore/space/concrete/_hermitian.py +++ b/spacecore/space/concrete/_hermitian.py @@ -5,9 +5,10 @@ from ..checks import HermitianCheck, SquareMatrixCheck from ..base import EuclideanJordanAlgebraSpace, StarSpace from ._dense_coordinate import DenseCoordinateSpace +from ..._check_policy import CheckLevel from ..._checks import checked_method from ...types import DenseArray -from ...backend import Context +from ...contextual import Context class HermitianSpace(DenseCoordinateSpace, StarSpace, EuclideanJordanAlgebraSpace): @@ -38,6 +39,11 @@ class HermitianSpace(DenseCoordinateSpace, StarSpace, EuclideanJordanAlgebraSpac Whether membership checks enforce Hermitian structure. ctx : Context, str, or None, optional Backend context specification. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. Attributes ---------- @@ -45,6 +51,11 @@ class HermitianSpace(DenseCoordinateSpace, StarSpace, EuclideanJordanAlgebraSpac Matrix dimension. """ + # Hermitian matrices with complex entries form a *real* vector space of + # dimension n^2 -- their complex span is all of Mat(n, C). The dtype says + # "complex" (correctly, about the entries); the scalars are real either way. + declared_scalar_field = "real" + def __init__( self, n: int, @@ -52,12 +63,13 @@ def __init__( rtol: float = 0.0, enforce_herm: bool = True, ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ): if n <= 0: raise ValueError("n must be positive.") shape = (n, n) - super(HermitianSpace, self).__init__(shape, ctx) + super(HermitianSpace, self).__init__(shape, ctx, check_level=check_level) self.atol = atol self.rtol = rtol @@ -89,6 +101,30 @@ def _local_checks(self): ), ) + @checked_method(in_space="self", arg_positions=(1,)) + def scale(self, a: Any, x: DenseArray) -> DenseArray: + """Return ``a * x``, rejecting a multiplier that would leave the space. + + Hermitian matrices are closed under **real** scaling only: for Hermitian + ``H``, ``(aH)* = conj(a) H``, which equals ``aH`` exactly when ``a`` is + real, and ``i H`` is anti-Hermitian. The inherited + :meth:`~spacecore.space.DenseCoordinateSpace.scale` validates only its + *input*, so without this override ``scale(1j, H)`` would return a + skew-Hermitian array still typed as an element of ``Herm(n)`` and the + error would surface at some unrelated later membership check — or never, + at ``check_level="none"``. + + Only a *provably* non-real multiplier is rejected, so a traced scalar + under ``jax.jit`` passes through unchanged. + """ + self.check_scalar(a) + return self._scale_core(a, x) + + def scale_batch(self, a: Any, x: DenseArray) -> DenseArray: + """Return the leading-axis batch scalar product, with the same guard as :meth:`scale`.""" + self.check_scalar(a) + return super().scale_batch(a, x) + def is_hermitian(self, x: DenseArray) -> bool: """Return whether ``x`` satisfies this space's Hermitian check.""" return HermitianCheck( diff --git a/spacecore/space/concrete/_stacked.py b/spacecore/space/concrete/_stacked.py index 6717bd0..a4b7677 100644 --- a/spacecore/space/concrete/_stacked.py +++ b/spacecore/space/concrete/_stacked.py @@ -12,8 +12,11 @@ ) from ._tree_space import TreeSpace, _space_capabilities from ..._checks import checked_method -from ..._contextual import resolve_context_priority -from ...backend import BackendOps, Context, jax_pytree_class +from ..._check_policy import CheckLevel +from ...contextual import resolve_context_priority +from ...contextual import Context +from ...backend import BackendOps +from ...backend import PyTreeNode from ...types import DenseArray from ..checks import BackendCheck, DTypeCheck, FieldCheck, ShapeCheck @@ -67,8 +70,7 @@ def _require_base(base: Space, capability: type, owner: str) -> None: ) -@jax_pytree_class -class StackedSpace(CoordinateSpace): +class StackedSpace(PyTreeNode, CoordinateSpace): """ Leading-axis copies of a coordinate leaf space. @@ -85,9 +87,20 @@ class StackedSpace(CoordinateSpace): ctx : Context, str, or None, optional Context specification. If omitted, the context is resolved from ``base``. + check_level : {"none", "cheap", "standard", "strict"}, optional + Runtime validation policy for this object. When omitted, the ambient + default (see :func:`spacecore.get_check_level`) is used. Unlike the + backend/dtype context, the validation policy is a property of the bound + object, not of the :class:`Context`. """ - def __new__(cls, base: Space, count: int, ctx: Context | str | None = None): + def __new__( + cls, + base: Space, + count: int, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ): if cls is StackedSpace: base = _validate_stacked_base(base) _validate_count(count) @@ -95,13 +108,21 @@ def __new__(cls, base: Space, count: int, ctx: Context | str | None = None): cls = _stacked_class_for(_stacked_capabilities(base.convert(resolved_ctx))) return super(StackedSpace, cls).__new__(cls) - def __init__(self, base: Space, count: int, ctx: Context | str | None = None) -> None: + def __init__( + self, + base: Space, + count: int, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, + ) -> None: base = _validate_stacked_base(base, type(self).__name__) count = _validate_count(count, type(self).__name__) ctx = resolve_context_priority(ctx, base) self.base = base.convert(ctx) self.count = count - super().__init__((self.count,) + tuple(self.base.shape), ctx) + super().__init__( + (self.count,) + tuple(self.base.shape), ctx, check_level=check_level + ) def _eq_algebra(self, other: Any) -> bool: # Tier 2: count + base. ``base == other.base`` is load-bearing — the @@ -338,37 +359,33 @@ def unit(self) -> Any: return self.ops.broadcast_to(base.unit(), self.shape) -@jax_pytree_class class _StackedInnerProductSpace(_StackedInnerProductMixin, StackedSpace, InnerProductSpace): """Stacked space whose base supports an inner product.""" - def __init__(self, base, count, ctx=None): + def __init__(self, base, count, ctx=None, check_level=None): base = _validate_stacked_base(base, type(self).__name__) _require_base(base, InnerProductSpace, type(self).__name__) - super().__init__(base, count, ctx) + super().__init__(base, count, ctx, check_level=check_level) -@jax_pytree_class class _StackedStarSpace(_StackedStarMixin, StackedSpace, StarSpace): """Stacked space whose base supports a star operation.""" - def __init__(self, base, count, ctx=None): + def __init__(self, base, count, ctx=None, check_level=None): base = _validate_stacked_base(base, type(self).__name__) _require_base(base, StarSpace, type(self).__name__) - super().__init__(base, count, ctx) + super().__init__(base, count, ctx, check_level=check_level) -@jax_pytree_class class _StackedJordanAlgebraSpace(_StackedJordanMixin, StackedSpace, JordanAlgebraSpace): """Stacked space whose base supports Jordan algebra operations.""" - def __init__(self, base, count, ctx=None): + def __init__(self, base, count, ctx=None, check_level=None): base = _validate_stacked_base(base, type(self).__name__) _require_base(base, JordanAlgebraSpace, type(self).__name__) - super().__init__(base, count, ctx) + super().__init__(base, count, ctx, check_level=check_level) -@jax_pytree_class class _StackedEuclideanJordanAlgebraSpace( _StackedInnerProductMixin, _StackedJordanMixin, @@ -377,14 +394,13 @@ class _StackedEuclideanJordanAlgebraSpace( ): """Stacked space whose base supports Euclidean Jordan algebra operations.""" - def __init__(self, base, count, ctx=None): + def __init__(self, base, count, ctx=None, check_level=None): base = _validate_stacked_base(base, type(self).__name__) _require_base(base, EuclideanJordanAlgebraSpace, type(self).__name__) - super().__init__(base, count, ctx) + super().__init__(base, count, ctx, check_level=check_level) _require_base(self.base, EuclideanJordanAlgebraSpace, type(self).__name__) -@jax_pytree_class class _StackedInnerProductStarSpace( _StackedInnerProductMixin, _StackedStarMixin, @@ -395,7 +411,6 @@ class _StackedInnerProductStarSpace( """Stacked implementation for inner-product plus star capability.""" -@jax_pytree_class class _StackedInnerProductJordanSpace( _StackedInnerProductMixin, _StackedJordanMixin, @@ -406,7 +421,6 @@ class _StackedInnerProductJordanSpace( """Stacked implementation for inner-product plus Jordan capability.""" -@jax_pytree_class class _StackedStarJordanSpace( _StackedStarMixin, _StackedJordanMixin, @@ -417,7 +431,6 @@ class _StackedStarJordanSpace( """Stacked implementation for star plus Jordan capability.""" -@jax_pytree_class class _StackedInnerProductStarJordanSpace( _StackedInnerProductMixin, _StackedStarMixin, @@ -430,7 +443,6 @@ class _StackedInnerProductStarJordanSpace( """Stacked implementation for inner-product, star, and Jordan capability.""" -@jax_pytree_class class _StackedEuclideanJordanStarSpace( _StackedStarMixin, _StackedEuclideanJordanAlgebraSpace, diff --git a/spacecore/space/concrete/_tree_space.py b/spacecore/space/concrete/_tree_space.py index bdae278..b04a3d4 100644 --- a/spacecore/space/concrete/_tree_space.py +++ b/spacecore/space/concrete/_tree_space.py @@ -5,10 +5,12 @@ import optree -from ..._check_policy import normalize_check_level +from ..._check_policy import minimum_check_level, normalize_check_level from ..._checks import checked_method -from ..._contextual import resolve_context_priority -from ...backend import BackendOps, CheckLevel, Context, jax_pytree_class +from ...contextual import resolve_context_priority +from ...contextual import Context +from ...backend import BackendOps, CheckLevel +from ...backend import PyTreeNode from ...types import DenseArray from ..base import ( CoordinateSpace, @@ -98,16 +100,8 @@ def _format_path(path: tuple[Any, ...]) -> str: return result -def _context_with_check_level(ctx: Context, check_level: CheckLevel | str | None) -> Context: - """Return ``ctx`` with an explicit validation policy when requested.""" - if check_level is None or ctx.check_level == check_level: - return ctx - return Context(ctx.ops, dtype=ctx.dtype, check_level=normalize_check_level(check_level)) - - -@jax_pytree_class @dataclass(frozen=True, eq=False) -class TreeElement: +class TreeElement(PyTreeNode): r""" Bind ordered leaves to a :class:`TreeSpace`. @@ -252,8 +246,7 @@ def validation_message( return "Invalid TreeSpace leaf." -@jax_pytree_class -class TreeSpace(CoordinateSpace): +class TreeSpace(PyTreeNode, CoordinateSpace): r""" Represent a finite direct product as a Python tree. @@ -336,7 +329,6 @@ def __new__( if cls is TreeSpace: spaces = _validate_leaf_spaces(leaf_spaces) resolved_ctx = resolve_context_priority(ctx, *spaces) - resolved_ctx = _context_with_check_level(resolved_ctx, check_level) converted = tuple(space.convert(resolved_ctx) for space in spaces) cls = _TREE_REGISTRY.get(_tree_capabilities(converted), TreeSpace) return super(TreeSpace, cls).__new__(cls) @@ -351,7 +343,6 @@ def __init__( ) -> None: spaces = _validate_leaf_spaces(leaf_spaces, type(self).__name__) resolved_ctx = resolve_context_priority(ctx, *spaces) - resolved_ctx = _context_with_check_level(resolved_ctx, check_level) treespec = treedef if isinstance(treedef, optree.PyTreeSpec) else optree.tree_structure(treedef) if treespec.num_leaves != len(spaces): raise ValueError( @@ -372,6 +363,12 @@ def __init__( slice(offsets[index], offsets[index + 1]) for index in range(len(self._dims)) ) super(TreeSpace, self).__init__((offsets[-1],), resolved_ctx) + if check_level is not None: + self._check_level = normalize_check_level(check_level) + else: + leaf_levels = tuple(space.check_level for space in uniform_spaces) + if leaf_levels: + self._check_level = minimum_check_level(leaf_levels) @classmethod def from_leaf_spaces( @@ -711,9 +708,8 @@ def star(self, x: Any) -> Any: ) -@jax_pytree_class @dataclass(frozen=True) -class TreeSpectralDecomposition: +class TreeSpectralDecomposition(PyTreeNode): """ Store leafwise Jordan spectral data in deterministic leaf order. @@ -986,9 +982,10 @@ class _TreeEuclideanJordanStarSpace( } ) -for _tree_type in set(_TREE_REGISTRY.values()): - jax_pytree_class(_tree_type) - +# The capability-dispatch subclasses above register themselves: each derives from +# ``TreeSpace``, so ``PyTreeNode.__init_subclass__`` fires at class creation. The +# explicit loop this replaces existed because a *decorator* cannot reach classes +# built through a dispatch table — inheritance can. __all__ = [ "TreeElement", diff --git a/tests/backend/_references.py b/tests/backend/_references.py index eadcd86..d2b5874 100644 --- a/tests/backend/_references.py +++ b/tests/backend/_references.py @@ -25,6 +25,103 @@ def test_some_op(backend_ops, conformance_dtype): from typing import Any, Sequence +import numpy as _np + + +# ---------------------------------------------------------------------- +# Pytree / axis helpers for the control-flow references +# ---------------------------------------------------------------------- +# Deliberately independent of ``EagerControlFlowMixin``'s ``_tree_*`` helpers: +# ``scan`` and ``vmap`` keep most of their logic in pytree walking, so a +# reference that reused those helpers would be comparing them to themselves. + + +def _ref_to_numpy(x: Any) -> Any: + """Best-effort conversion of one backend leaf to a NumPy array.""" + if hasattr(x, "detach"): # torch + x = x.detach().cpu() + return _np.asarray(x) + + +def _ref_leaves(tree: Any) -> list[Any]: + """Flatten a dict/tuple/list pytree to its leaves, in structure order.""" + if isinstance(tree, dict): + return [leaf for v in tree.values() for leaf in _ref_leaves(v)] + if isinstance(tree, (tuple, list)): + return [leaf for v in tree for leaf in _ref_leaves(v)] + return [tree] + + +def _ref_tree_index(tree: Any, i: int) -> Any: + """Take ``leaf[i]`` along axis 0 of every leaf, preserving structure.""" + if isinstance(tree, dict): + return {k: _ref_tree_index(v, i) for k, v in tree.items()} + if isinstance(tree, tuple): + return tuple(_ref_tree_index(v, i) for v in tree) + if isinstance(tree, list): + return [_ref_tree_index(v, i) for v in tree] + return tree[i] + + +def _ref_tree_stack(steps: Sequence[Any]) -> Any: + """Stack a list of same-structure pytrees leafwise along a new axis 0.""" + if not steps: + return () + first = steps[0] + if isinstance(first, dict): + return {k: _ref_tree_stack([s[k] for s in steps]) for k in first} + if isinstance(first, tuple): + return tuple(_ref_tree_stack([s[i] for s in steps]) for i in range(len(first))) + if isinstance(first, list): + return [_ref_tree_stack([s[i] for s in steps]) for i in range(len(first))] + return _np.stack([_ref_to_numpy(s) for s in steps], axis=0) + + +def _ref_leading_length(tree: Any) -> int: + """Length of axis 0 of the first leaf — how ``scan`` infers its length.""" + return int(_np.shape(_ref_to_numpy(_ref_leaves(tree)[0]))[0]) + + +def _ref_axis_size(arg: Any, axis: Any) -> int | None: + """Size of ``arg`` along ``axis``; ``None`` when the argument is unmapped.""" + if axis is None: + return None + if isinstance(arg, tuple): + axes = axis if isinstance(axis, (tuple, list)) else (axis,) * len(arg) + for sub, sub_axis in zip(arg, axes): + size = _ref_axis_size(sub, sub_axis) + if size is not None: + return size + return None + return int(_np.shape(_ref_to_numpy(arg))[int(axis)]) + + +def _ref_axis_take(arg: Any, axis: Any, i: int) -> Any: + """Slice index ``i`` out of ``axis``; pass through when ``axis`` is None.""" + if axis is None: + return arg + if isinstance(arg, tuple): + axes = axis if isinstance(axis, (tuple, list)) else (axis,) * len(arg) + return tuple(_ref_axis_take(sub, a, i) for sub, a in zip(arg, axes)) + return _np.take(_ref_to_numpy(arg), i, axis=int(axis)) + + +def _ref_stack_at(outputs: Sequence[Any], out_axes: Any) -> Any: + """Stack per-call outputs along ``out_axes`` (``None`` keeps the first).""" + first = outputs[0] + if isinstance(first, tuple): + axes = ( + out_axes + if isinstance(out_axes, (tuple, list)) + else (out_axes,) * len(first) + ) + return tuple( + _ref_stack_at([o[i] for o in outputs], a) for i, a in enumerate(axes) + ) + if out_axes is None: + return first + return _np.stack([_ref_to_numpy(o) for o in outputs], axis=int(out_axes)) + class ReferenceOps: """Calls the native library a backend wraps. The truth, not a wrapper. @@ -649,6 +746,96 @@ def while_loop(self, cond_fun, body_fun, init_val): def cond(self, pred: bool, true_fun, false_fun, *operands): return true_fun(*operands) if bool(pred) else false_fun(*operands) + def scan(self, f, init, xs, length=None, reverse=False, unroll=1): + """Python-loop reference for ``scan``. + + Stacks per-step outputs with NumPy rather than the backend's ``stack``, + and walks pytrees with its own helpers, so the reference shares no code + with the implementation under test — the eager ``scan`` keeps most of + its logic in exactly that pytree walking. ``unroll`` is accepted for + signature parity and by definition cannot change the result. + """ + carry = init + if xs is None: + if length is None: + raise ValueError("scan(xs=None) requires an explicit `length`.") + n = int(length) + + def take(_i): + return None + else: + n = int(length) if length is not None else _ref_leading_length(xs) + + def take(i): + return _ref_tree_index(xs, i) + + steps = range(n - 1, -1, -1) if reverse else range(n) + ys: list[Any] = [] + for i in steps: + carry, y = f(carry, take(i)) + ys.append(y) + if reverse: + ys.reverse() + return carry, _ref_tree_stack(ys) + + def vmap(self, fn, in_axes=0, out_axes=0): + """Python-loop reference for ``vmap``: slice, call, stack. + + The definition of vectorization written out. Axis ``None`` means + "pass this argument through unsliced"; when every argument has axis + ``None`` there is nothing to map over and ``fn`` is called once. + """ + + def mapped(*args: Any) -> Any: + axes = ( + tuple(in_axes) + if isinstance(in_axes, (tuple, list)) + else (in_axes,) * len(args) + ) + size = None + for arg, axis in zip(args, axes): + size = _ref_axis_size(arg, axis) + if size is not None: + break + if size is None: + return fn(*args) + outputs = [ + fn(*(_ref_axis_take(arg, axis, i) for arg, axis in zip(args, axes))) + for i in range(size) + ] + return _ref_stack_at(outputs, out_axes) + + return mapped + + def vectorize(self, pyfunc, *, excluded=None, signature=None): + """Elementwise reference for ``vectorize`` (no ``signature`` support). + + Broadcasts the non-excluded arguments and calls ``pyfunc`` once per + index — the semantics :func:`numpy.vectorize` documents, spelled out as + a loop instead of delegated to it. + """ + if signature is not None: + raise NotImplementedError( + "the reference vectorize covers the elementwise case only" + ) + skip = set() if excluded is None else set(excluded) + + def vectorized(*args: Any) -> Any: + mapped_idx = [i for i in range(len(args)) if i not in skip] + arrays = _np.broadcast_arrays( + *[_np.asarray(_ref_to_numpy(args[i])) for i in mapped_idx] + ) + shape = arrays[0].shape if arrays else () + out = _np.empty(shape, dtype=object) + for idx in _np.ndindex(*shape): + call = list(args) + for slot, arr in zip(mapped_idx, arrays): + call[slot] = arr[idx] + out[idx] = pyfunc(*call) + return _np.array(out.tolist()) + + return vectorized + # ------------------------------------------------------------------ # Sparse # ------------------------------------------------------------------ diff --git a/tests/backend/test_context.py b/tests/backend/test_context.py index 579b349..2dcf206 100644 --- a/tests/backend/test_context.py +++ b/tests/backend/test_context.py @@ -17,7 +17,9 @@ 7. ``assert_dense`` / ``assert_sparse`` gates. 8. ``convert`` dispatch (dense → ``asarray``, sparse → ``assparse``). -9. ``check_level`` normalization and the deprecated ``enable_checks`` alias. +9. Validation policy is *not* a context property: ``check_level`` is carried by + the context-bound object, and the deprecated ``enable_checks`` Boolean + survives only as the ``normalize_check_level`` shim. Generic per-op behavior lives in :mod:`tests.backend.test_operations`; this module pins the ``Context`` API only. @@ -29,6 +31,7 @@ import scipy.sparse as sps import spacecore as sc +from spacecore._check_policy import level_to_enabled, normalize_check_level from tests._helpers import has_cupy, has_jax, has_torch, to_numpy from tests.backend._conformance import ( @@ -150,15 +153,15 @@ def test_rejects_unknown_ops(self): # 2. Equality and hash # =========================================================================== class TestEqualityAndHash: - def test_equality_same_family_same_dtype(self): - a = sc.Context(sc.NumpyOps(), dtype=np.float64) - b = sc.Context(sc.NumpyOps(), dtype=np.float64) - assert a == b - - def test_inequality_different_dtype(self): - a = sc.Context(sc.NumpyOps(), dtype=np.float32) - b = sc.Context(sc.NumpyOps(), dtype=np.float64) - assert a != b + """Cross-backend equality only. + + The full value-object equality/hash contract — reflexivity, symmetry, + transitivity, hash-consistency, frozen-immutability, dict-key usability, + and foreign-type inequality — lives in + ``tests/context/test_context_contracts.py`` (via ``assert_equality_contract``). + This keeps only the one case that requires a *real* second backend, which + the synthetic single-backend contract test cannot exercise. + """ def test_inequality_different_family(self): if not has_jax(): @@ -169,30 +172,6 @@ def test_inequality_different_family(self): b = sc.Context(sc.JaxOps(), dtype=dt) assert a != b - def test_inequality_different_check_level(self): - a = sc.Context(sc.NumpyOps(), check_level="standard") - b = sc.Context(sc.NumpyOps(), check_level="none") - assert a != b - - def test_inequality_against_non_context(self): - ctx = sc.Context(sc.NumpyOps()) - assert ctx != "Context" - assert ctx != sc.NumpyOps() - assert (ctx == 42) is False - - def test_hashable_and_dict_keyable(self): - a = sc.Context(sc.NumpyOps(), dtype=np.float64) - b = sc.Context(sc.NumpyOps(), dtype=np.float64) - table = {a: "first"} - table[b] = "second" - assert table[a] == "second" - assert len(table) == 1 - - def test_is_frozen(self): - ctx = sc.Context(sc.NumpyOps()) - with pytest.raises((AttributeError, Exception)): - ctx.dtype = np.float32 # type: ignore[misc] - # =========================================================================== # 3. Context.dtype defaulting per family @@ -346,12 +325,13 @@ def test_refused_when_allow_sparse_false(self): # 6. Repr stability # =========================================================================== class TestRepr: - def test_repr_includes_family_dtype_check_level(self): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="standard") + def test_repr_includes_family_and_dtype_only(self): + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) r = repr(ctx) assert "Context(" in r assert "NumpyOps" in r - assert "check_level='standard'" in r + assert "float64" in r + assert "check_level" not in r def test_repr_is_deterministic(self): a = repr(sc.Context(sc.NumpyOps(), dtype=np.float64)) @@ -409,27 +389,32 @@ def test_rejects_unknown_type(self): ctx.convert([1.0, 2.0]) -# ---- 9. check_level normalization + enable_checks deprecation alias ------- +# ---- 9. check_level lives on the bound object, not on the Context --------- class TestCheckLevel: + def test_context_carries_no_validation_policy(self): + ctx = sc.Context(sc.NumpyOps()) + assert not hasattr(ctx, "check_level") + assert not hasattr(ctx, "enable_checks") + def test_default_check_level_is_standard(self): ctx = sc.Context(sc.NumpyOps()) - assert ctx.check_level == "standard" + assert sc.DenseCoordinateSpace((2,), ctx).check_level == "standard" @pytest.mark.parametrize("level", ["none", "cheap", "standard", "strict"]) def test_explicit_check_level(self, level): - ctx = sc.Context(sc.NumpyOps(), check_level=level) - assert ctx.check_level == level + ctx = sc.Context(sc.NumpyOps()) + assert sc.DenseCoordinateSpace((2,), ctx, check_level=level).check_level == level def test_enable_checks_true_maps_to_standard(self): with pytest.warns(DeprecationWarning): - ctx = sc.Context(sc.NumpyOps(), enable_checks=True) - assert ctx.check_level == "standard" + level = normalize_check_level(enable_checks=True, warn_legacy=True) + assert level == "standard" def test_enable_checks_false_maps_to_none(self): with pytest.warns(DeprecationWarning): - ctx = sc.Context(sc.NumpyOps(), enable_checks=False) - assert ctx.check_level == "none" + level = normalize_check_level(enable_checks=False, warn_legacy=True) + assert level == "none" - def test_enable_checks_property_is_legacy_view(self): - assert sc.Context(sc.NumpyOps(), check_level="none").enable_checks is False - assert sc.Context(sc.NumpyOps(), check_level="standard").enable_checks is True + def test_enable_checks_is_the_legacy_view_of_a_level(self): + assert level_to_enabled("none") is False + assert level_to_enabled("standard") is True diff --git a/tests/backend/test_jax_pytree_class.py b/tests/backend/test_jax_pytree_class.py deleted file mode 100644 index ca92e40..0000000 --- a/tests/backend/test_jax_pytree_class.py +++ /dev/null @@ -1,165 +0,0 @@ -"""Tests for the :func:`spacecore.jax_pytree_class` decorator. - -The decorator registers a class as a JAX PyTree node so that instances can -flow through ``jax.tree_util``, ``jax.jit``, ``jax.vmap``, and ``jax.grad``. -When JAX is not installed, the decorator is a no-op. -""" -from __future__ import annotations - -import pytest - -import spacecore as sc - -from tests._helpers import has_jax - - -@pytest.mark.skipif(not has_jax(), reason="jax is not installed") -def test_jax_pytree_class_is_noop_when_import_fails(monkeypatch): - """backend-001: force ``from jax import tree_util`` to raise, covering the - ``except Exception: return klass`` branch of the decorator (jax/_pytree.py - lines 24-27). - - When JAX is installed we cannot truly uninstall it, so we make the symbol - the decorator imports unavailable: deleting ``jax.tree_util`` makes - ``from jax import tree_util`` raise ``ImportError``. The decorator must - swallow it and return the class unchanged, and the class must remain a - usable plain class. - """ - import jax - - # Removing the attribute makes ``from jax import tree_util`` raise. - monkeypatch.delattr(jax, "tree_util", raising=True) - - @sc.jax_pytree_class - class Foo: - def __init__(self, a: int, b: int) -> None: - self.a = a - self.b = b - - # Decorator returned the class unchanged (no registration happened). - inst = Foo(1, 2) - assert isinstance(inst, Foo) - assert (inst.a, inst.b) == (1, 2) - - -def test_jax_pytree_class_is_noop_when_import_fails_no_jax(monkeypatch): - """backend-001 (no-JAX path): if ``jax`` itself is unimportable the - decorator returns the class unchanged. Simulated by blocking the import. - """ - import builtins - - real_import = builtins.__import__ - - def fake_import(name, *args, **kwargs): - if name == "jax" or name.startswith("jax."): - raise ImportError("simulated: jax not installed") - return real_import(name, *args, **kwargs) - - monkeypatch.setattr(builtins, "__import__", fake_import) - - @sc.jax_pytree_class - class Bar: - def __init__(self, v: int) -> None: - self.v = v - - inst = Bar(5) - assert isinstance(inst, Bar) - assert inst.v == 5 - - -@pytest.mark.skipif(not has_jax(), reason="jax is not installed") -def test_jax_pytree_class_round_trips_via_tree_util(): - """A decorated class with proper flatten/unflatten round-trips through - ``jax.tree_util.tree_flatten`` and ``tree_unflatten``. - """ - import jax.tree_util as jtu - - @sc.jax_pytree_class - class Pair: - def __init__(self, x: float, y: float) -> None: - self.x = x - self.y = y - - def tree_flatten(self): - return (self.x, self.y), None - - @classmethod - def tree_unflatten(cls, _aux, children): - x, y = children - inst = cls.__new__(cls) - inst.x = x - inst.y = y - return inst - - def __eq__(self, other): - return isinstance(other, Pair) and self.x == other.x and self.y == other.y - - inst = Pair(1.0, 2.0) - leaves, treedef = jtu.tree_flatten(inst) - assert tuple(leaves) == (1.0, 2.0) - rebuilt = jtu.tree_unflatten(treedef, leaves) - assert rebuilt == inst - - -@pytest.mark.skipif(not has_jax(), reason="jax is not installed") -def test_jax_pytree_class_supports_tree_map(): - """``jax.tree_util.tree_map`` applies a leaf transformation across the - structure when the class is registered. - """ - import jax.tree_util as jtu - - @sc.jax_pytree_class - class Vec3: - def __init__(self, a, b, c): - self.a = a - self.b = b - self.c = c - - def tree_flatten(self): - return (self.a, self.b, self.c), None - - @classmethod - def tree_unflatten(cls, _aux, children): - inst = cls.__new__(cls) - inst.a, inst.b, inst.c = children - return inst - - v = Vec3(1.0, 2.0, 3.0) - doubled = jtu.tree_map(lambda x: x * 2.0, v) - assert doubled.a == 2.0 and doubled.b == 4.0 and doubled.c == 6.0 - - -@pytest.mark.skipif(not has_jax(), reason="jax is not installed") -def test_backend_ops_classes_are_themselves_pytree_compatible(): - """Every ``*Ops`` class is decorated, so ``tree_map`` does not error on - a single ``BackendOps`` instance.""" - import jax.tree_util as jtu - - ops = sc.NumpyOps() - # NumpyOps as a leaf: tree_map identity-mapping a leaf returns the leaf. - result = jtu.tree_map(lambda x: x, ops) - assert isinstance(result, sc.NumpyOps) - - -@pytest.mark.skipif(not has_jax(), reason="jax is not installed") -def test_jax_pytree_class_handles_redundant_registration(): - """Re-decorating an already-registered class is tolerated (catches the - JAX ``ValueError`` for duplicate registration internally). - """ - @sc.jax_pytree_class - class Once: - def __init__(self, v): - self.v = v - - def tree_flatten(self): - return (self.v,), None - - @classmethod - def tree_unflatten(cls, _aux, children): - inst = cls.__new__(cls) - inst.v = children[0] - return inst - - # Re-application must not raise. - Reapplied = sc.jax_pytree_class(Once) - assert Reapplied is Once diff --git a/tests/backend/test_optional_guard.py b/tests/backend/test_optional_guard.py new file mode 100644 index 0000000..50f2440 --- /dev/null +++ b/tests/backend/test_optional_guard.py @@ -0,0 +1,188 @@ +"""Guard behavior for optional backend imports (``backend._optional``). + +An optional backend can fail to load in several distinct ways, and only one of +them means "not installed": + +* **absent** — the dependency is simply not installed (``ModuleNotFoundError`` + naming the dependency itself); +* **broken / shadowed** — the dependency *is* importable but its import fails + (a partial install, an ABI mismatch, or a namespace shim), raising a plain + ``ImportError``; +* **transitive-missing** — a *different* module the backend needs is absent + (``ModuleNotFoundError`` naming something other than the dependency). + +Only *absent* is a routine condition; the others are installation faults. In +every case loading an optional backend must degrade gracefully rather than +abort ``import spacecore`` — crashing the whole library over an optional +backend the user was never going to use is the wrong failure mode (Ousterhout, +*A Philosophy of Software Design*: crash on invariant violation, handle +environmental errors gracefully; Hunt & Thomas, *The Pragmatic Programmer*: +fail, but do not corrupt). + +These tests replace ``importlib.import_module`` with a *Test Stub* (Meszaros, +*xUnit Test Patterns*) so the SUT is isolated from whatever backends happen to +be installed — the suite must be *Repeatable* regardless of the environment +(same source, principle: isolate the SUT from irrelevant dependencies), and +must not become an *Erratic Test*. +""" +from __future__ import annotations + +import warnings + +import pytest + +import spacecore as sc +from spacecore.backend import _optional + + +@pytest.fixture(autouse=True) +def _fresh_discovery_cache(): + """Run every test in this module against uncached backend discovery. + + ``available_ops`` is memoized for the process, which is right in production + (availability cannot change within a run, and re-probing makes one broken + install warn repeatedly) but wrong here: these tests monkeypatch the import + machinery and the entry points, then assert on what discovery *now* returns. + Against a warm cache they read the pre-patch answer — not merely failing, but + in one case passing vacuously, as ``test_external_cannot_shadow_builtin_family`` + did until this fixture was added. + + Clearing afterwards as well stops a patched result leaking into a later test. + """ + _optional.available_ops.cache_clear() + yield + _optional.available_ops.cache_clear() + + +def _raises(exc): + def _stub(*_args, **_kwargs): + raise exc + + return _stub + + +class _FakeEntryPoint: + """A minimal stand-in for an ``importlib.metadata`` EntryPoint. + + Test Stub (Meszaros, *xUnit Test Patterns*): lets discovery be exercised + without installing a real distribution, keeping the test Repeatable. + """ + + def __init__(self, name, target): + self.name = name + self._target = target + + def load(self): + if isinstance(self._target, BaseException): + raise self._target + return self._target + + +class _ExternalOps(sc.NumpyOps): + """A pretend third-party backend advertised via an entry point.""" + + _family = "external_demo" + + +class TestImportBackendGuard: + def test_absent_dependency_returns_none_without_warning(self, monkeypatch): + # ModuleNotFoundError naming the dependency itself == "not installed". + monkeypatch.setattr( + _optional.importlib, + "import_module", + _raises(ModuleNotFoundError("No module named 'jax'", name="jax")), + ) + with warnings.catch_warnings(): + warnings.simplefilter("error") # any warning would fail the test + assert _optional.import_backend(".jax", "jax") is None + + def test_broken_backend_warns_and_skips(self, monkeypatch): + # Installed-but-broken raises a plain ImportError (not ModuleNotFound). + monkeypatch.setattr( + _optional.importlib, + "import_module", + _raises(ImportError("cannot import name 'abs' from partially initialized module")), + ) + with pytest.warns(UserWarning): + assert _optional.import_backend(".cupy", "cupy") is None + + def test_transitive_missing_dependency_warns_and_skips(self, monkeypatch): + # ModuleNotFoundError for a *different* module (a transitive dependency) + # must not be treated as fatal / re-raised. + monkeypatch.setattr( + _optional.importlib, + "import_module", + _raises(ModuleNotFoundError("No module named 'cupyx'", name="cupyx")), + ) + with pytest.warns(UserWarning): + assert _optional.import_backend(".cupy", "cupy") is None + + def test_successful_import_is_returned(self, monkeypatch): + sentinel = object() + monkeypatch.setattr(_optional.importlib, "import_module", lambda *a, **k: sentinel) + assert _optional.import_backend(".jax", "jax") is sentinel + + +class TestAvailableOpsResilience: + def test_numpy_survives_when_every_optional_backend_is_broken(self, monkeypatch): + # Even if all optional backends fail to import, NumpyOps remains and no + # exception escapes. + monkeypatch.setattr( + _optional.importlib, "import_module", _raises(ImportError("boom")) + ) + with warnings.catch_warnings(): + warnings.simplefilter("ignore") + ops = _optional.available_ops() + names = [o.__name__ for o in ops] + assert names == ["NumpyOps"] + + def test_available_ops_starts_with_numpy(self): + # Guard the real environment: NumpyOps is always first, unconditionally. + assert _optional.available_ops()[0] is sc.NumpyOps + + +class TestEntryPointDiscovery: + """External backends advertised via ``spacecore.backends`` entry points. + + Discovery makes the backend layer *open for extension* (Open-Closed + Principle): a package registers a backend by installation, not by editing + SpaceCore. Loading is defensive — a broken or wrong-typed plugin is skipped, + never fatal (Ousterhout, *A Philosophy of Software Design*: handle + environmental errors gracefully). + """ + + def _patch(self, monkeypatch, entry_points): + monkeypatch.setattr(_optional, "_entry_points", lambda group: list(entry_points)) + + def test_valid_external_backend_is_discovered(self, monkeypatch): + self._patch(monkeypatch, [_FakeEntryPoint("external_demo", _ExternalOps)]) + assert _ExternalOps in _optional.discover_entry_point_ops() + + def test_external_backend_appears_in_available_ops(self, monkeypatch): + self._patch(monkeypatch, [_FakeEntryPoint("external_demo", _ExternalOps)]) + families = [o._family for o in _optional.available_ops()] + assert families[0] == "numpy" # built-in still first + assert "external_demo" in families + + def test_non_backendops_target_is_skipped_with_warning(self, monkeypatch): + self._patch(monkeypatch, [_FakeEntryPoint("bogus", object)]) + with pytest.warns(UserWarning): + assert _optional.discover_entry_point_ops() == [] + + def test_load_failure_is_non_fatal(self, monkeypatch): + self._patch(monkeypatch, [_FakeEntryPoint("broken", ImportError("boom"))]) + with pytest.warns(UserWarning): + assert _optional.discover_entry_point_ops() == [] + + def test_external_cannot_shadow_builtin_family(self, monkeypatch): + # A plugin advertising the "numpy" family must not displace the built-in. + class _RogueNumpy(sc.NumpyOps): + _family = "numpy" + + self._patch(monkeypatch, [_FakeEntryPoint("rogue", _RogueNumpy)]) + numpy_entries = [o for o in _optional.available_ops() if o._family == "numpy"] + assert numpy_entries == [sc.NumpyOps] + + def test_entry_points_helper_returns_a_list(self): + # The version-compat shim returns a list on the running interpreter. + assert isinstance(_optional._entry_points(_optional._ENTRY_POINT_GROUP), list) diff --git a/tests/backend/test_pytree_registry.py b/tests/backend/test_pytree_registry.py new file mode 100644 index 0000000..36e868a --- /dev/null +++ b/tests/backend/test_pytree_registry.py @@ -0,0 +1,429 @@ +"""Tests for the backend-neutral container registry (:mod:`spacecore.backend._container`). + +The wiring contract is exercised against a *fake* registrar rather than a real +backend: a registrar is just ``Callable[[type], None]`` and +:class:`PyTreeRegistry` is instantiable, so the whole class x backend +cross-product logic is testable with no backend installed and no arrays. Only +the integration test that asserts against a real backend's registry needs one. +""" +from __future__ import annotations + +import inspect +from abc import abstractmethod + +import pytest + +from spacecore.backend._container import PyTreeNode, PyTreeRegistry, registry as global_registry + +from tests._helpers import has_jax, has_torch + + +class _Recorder: + """Fake tree protocol: records the classes handed to it.""" + + def __init__(self) -> None: + self.seen: list[type] = [] + + def __call__(self, cls: type) -> None: + self.seen.append(cls) + + +class _Alpha: + pass + + +class _Beta: + pass + + +# -------------------------------------------------------------------------- +# Two-axis back-fill +# -------------------------------------------------------------------------- + +def test_class_then_backend_wires_the_pair(): + """Backend arriving after the class back-fills it (backend loaded late).""" + reg, rec = PyTreeRegistry(), _Recorder() + reg.register_class(_Alpha) + assert rec.seen == [] + reg.register_backend("fake", rec) + assert rec.seen == [_Alpha] + assert reg.is_wired("fake", _Alpha) + + +def test_backend_then_class_wires_the_pair(): + """Class arriving after the backend forward-fills (the normal import order).""" + reg, rec = PyTreeRegistry(), _Recorder() + reg.register_backend("fake", rec) + reg.register_class(_Alpha) + assert rec.seen == [_Alpha] + assert reg.is_wired("fake", _Alpha) + + +def test_wiring_is_order_independent(): + """Either arrival order must reach the same set of wired pairs.""" + a, rec_a = PyTreeRegistry(), _Recorder() + a.register_class(_Alpha) + a.register_class(_Beta) + a.register_backend("fake", rec_a) + + b, rec_b = PyTreeRegistry(), _Recorder() + b.register_backend("fake", rec_b) + b.register_class(_Alpha) + b.register_class(_Beta) + + assert set(rec_a.seen) == set(rec_b.seen) == {_Alpha, _Beta} + assert a.registered_classes() == b.registered_classes() + + +def test_every_backend_gets_every_class(): + """The registry maintains the full cross-product, not just the last pair.""" + reg = PyTreeRegistry() + first, second = _Recorder(), _Recorder() + reg.register_backend("one", first) + reg.register_class(_Alpha) + reg.register_backend("two", second) # back-fills _Alpha + reg.register_class(_Beta) # forward-fills into both + + assert set(first.seen) == {_Alpha, _Beta} + assert set(second.seen) == {_Alpha, _Beta} + assert set(reg.backends()) == {"one", "two"} + + +# -------------------------------------------------------------------------- +# Idempotency +# -------------------------------------------------------------------------- + +def test_duplicate_class_registration_wires_once(): + reg, rec = PyTreeRegistry(), _Recorder() + reg.register_backend("fake", rec) + reg.register_class(_Alpha) + reg.register_class(_Alpha) + assert rec.seen == [_Alpha] + assert reg.registered_classes() == (_Alpha,) + + +def test_duplicate_backend_registration_wires_once(): + """Re-installing a protocol must not re-register (backends raise on that).""" + reg, rec = PyTreeRegistry(), _Recorder() + reg.register_class(_Alpha) + reg.register_backend("fake", rec) + reg.register_backend("fake", rec) + assert rec.seen == [_Alpha] + + +# -------------------------------------------------------------------------- +# Failure containment: import must survive a broken registrar +# -------------------------------------------------------------------------- + +def test_failing_registrar_warns_instead_of_raising(): + """A backend that rejects registration must not abort ``import spacecore``.""" + def boom(cls: type) -> None: + raise ValueError("duplicate registration") + + reg = PyTreeRegistry() + reg.register_class(_Alpha) + with pytest.warns(RuntimeWarning, match="duplicate registration"): + reg.register_backend("fake", boom) + + +def test_failing_registrar_is_not_retried(): + """The pair is marked done on failure, so there is no repeat attempt or spam.""" + calls: list[type] = [] + + def boom(cls: type) -> None: + calls.append(cls) + raise ValueError("nope") + + reg = PyTreeRegistry() + with pytest.warns(RuntimeWarning): + reg.register_backend("fake", boom) + reg.register_class(_Alpha) + reg.register_class(_Alpha) # second attempt must be a no-op + assert calls == [_Alpha] + + +# -------------------------------------------------------------------------- +# The ABCMeta gate: which subclasses auto-register +# -------------------------------------------------------------------------- + +def test_concrete_subclass_is_auto_registered(): + class Concrete(PyTreeNode): + def __init__(self, x): + self.x = x + + def tree_flatten(self): + return (self.x,), None + + @classmethod + def tree_unflatten(cls, aux, children): + return cls(children[0]) + + assert Concrete in global_registry.registered_classes() + + +def test_subclass_without_the_protocol_is_not_registered(): + """A subclass that has not implemented ``tree_flatten`` must be skipped.""" + class StillAbstract(PyTreeNode): + pass + + assert StillAbstract not in global_registry.registered_classes() + + +def test_subclass_with_protocol_but_abstract_elsewhere_is_registered(): + """Flattenable-but-abstract-for-other-reasons still registers. + + Pins the semantics of the gate: it asks "is this class flattenable yet?", + not "is it fully concrete?". ``inspect.isabstract`` would answer the latter + and skip this class, but a registrar only records the type and consults the + two protocol methods, both of which are present here. + """ + class Extra(PyTreeNode): + @abstractmethod + def apply(self): # unrelated abstract method keeps the class abstract + ... + + def tree_flatten(self): + return (), None + + @classmethod + def tree_unflatten(cls, aux, children): + return cls() + + assert inspect.isabstract(Extra) # genuinely still abstract + assert Extra in global_registry.registered_classes() # yet registered + + +def test_pytree_node_itself_is_not_registered(): + assert PyTreeNode not in global_registry.registered_classes() + + +def test_abstract_subclass_cannot_be_instantiated(): + class StillAbstract(PyTreeNode): + pass + + with pytest.raises(TypeError): + StillAbstract() + + +# -------------------------------------------------------------------------- +# Mixin composition +# -------------------------------------------------------------------------- + +def test_init_subclass_chain_is_cooperative(): + """PyTreeNode must not swallow ``__init_subclass__`` for sibling mixins.""" + observed: list[str] = [] + + class OtherMixin: + def __init_subclass__(cls, **kwargs): + super().__init_subclass__(**kwargs) + observed.append(cls.__name__) + + class Combined(PyTreeNode, OtherMixin): + def tree_flatten(self): + return (), None + + @classmethod + def tree_unflatten(cls, aux, children): + return cls() + + assert observed == ["Combined"] # sibling hook still ran + assert Combined in global_registry.registered_classes() # and ours did too + + +# -------------------------------------------------------------------------- +# Completeness against the real backend +# -------------------------------------------------------------------------- + +@pytest.mark.skipif(not has_jax(), reason="jax is not installed") +def test_every_registered_class_reached_jax(): + """Each class the registry recorded really is a JAX pytree node. + + This is the guard the old decorator could not offer: it kept no list of what + it had decorated, so "did anything get missed?" was unanswerable. The probe + is JAX's own public API — re-registering an already-registered type raises. + """ + import jax + + import spacecore # noqa: F401 - ensure every container module has imported + + recorded = global_registry.registered_classes() + assert len(recorded) > 40, f"suspiciously few containers registered: {len(recorded)}" + + unreached = [] + for cls in recorded: + try: + jax.tree_util.register_pytree_node_class(cls) + except ValueError: + continue # already registered - the expected outcome + unreached.append(f"{cls.__module__}.{cls.__qualname__}") + assert not unreached, f"recorded but never reached JAX: {unreached}" + + +@pytest.mark.skipif(not has_jax(), reason="jax is not installed") +def test_jax_protocol_is_installed_on_import(): + import spacecore # noqa: F401 + + assert "jax" in global_registry.backends() + + +@pytest.mark.skipif(not has_jax(), reason="jax is not installed") +def test_registered_class_round_trips_through_jax_tree_util(): + """A PyTreeNode subclass decomposes and rebuilds via ``jax.tree_util``. + + Re-homed from the retired ``test_jax_pytree_class.py``: the capability it + covered is unchanged, only the mechanism that grants it. + """ + import jax.tree_util as jtu + + class JaxPair(PyTreeNode): + def __init__(self, x, y): + self.x, self.y = x, y + + def tree_flatten(self): + return (self.x, self.y), None + + @classmethod + def tree_unflatten(cls, aux, children): + return cls(*children) + + def __eq__(self, other): + return isinstance(other, JaxPair) and (self.x, self.y) == (other.x, other.y) + + inst = JaxPair(1.0, 2.0) + leaves, treedef = jtu.tree_flatten(inst) + assert tuple(leaves) == (1.0, 2.0) # decomposed, not an opaque leaf + assert jtu.tree_unflatten(treedef, leaves) == inst + + +@pytest.mark.skipif(not has_jax(), reason="jax is not installed") +def test_registered_class_supports_tree_map(): + """``tree_map`` reaches the leaves of a registered container.""" + import jax.tree_util as jtu + + class JaxVec3(PyTreeNode): + def __init__(self, a, b, c): + self.a, self.b, self.c = a, b, c + + def tree_flatten(self): + return (self.a, self.b, self.c), None + + @classmethod + def tree_unflatten(cls, aux, children): + return cls(*children) + + doubled = jtu.tree_map(lambda v: v * 2.0, JaxVec3(1.0, 2.0, 3.0)) + assert (doubled.a, doubled.b, doubled.c) == (2.0, 4.0, 6.0) + + +@pytest.mark.skipif(not has_torch(), reason="torch is not installed") +def test_torch_protocol_is_installed_on_import(): + import spacecore # noqa: F401 + + assert "torch" in global_registry.backends() + + +@pytest.mark.skipif(not has_torch(), reason="torch is not installed") +def test_every_registered_class_reached_torch(): + """Torch's registry received every container — the second-backend guard. + + Mirrors the JAX completeness check; together they show the seam scales to a + second backend with no per-class edits. + """ + import torch.utils._pytree as torch_pytree + + import spacecore # noqa: F401 + + unreached = [] + for cls in global_registry.registered_classes(): + try: + torch_pytree.register_pytree_node( + cls, lambda o: ([], None), lambda children, aux: None + ) + except ValueError: + continue # already registered - the expected outcome + unreached.append(f"{cls.__module__}.{cls.__qualname__}") + assert not unreached, f"recorded but never reached torch: {unreached}" + + +@pytest.mark.skipif(not has_torch(), reason="torch is not installed") +def test_torch_backed_container_round_trips_and_maps(): + """A real Torch-backed operator decomposes, rebuilds, and maps. + + Exercises the adapter's translation: Torch wants ``list`` children and calls + ``unflatten(children, context)`` — the reverse of the ``(aux, children)`` + order :class:`PyTreeNode` defines. + """ + import torch + import torch.utils._pytree as torch_pytree + + import spacecore as sc + + ctx = sc.Context(sc.TorchOps(), dtype=torch.float64) + space = sc.DenseCoordinateSpace((3,), ctx=ctx) + op = sc.DenseLinOp(torch.eye(3, dtype=torch.float64), space, space) + + leaves, spec = torch_pytree.tree_flatten(op) + assert leaves and torch.is_tensor(leaves[0]) # decomposed, not opaque + assert torch_pytree.tree_unflatten(leaves, spec) == op + + doubled = torch_pytree.tree_map(lambda t: t * 2 if torch.is_tensor(t) else t, op) + assert torch.allclose(doubled.to_matrix(), 2 * torch.eye(3, dtype=torch.float64)) + + +@pytest.mark.skipif(not has_torch(), reason="torch is not installed") +def test_torch_registration_also_covers_cxx_pytree(): + """One registration serves both Torch pytree implementations. + + Torch mirrors registrations into ``_cxx_pytree`` (the optree-backed variant + selected by ``PYTORCH_USE_CXX_PYTREE=1``), so the adapter registers once — + registering with both explicitly would raise on the second call. + """ + import torch + import torch.utils._cxx_pytree as cxx_pytree + + import spacecore as sc + + ctx = sc.Context(sc.TorchOps(), dtype=torch.float64) + space = sc.DenseCoordinateSpace((3,), ctx=ctx) + op = sc.DenseLinOp(torch.eye(3, dtype=torch.float64), space, space) + + leaves, _ = cxx_pytree.tree_flatten(op) + assert leaves and torch.is_tensor(leaves[0]) + + +@pytest.mark.skipif(not has_jax(), reason="jax is not installed") +def test_backend_ops_are_deliberately_not_containers(): + """``*Ops`` classes are opaque leaves, not pytree nodes — by design. + + Replaces a test that asserted the opposite in its docstring yet passed + vacuously: ``tree_map`` returns an *unregistered* object unchanged because + it is a leaf, so the old assertion held whether or not registration existed. + """ + import jax.tree_util as jtu + import spacecore as sc + + ops = sc.NumpyOps() + assert type(ops) not in global_registry.registered_classes() + leaves, _ = jtu.tree_flatten(ops) + assert leaves == [ops] # exactly one leaf: itself + + +def test_round_trip_contract(): + class Pair(PyTreeNode): + def __init__(self, x, y): + self.x, self.y = x, y + + def tree_flatten(self): + return (self.x, self.y), "meta" + + @classmethod + def tree_unflatten(cls, aux, children): + assert aux == "meta" + return cls(*children) + + def __eq__(self, other): + return isinstance(other, Pair) and (self.x, self.y) == (other.x, other.y) + + original = Pair(1.0, 2.0) + children, aux = original.tree_flatten() + assert Pair.tree_unflatten(aux, children) == original diff --git a/tests/bench/test_bench_smoke.py b/tests/bench/test_bench_smoke.py index 7e843cb..d18fa84 100644 --- a/tests/bench/test_bench_smoke.py +++ b/tests/bench/test_bench_smoke.py @@ -144,8 +144,8 @@ def test_bare_inputs_are_backend_native(name, backend): def test_none_check_level_skips_membership_checks(monkeypatch): import spacecore as sc - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - space = sc.DenseCoordinateSpace((3,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + space = sc.DenseCoordinateSpace((3,), ctx, check_level="none") monkeypatch.setattr( space, "_check_member", @@ -158,8 +158,8 @@ def test_none_check_level_skips_membership_checks(monkeypatch): def test_checked_space_methods_match_unchecked_cores(): import spacecore as sc - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="cheap") - space = sc.DenseCoordinateSpace((3,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + space = sc.DenseCoordinateSpace((3,), ctx, check_level="cheap") x = ctx.asarray([1.0, 2.0, 3.0]) y = ctx.asarray([4.0, 5.0, 6.0]) np.testing.assert_allclose(space.add(x, y), space._add_core(x, y)) @@ -170,8 +170,8 @@ def test_checked_space_methods_match_unchecked_cores(): def test_checked_dense_linop_methods_match_unchecked_cores(): import spacecore as sc - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="cheap") - space = sc.DenseCoordinateSpace((2,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + space = sc.DenseCoordinateSpace((2,), ctx, check_level="cheap") matrix = ctx.asarray([[2.0, 1.0], [0.0, 3.0]]) op = sc.DenseLinOp(matrix, space, space, ctx) x = ctx.asarray([1.0, 2.0]) @@ -554,6 +554,8 @@ def test_run_probes_max_size_can_eliminate_every_case(monkeypatch): def test_run_probes_builds_cases_in_each_declared_check_level(monkeypatch): + import spacecore as sc + from bench._operations import _backend_ctx from bench._probes import Probe, ProbeCase from bench._run import run_probes @@ -562,7 +564,7 @@ def test_run_probes_builds_cases_in_each_declared_check_level(monkeypatch): def factory(backend, seed, size): ctx = _backend_ctx(backend) - built_levels.append(ctx.check_level) + built_levels.append(sc.get_check_level()) x = ctx.asarray([1.0]) return ProbeCase( bare_label="x + x", diff --git a/tests/conftest.py b/tests/conftest.py index 93849fa..412561f 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -2,6 +2,24 @@ import sys from pathlib import Path +import pytest + ROOT = Path(__file__).resolve().parents[1] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) + + +@pytest.fixture(autouse=True) +def _restore_ambient_check_level(): + """Restore the ambient validation level after each test. + + ``check_level`` moved off :class:`Context` onto the bound object; a test may + set the process-wide ambient default via ``sc.set_check_level(...)`` to build + objects at a given level. This fixture snapshots and restores it so that + setting cannot leak between tests (avoids the Erratic Test smell). + """ + import spacecore as sc + + previous = sc.get_check_level() + yield + sc.set_check_level(previous) diff --git a/tests/context/_contracts.py b/tests/context/_contracts.py new file mode 100644 index 0000000..b8410fc --- /dev/null +++ b/tests/context/_contracts.py @@ -0,0 +1,101 @@ +"""Reusable custom assertions for ``spacecore.contextual`` contract tests. + +These are xUnit *Custom Assertion* / *Test Utility Method* helpers (Meszaros, +*xUnit Test Patterns*, 2007, ch. 21): domain-specific assertions that check one +named property and report *which* property failed, so the individual tests stay +free of *Assertion Roulette* and *Test Code Duplication* (same source, test +smells). + +The relation/equality contracts encoded here are the standard value-object +contracts (reflexive / symmetric / transitive, hash-consistency) that +*Design by Contract* asks us to test as a contract rather than an +implementation (Hunt & Thomas, *The Pragmatic Programmer*, tip 37). +""" +from __future__ import annotations + +from typing import Any, Callable, Sequence + + +def assert_equivalence_relation( + rel: Callable[[Any, Any], bool], + *, + equivalent: Sequence[Any], + unrelated: Sequence[Any], +) -> None: + """Assert ``rel`` is an equivalence relation over the given samples. + + Parameters + ---------- + rel: + Binary predicate under test (e.g. ``Context.same_math``). + equivalent: + Two or more items that must all be pairwise related (one equivalence + class). Their pairwise-relatedness exercises transitivity. + unrelated: + Items that must *not* be related to the equivalence class, so the + relation is shown to discriminate rather than return ``True`` always. + """ + equivalent = list(equivalent) + unrelated = list(unrelated) + assert len(equivalent) >= 2, "need >= 2 equivalent samples to test transitivity" + everyone = equivalent + unrelated + + # Reflexive. + for x in everyone: + assert rel(x, x), f"not reflexive: rel(x, x) is False for {x!r}" + + # Symmetric. + for i, a in enumerate(everyone): + for b in everyone[i + 1 :]: + assert rel(a, b) == rel(b, a), f"not symmetric for {a!r} and {b!r}" + + # Every pair inside the class is related (=> transitivity within the class). + for i, a in enumerate(equivalent): + for b in equivalent[i + 1 :]: + assert rel(a, b), f"expected equivalent but rel is False: {a!r} ~ {b!r}" + + # The class is discriminated from the unrelated items. + ref = equivalent[0] + for u in unrelated: + assert not rel(ref, u), f"expected unrelated but rel is True: {ref!r} ~ {u!r}" + + +def assert_equality_contract( + *, + equal: Sequence[Any], + distinct: Sequence[Any], +) -> None: + """Assert the Python equality/hash contract for a value object. + + Checks reflexivity, symmetry, hash-consistency for the ``equal`` group, and + inequality against every item in ``distinct`` and against foreign types. + + Parameters + ---------- + equal: + Two or more objects that must compare ``==`` and hash-equal (an + "equal but distinct instance" set — build them independently rather + than aliasing one object). + distinct: + Objects that must each compare ``!=`` to the first ``equal`` item. + """ + equal = list(equal) + distinct = list(distinct) + assert len(equal) >= 2, "need >= 2 independently-built equal instances" + a = equal[0] + + # Reflexive. + assert a == a, f"not reflexive: {a!r} != itself" + + # Symmetric equality + hash-consistency across the equal group. + for b in equal[1:]: + assert a == b and b == a, f"equality not symmetric for {a!r} and {b!r}" + assert hash(a) == hash(b), f"hash inconsistent for equal {a!r} and {b!r}" + + # Inequality against genuinely different contexts, both directions. + for d in distinct: + assert a != d and d != a, f"expected inequality: {a!r} vs {d!r}" + + # Foreign types compare unequal, never raise. + assert (a == object()) is False, "equality with a foreign object should be False" + assert (a == None) is False # noqa: E711 - explicitly exercising __eq__(None) diff --git a/tests/context/conftest.py b/tests/context/conftest.py index 3477d26..30d1b4e 100644 --- a/tests/context/conftest.py +++ b/tests/context/conftest.py @@ -4,6 +4,11 @@ state into one another. The most common leak is :func:`spacecore.set_context` mutating the default context; without an explicit reset, a later test sees an unexpected default and fails spuriously. + +Book justification: ``preserve_default_context`` is a Fresh Fixture with +automated teardown — the standard remedy for the Erratic Test smell that shared +process-global state produces, keeping tests independent and order-insensitive +(Meszaros, *xUnit Test Patterns*). """ from __future__ import annotations diff --git a/tests/context/test_ambient_scoping.py b/tests/context/test_ambient_scoping.py new file mode 100644 index 0000000..96f1438 --- /dev/null +++ b/tests/context/test_ambient_scoping.py @@ -0,0 +1,216 @@ +"""Tests for the scoped ambient state behind ``use_context`` / ``use_check_level``. + +The ambient default context and validation level are *swappable* global state. +They used to be plain module globals mutated by save-and-restore, which is the +Global Data smell in its most dangerous form (Fowler, *Refactoring* 2e, ch. 3: +"Global data is especially nasty when it's mutable"). They now live in +:class:`contextvars.ContextVar`, so an override is scoped to the current thread / +async task and unwinds through a ``Token``. + +Note the asymmetry these tests pin, which is the whole point of the design: the +*registry* half of ``Contextual`` (``register_ops`` / ``available_ops``) is +monotonic and stays process-wide, while only the *policy* half is scoped. +Scoping the registry would be a bug — a backend registered inside a ``with`` +block would vanish on exit. + +Concurrency assertions use an explicit :class:`threading.Barrier` rather than +sleeps so the overlap is forced, not raced for (Meszaros, *xUnit Test Patterns*: +avoid the Erratic Test smell). Each test states one property. +""" +from __future__ import annotations + +import asyncio +import threading +from contextvars import copy_context + +import numpy as np +import pytest + +import spacecore as sc +from spacecore.backend import ops_registry + + +def _ctx(dtype): + """A context distinguishable by dtype alone, so no second backend is needed.""" + return sc.Context(sc.NumpyOps(), dtype=dtype) + + +F32 = np.dtype(np.float32) +F64 = np.dtype(np.float64) + + +# --------------------------------------------------------------------------- +# Nesting and unwinding +# --------------------------------------------------------------------------- +class TestNesting: + def test_check_level_unwinds_to_each_enclosing_level(self): + outer = sc.get_check_level() + with sc.use_check_level("strict"): + assert sc.get_check_level() == "strict" + with sc.use_check_level("none"): + assert sc.get_check_level() == "none" + assert sc.get_check_level() == "strict" + assert sc.get_check_level() == outer + + def test_context_unwinds_to_each_enclosing_context(self): + outer = sc.get_context() + with sc.use_context(_ctx(F32)): + assert sc.get_context().dtype == F32 + with sc.use_context(_ctx(F64)): + assert sc.get_context().dtype == F64 + assert sc.get_context().dtype == F32 + assert sc.get_context() == outer + + def test_override_unwinds_on_exception(self): + outer = sc.get_check_level() + with pytest.raises(RuntimeError): + with sc.use_check_level("strict"): + raise RuntimeError("boom") + assert sc.get_check_level() == outer + + +# --------------------------------------------------------------------------- +# ``set_*`` moves the baseline; ``use_*`` installs a scoped override +# --------------------------------------------------------------------------- +class TestBaselineVersusOverride: + def test_set_check_level_does_not_disturb_an_active_override(self): + with sc.use_check_level("strict"): + sc.set_check_level("none") # moves the baseline only + assert sc.get_check_level() == "strict" + assert sc.get_check_level() == "none" # ...revealed once the scope exits + + def test_set_context_from_a_worker_thread_is_visible_in_main(self): + """``set_context`` is process-wide by contract, unlike ``use_context``.""" + previous = sc.get_context() + try: + t = threading.Thread(target=lambda: sc.set_context(_ctx(F32))) + t.start() + t.join() + assert sc.get_context().dtype == F32 + finally: + sc.set_context(previous) + + def test_use_context_in_a_worker_thread_is_not_visible_in_main(self): + started, release = threading.Event(), threading.Event() + + def worker(): + with sc.use_context(_ctx(F32)): + started.set() + release.wait(timeout=5) + + t = threading.Thread(target=worker) + t.start() + try: + assert started.wait(timeout=5) + assert sc.get_context().dtype != F32 # main keeps the baseline + finally: + release.set() + t.join() + + +# --------------------------------------------------------------------------- +# Isolation between concurrent scopes — the bug this design fixes +# --------------------------------------------------------------------------- +class TestConcurrentIsolation: + def test_threads_do_not_clobber_each_others_check_level(self): + """Two overlapping scopes, with the interleaving pinned by events rather + than left to the scheduler: + + A enters "strict" -> B enters "none" -> A reads -> B reads + + Over a shared global, A's read at step 3 returns B's "none". Each thread + now reads its own ``ContextVar``, so both see what they set. + """ + seen: dict[str, str] = {} + a_in, b_in, a_read = threading.Event(), threading.Event(), threading.Event() + + def thread_a() -> None: + with sc.use_check_level("strict"): + a_in.set() + assert b_in.wait(timeout=5) # B is now inside its own scope + seen["a"] = sc.get_check_level() + a_read.set() + + def thread_b() -> None: + assert a_in.wait(timeout=5) # A is already inside its scope + with sc.use_check_level("none"): + b_in.set() + assert a_read.wait(timeout=5) # hold the scope open past A's read + seen["b"] = sc.get_check_level() + + threads = [threading.Thread(target=thread_a), threading.Thread(target=thread_b)] + for t in threads: + t.start() + for t in threads: + t.join(timeout=10) + + assert seen == {"a": "strict", "b": "none"} + + def test_interleaved_coroutines_keep_their_own_check_level(self): + """Each asyncio task copies the ambient context at creation, so suspending + inside a ``use_check_level`` block cannot clobber a sibling task.""" + + async def task(level: str, out: dict, key: str) -> None: + with sc.use_check_level(level): + await asyncio.sleep(0) # suspend *inside* the scope + out[key] = sc.get_check_level() + + async def main() -> dict: + out: dict[str, str] = {} + await asyncio.gather(task("strict", out, "x"), task("none", out, "y")) + return out + + outer = sc.get_check_level() + assert asyncio.run(main()) == {"x": "strict", "y": "none"} + assert sc.get_check_level() == outer # and nothing leaked out + + +# --------------------------------------------------------------------------- +# Thread non-inheritance is deliberate, and has a documented escape hatch +# --------------------------------------------------------------------------- +class TestThreadInheritance: + def test_spawned_thread_sees_the_baseline_not_the_enclosing_override(self): + """Pinned so the documented behaviour cannot change silently: a scoped + override must not leak into threads spawned inside it. See the caveat in + ``use_context``'s docstring.""" + box: dict[str, str] = {} + baseline = sc.get_check_level() + with sc.use_check_level("strict"): + t = threading.Thread(target=lambda: box.update(v=sc.get_check_level())) + t.start() + t.join() + assert box["v"] == baseline + + def test_copy_context_propagates_the_override_on_purpose(self): + """The documented way to opt in, for thread pools.""" + box: dict[str, str] = {} + with sc.use_check_level("strict"): + snapshot = copy_context() + t = threading.Thread( + target=lambda: snapshot.run(lambda: box.update(v=sc.get_check_level())) + ) + t.start() + t.join() + assert box["v"] == "strict" + + +# --------------------------------------------------------------------------- +# The registry half stays process-wide +# --------------------------------------------------------------------------- +def test_backend_registry_is_not_scoped_by_use_context(): + """``available_ops`` is monotonic state, not ambient policy: entering and + leaving a scope must not add or remove backends.""" + + before = dict(ops_registry.classes()) + with sc.use_context(_ctx(F32)): + assert dict(ops_registry.classes()) == before + assert dict(ops_registry.classes()) == before + + +def test_scoped_context_is_honoured_by_object_construction(): + """The override must reach every resolution path, not just ``get_context``: + ``normalize_context(None)`` inside ``ContextBound.__init__`` reads through the + same property.""" + with sc.use_context(_ctx(F32)): + space = sc.DenseCoordinateSpace((2,)) + assert space.ctx.dtype == F32 diff --git a/tests/context/test_check_policy.py b/tests/context/test_check_policy.py index 1d793fc..8b8e6be 100644 --- a/tests/context/test_check_policy.py +++ b/tests/context/test_check_policy.py @@ -1,40 +1,54 @@ +"""Behavioral tests for what each ``check_level`` actually enforces. + +Drives real spaces, linear operators, functionals, and solvers to pin which +invariants each level (``none``/``cheap``/``standard``/``strict``) checks or +skips end-to-end. The level is a property of the context-bound object, so every +case sets it through the object's ``check_level=`` keyword rather than through +the :class:`spacecore.Context`. + +Book justification: ``check_level`` is a *contract* about which invariants are +enforced, so it is tested as a contract rather than by inspecting internals +(Hunt & Thomas, *The Pragmatic Programmer*, tip 37). Strict levels must assert +states believed impossible while lenient levels must not — assertive +programming / fail-fast (same source, tips 38-39; Ousterhout, *A Philosophy of +Software Design*, "when to crash"). Numerical precision is treated as part of +correctness across levels (Irving et al., *Research Software Engineering with +Python*). +""" import numpy as np import pytest import spacecore as sc +from spacecore._check_policy import normalize_check_level -def _ctx(level: sc.CheckLevel, dtype=np.float64) -> sc.Context: - return sc.Context(sc.NumpyOps(), dtype=dtype, check_level=level) +def _ctx(dtype=np.float64) -> sc.Context: + return sc.Context(sc.NumpyOps(), dtype=dtype) def test_check_level_public_api_and_legacy_mapping(): + ctx = _ctx() + assert sc.CHECK_LEVELS == ("none", "cheap", "standard", "strict") - assert sc.Context(sc.NumpyOps()).check_level == "standard" - assert sc.Context(sc.NumpyOps(), check_level="cheap").check_level == "cheap" - assert sc.normalize_context("numpy", check_level="strict").check_level == "strict" + assert sc.DenseCoordinateSpace((2,), ctx).check_level == "standard" + assert sc.DenseCoordinateSpace((2,), ctx, check_level="cheap").check_level == "cheap" + assert sc.DenseCoordinateSpace((2,), ctx, check_level="strict").check_level == "strict" + assert normalize_check_level(enable_checks=True) == "standard" + assert normalize_check_level(enable_checks=False) == "none" with pytest.warns(DeprecationWarning, match="enable_checks"): - checked = sc.Context(sc.NumpyOps(), enable_checks=True) - with pytest.warns(DeprecationWarning, match="enable_checks"): - unchecked = sc.Context(sc.NumpyOps(), enable_checks=False) - - assert checked.check_level == "standard" - assert checked.enable_checks is True - assert unchecked.check_level == "none" - assert unchecked.enable_checks is False + assert normalize_check_level(enable_checks=True, warn_legacy=True) == "standard" with pytest.raises(TypeError, match="either check_level or enable_checks"): - sc.Context(sc.NumpyOps(), enable_checks=True, check_level="strict") + normalize_check_level("strict", enable_checks=True) with pytest.raises(ValueError, match="Unknown check_level"): - sc.Context(sc.NumpyOps(), check_level="fast") + normalize_check_level("fast") -def test_inferred_context_uses_the_least_expensive_source_level(): - strict_ctx = _ctx("strict") - cheap_ctx = _ctx("cheap") - strict_space = sc.DenseCoordinateSpace((1,), strict_ctx) - cheap_space = sc.DenseCoordinateSpace((1,), cheap_ctx) +def test_derived_object_uses_the_least_expensive_source_level(): + ctx = _ctx() + strict_space = sc.DenseCoordinateSpace((1,), ctx, check_level="strict") + cheap_space = sc.DenseCoordinateSpace((1,), ctx, check_level="cheap") product = sc.TreeSpace.from_leaf_spaces((strict_space, cheap_space)) @@ -42,9 +56,9 @@ def test_inferred_context_uses_the_least_expensive_source_level(): def test_none_skips_optional_space_linop_and_batched_checks(): - ctx = _ctx("none") - space = sc.DenseCoordinateSpace((2,), ctx) - identity = sc.IdentityLinOp(space, ctx) + ctx = _ctx() + space = sc.DenseCoordinateSpace((2,), ctx, check_level="none") + identity = sc.IdentityLinOp(space, ctx, check_level="none") invalid = ctx.asarray([1.0, 2.0, 3.0]) invalid_batch = ctx.asarray([[1.0, 2.0, 3.0]]) @@ -54,8 +68,8 @@ def test_none_skips_optional_space_linop_and_batched_checks(): def test_cheap_checks_shape_dtype_backend_and_tree_structure_only(): - ctx = _ctx("cheap", np.float32) - vector = sc.DenseCoordinateSpace((2,), ctx) + ctx = _ctx(np.float32) + vector = sc.DenseCoordinateSpace((2,), ctx, check_level="cheap") with pytest.raises(sc.SpaceValidationError, match="Expected shape"): vector.check_member(np.asarray([1.0, 2.0, 3.0], dtype=np.float32)) @@ -71,21 +85,20 @@ def test_cheap_checks_shape_dtype_backend_and_tree_structure_only(): def test_standard_adds_recursive_and_hermitian_membership(): - ctx = _ctx("standard") - vector = sc.DenseCoordinateSpace((2,), ctx) + ctx = _ctx() + vector = sc.DenseCoordinateSpace((2,), ctx, check_level="standard") product = sc.TreeSpace.from_leaf_spaces((vector, vector), ctx) with pytest.raises(sc.SpaceValidationError, match=r"\$\[0\]"): product.check_member((ctx.asarray([1.0]), ctx.asarray([2.0, 3.0]))) - hermitian = sc.HermitianSpace(2, ctx=ctx) + hermitian = sc.HermitianSpace(2, ctx=ctx, check_level="standard") with pytest.raises(sc.SpaceValidationError, match="not Hermitian"): hermitian.check_member(ctx.asarray([[1.0, 2.0], [0.0, 1.0]])) - cheap_ctx = _ctx("cheap") - cheap_hermitian = sc.HermitianSpace(2, ctx=cheap_ctx) - cheap_product = sc.TreeSpace.from_leaf_spaces((cheap_hermitian,), cheap_ctx) - cheap_product.check_member((cheap_ctx.asarray([[1.0, 2.0], [0.0, 1.0]]),)) + cheap_hermitian = sc.HermitianSpace(2, ctx=ctx, check_level="cheap") + cheap_product = sc.TreeSpace.from_leaf_spaces((cheap_hermitian,), ctx) + cheap_product.check_member((ctx.asarray([[1.0, 2.0], [0.0, 1.0]]),)) standard_product = sc.TreeSpace.from_leaf_spaces((hermitian,), ctx) with pytest.raises(sc.SpaceValidationError, match=r"\$\[0\].*not Hermitian"): @@ -93,9 +106,9 @@ def test_standard_adds_recursive_and_hermitian_membership(): def test_checked_method_and_batched_validation_follow_cheap_policy(): - ctx = _ctx("cheap") - space = sc.DenseCoordinateSpace((2,), ctx) - identity = sc.IdentityLinOp(space, ctx) + ctx = _ctx() + space = sc.DenseCoordinateSpace((2,), ctx, check_level="cheap") + identity = sc.IdentityLinOp(space, ctx, check_level="cheap") with pytest.raises(sc.SpaceValidationError, match="Expected shape"): identity.apply(ctx.asarray([1.0, 2.0, 3.0])) @@ -104,52 +117,58 @@ def test_checked_method_and_batched_validation_follow_cheap_policy(): def test_functional_scalar_output_shape_is_standard(): - cheap_ctx = _ctx("cheap") - cheap_space = sc.DenseCoordinateSpace((2,), cheap_ctx) + ctx = _ctx() + cheap_space = sc.DenseCoordinateSpace((2,), ctx, check_level="cheap") cheap_functional = sc.MatrixFreeLinearFunctional( - lambda _x: cheap_ctx.asarray([1.0]), cheap_space, cheap_ctx + lambda _x: ctx.asarray([1.0]), cheap_space, ctx, check_level="cheap" ) - assert cheap_functional.value(cheap_ctx.asarray([1.0, 2.0])).shape == (1,) + assert cheap_functional.value(ctx.asarray([1.0, 2.0])).shape == (1,) - standard_ctx = _ctx("standard") - standard_space = sc.DenseCoordinateSpace((2,), standard_ctx) + standard_space = sc.DenseCoordinateSpace((2,), ctx, check_level="standard") standard_functional = sc.MatrixFreeLinearFunctional( - lambda _x: standard_ctx.asarray([1.0]), standard_space, standard_ctx + lambda _x: ctx.asarray([1.0]), standard_space, ctx, check_level="standard" ) - with pytest.raises(ValueError, match="scalar batch output"): - standard_functional.value(standard_ctx.asarray([1.0, 2.0])) + with pytest.raises(ValueError, match="scalar output"): + standard_functional.value(ctx.asarray([1.0, 2.0])) def test_strict_matrix_free_adjoint_probe_is_strict_only(): - standard_ctx = _ctx("standard") - standard_space = sc.DenseCoordinateSpace((2,), standard_ctx) + ctx = _ctx() + standard_space = sc.DenseCoordinateSpace((2,), ctx, check_level="standard") sc.MatrixFreeLinOp( lambda x: x, - lambda y: standard_ctx.asarray([0.0, 0.0]), + lambda y: ctx.asarray([0.0, 0.0]), standard_space, standard_space, - standard_ctx, + ctx, + check_level="standard", ) - strict_ctx = _ctx("strict") - strict_space = sc.DenseCoordinateSpace((2,), strict_ctx) + strict_space = sc.DenseCoordinateSpace((2,), ctx, check_level="strict") with pytest.raises(ValueError, match="adjoint consistency check failed"): sc.MatrixFreeLinOp( lambda x: x, - lambda y: strict_ctx.asarray([0.0, 0.0]), + lambda y: ctx.asarray([0.0, 0.0]), strict_space, strict_space, - strict_ctx, + ctx, + check_level="strict", ) def test_strict_matrix_free_coordinate_adjoint_preserves_non_euclidean_metric(): - ctx = _ctx("strict") + ctx = _ctx() domain = sc.DenseCoordinateSpace( - (2,), ctx, geometry=sc.WeightedInnerProduct(ctx.asarray([2.0, 3.0])) + (2,), + ctx, + geometry=sc.WeightedInnerProduct(ctx.asarray([2.0, 3.0])), + check_level="strict", ) codomain = sc.DenseCoordinateSpace( - (2,), ctx, geometry=sc.WeightedInnerProduct(ctx.asarray([5.0, 7.0])) + (2,), + ctx, + geometry=sc.WeightedInnerProduct(ctx.asarray([5.0, 7.0])), + check_level="strict", ) op = sc.MatrixFreeLinOp.from_coordinate_adjoint( lambda x: x, @@ -164,24 +183,22 @@ def test_strict_matrix_free_coordinate_adjoint_preserves_non_euclidean_metric(): def test_linalg_keeps_square_invariant_and_adds_strict_cg_probe(): - none_ctx = _ctx("none") - domain = sc.DenseCoordinateSpace((2,), none_ctx) - codomain = sc.DenseCoordinateSpace((3,), none_ctx) - rectangular = sc.ZeroLinOp(domain, codomain, none_ctx) + ctx = _ctx() + domain = sc.DenseCoordinateSpace((2,), ctx, check_level="none") + codomain = sc.DenseCoordinateSpace((3,), ctx, check_level="none") + rectangular = sc.ZeroLinOp(domain, codomain, ctx, check_level="none") with pytest.raises(ValueError, match="square LinOp"): - sc.cg(rectangular, none_ctx.asarray([1.0, 1.0, 1.0]), maxiter=0) + sc.cg(rectangular, ctx.asarray([1.0, 1.0, 1.0]), maxiter=0) - standard_ctx = _ctx("standard") - standard_space = sc.DenseCoordinateSpace((2,), standard_ctx) + standard_space = sc.DenseCoordinateSpace((2,), ctx, check_level="standard") standard_negative = sc.DiagonalLinOp( - standard_ctx.asarray([-1.0, -1.0]), standard_space, standard_ctx + ctx.asarray([-1.0, -1.0]), standard_space, ctx, check_level="standard" ) - sc.cg(standard_negative, standard_ctx.asarray([1.0, 1.0]), maxiter=0) + sc.cg(standard_negative, ctx.asarray([1.0, 1.0]), maxiter=0) - strict_ctx = _ctx("strict") - strict_space = sc.DenseCoordinateSpace((2,), strict_ctx) + strict_space = sc.DenseCoordinateSpace((2,), ctx, check_level="strict") strict_negative = sc.DiagonalLinOp( - strict_ctx.asarray([-1.0, -1.0]), strict_space, strict_ctx + ctx.asarray([-1.0, -1.0]), strict_space, ctx, check_level="strict" ) with pytest.raises(ValueError, match="positive curvature"): - sc.cg(strict_negative, strict_ctx.asarray([1.0, 1.0]), maxiter=0) + sc.cg(strict_negative, ctx.asarray([1.0, 1.0]), maxiter=0) diff --git a/tests/context/test_check_policy_helpers.py b/tests/context/test_check_policy_helpers.py index 8ef0644..2712b70 100644 --- a/tests/context/test_check_policy_helpers.py +++ b/tests/context/test_check_policy_helpers.py @@ -1,5 +1,13 @@ """Tests for :mod:`spacecore._check_policy` pure helper functions. +Book justification: these are pure functions, pinned directly against their +contract rather than only through ``Context`` (Hunt & Thomas, *The Pragmatic +Programmer*, tip 37). The ordering and ``>=`` relations are covered as +exhaustive Parameterized truth tables to avoid Conditional Test Logic and +Assertion Roulette (Meszaros, *xUnit Test Patterns*); normalization that +resolves rather than raises reflects "define errors out of existence" +(Ousterhout, *A Philosophy of Software Design*). + These functions back the public ``check_level`` policy but are otherwise only exercised transitively through :class:`spacecore.Context`. This module unit tests them directly. diff --git a/tests/context/test_checked_method.py b/tests/context/test_checked_method.py index a7a0660..5ca9486 100644 --- a/tests/context/test_checked_method.py +++ b/tests/context/test_checked_method.py @@ -1,5 +1,13 @@ """Tests for :func:`spacecore.checked_method`. +Book justification: validation paths run rarely and are easy to leave +under-tested, so each is exercised explicitly (Ousterhout, *A Philosophy of +Software Design*: exceptions add rarely-run paths that must still be tested). +The decorator is driven through a recording Test Double so the wrapper is +isolated from any real receiver (Meszaros, *xUnit Test Patterns*), and the +metadata-preservation cases pin an API-consistency contract (Myers & Stylos, +*Improving API Usability*). + The decorator wraps a method so that selected positional arguments are validated against an input space and the return value against an output space, gated by the receiver's ``check_level`` policy. @@ -101,10 +109,10 @@ class _BatchedDemo: """Receiver using real spaces so the ``_check_batched`` branch fires.""" def __init__(self, check_level="cheap"): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level=check_level) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) self.ctx = ctx - self.dom = sc.DenseCoordinateSpace((3,), ctx) - self.cod = sc.DenseCoordinateSpace((3,), ctx) + self.dom = sc.DenseCoordinateSpace((3,), ctx, check_level=check_level) + self.cod = sc.DenseCoordinateSpace((3,), ctx, check_level=check_level) # The decorator reads ``self.check_level`` for gating. self.check_level = check_level self.out_result = ctx.asarray([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]]) diff --git a/tests/context/test_compatibility.py b/tests/context/test_compatibility.py index fea01a8..8b1c93b 100644 --- a/tests/context/test_compatibility.py +++ b/tests/context/test_compatibility.py @@ -1,33 +1,42 @@ """Tests for context compatibility / inference helpers. +Book justification: the resolver relations and inference are pinned as +contracts (Hunt & Thomas, *The Pragmatic Programmer*, tip 37). The synthetic +``_OtherFamilyOps`` is a Test-Specific Subclass standing in for a second backend +family, so cross-family behavior is covered without depending on an installed +optional backend — keeping the suite Repeatable and free of Erratic Tests +(Meszaros, *xUnit Test Patterns*). + Checklist section 2: context compatibility and inference. -Covers the free helpers that gate operator algebra and context resolution: +Covers the free helpers that gate operator algebra and context resolution. +(The ``Context.same_math`` / ``same_backend`` relation contracts live in +``test_context_contracts.py``; this file covers the resolver helpers.) -* :func:`spacecore._contextual._bound._same_math_context` — the - algebra-gating equality that ignores ``check_level``. * ``Contextual.are_compatible_values`` / ``are_compatible_ops`` — the family-mismatch logic for raw values and raw ``BackendOps``. * ``Contextual.infer_context`` / ``infer_contexts`` — the ``.ctx`` fast path, ``is_array`` matching, the ``get_dtype`` fallback, and the no-match ``None`` branch. -* ``Contextual.ctx_from_ops`` — dtype sanitization and check-level - normalization for a raw ``BackendOps`` instance. -* :func:`spacecore.normalize_context` — the deprecated ``enable_checks`` - legacy path and its ``DeprecationWarning``. +* ``Contextual.ctx_from_ops`` — dtype sanitization for a raw ``BackendOps`` + instance. Validation policy is not part of the context: the ambient level + seeds the bound object instead. +* ``normalize_check_level`` — the deprecated ``enable_checks`` legacy shim, + now the only surviving home of the Boolean switch, and the + ``DeprecationWarning`` it emits under ``warn_legacy=True``. References are independent: NumPy dtypes, explicit family strings, and the source contracts read from ``_state.py`` / ``_bound.py`` / ``_check_policy.py``. """ from __future__ import annotations -import warnings - import numpy as np +import pytest import spacecore as sc -from spacecore._contextual._bound import _same_math_context -from spacecore._contextual._state import Contextual +from spacecore._check_policy import normalize_check_level +from spacecore.contextual._contextual import Contextual +from spacecore.backend import OpsRegistry # A NumpyOps subclass with a distinct family. Backend equality and the @@ -38,33 +47,6 @@ class _OtherFamilyOps(sc.NumpyOps): _family = "other_family" -# =========================================================================== -# _same_math_context — gates operator algebra; ignores check_level -# =========================================================================== -class TestSameMathContext: - def test_differ_only_in_check_level_is_same(self): - """Two contexts differing ONLY in ``check_level`` share math context.""" - a = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - b = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="strict") - assert a != b # full equality is sensitive to check_level - assert _same_math_context(a, b) is True - - def test_differ_in_dtype_is_not_same(self): - a = sc.Context(sc.NumpyOps(), dtype=np.float32) - b = sc.Context(sc.NumpyOps(), dtype=np.float64) - assert _same_math_context(a, b) is False - - def test_differ_in_ops_family_is_not_same(self): - a = sc.Context(sc.NumpyOps(), dtype=np.float64) - b = sc.Context(_OtherFamilyOps(), dtype=np.float64) - assert _same_math_context(a, b) is False - - def test_identical_context_is_same(self): - a = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="cheap") - b = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="cheap") - assert _same_math_context(a, b) is True - - # =========================================================================== # are_compatible_ops — raw BackendOps family logic # =========================================================================== @@ -127,8 +109,8 @@ def test_context_argument_returns_itself(self): def test_ctx_fast_path(self): """An object exposing ``.ctx`` resolves via that attribute directly.""" - ctx = sc.Context(sc.NumpyOps(), dtype=np.float32, check_level="cheap") - bound = sc.DenseCoordinateSpace((2,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float32) + bound = sc.DenseCoordinateSpace((2,), ctx, check_level="cheap") out = self.state.infer_context(bound) assert out == ctx @@ -163,10 +145,10 @@ def get_dtype(self, x): class _FakeArray: dtype = np.dtype(np.float32) - state = Contextual() - # Replace the registry with just the fake backend so it is the sole - # match for the fake array. - state._available_ops = {"fake_only": _FakeOps} + # Resolve against a registry holding only the fake backend, so it is + # the sole match for the fake array. Injecting a registry is the + # supported seam; this used to overwrite private state. + state = Contextual(ops_registry=OpsRegistry([_FakeOps])) captured["x"] = _FakeArray() out = state.infer_context(captured["x"]) assert out is not None @@ -198,7 +180,7 @@ def test_all_uninferrable_yields_empty_tuple(self): # =========================================================================== -# ctx_from_ops — dtype sanitization + check-level normalization +# ctx_from_ops — dtype sanitization; policy lives on the bound object # =========================================================================== class TestCtxFromOps: def setup_method(self): @@ -214,14 +196,17 @@ def test_explicit_dtype_is_sanitized(self): out = self.state.ctx_from_ops(ops, dtype=np.float32) assert out.dtype == np.dtype(np.float32) - def test_default_check_level_is_none(self): - """``Contextual._default_check_level`` is 'none'.""" + def test_baseline_check_level_is_standard(self): + """``Contextual._baseline_check_level`` is 'standard', and it is what a + bound object built on a resolver-made context inherits.""" + assert self.state.get_check_level() == "standard" out = self.state.ctx_from_ops(sc.NumpyOps()) - assert out.check_level == "none" + assert sc.DenseCoordinateSpace((2,), out).check_level == "standard" def test_explicit_check_level_is_honored(self): - out = self.state.ctx_from_ops(sc.NumpyOps(), check_level="strict") - assert out.check_level == "strict" + out = self.state.ctx_from_ops(sc.NumpyOps()) + bound = sc.DenseCoordinateSpace((2,), out, check_level="strict") + assert bound.check_level == "strict" def test_returned_ops_match_input(self): ops = sc.NumpyOps() @@ -230,25 +215,19 @@ def test_returned_ops_match_input(self): # =========================================================================== -# normalize_context — legacy enable_checks path + DeprecationWarning +# normalize_check_level — legacy enable_checks shim + DeprecationWarning # =========================================================================== -class TestNormalizeContextEnableChecks: +class TestEnableChecksShim: def test_enable_checks_true_resolves_standard(self): - with warnings.catch_warnings(): - warnings.simplefilter("ignore", DeprecationWarning) - out = sc.normalize_context("numpy", enable_checks=True) - assert out.check_level == "standard" + assert normalize_check_level(enable_checks=True) == "standard" def test_enable_checks_false_resolves_none(self): - with warnings.catch_warnings(): - warnings.simplefilter("ignore", DeprecationWarning) - out = sc.normalize_context("numpy", enable_checks=False) - assert out.check_level == "none" - - def test_enable_checks_emits_deprecation_warning(self): - with warnings.catch_warnings(record=True) as caught: - warnings.simplefilter("always") - sc.normalize_context("numpy", enable_checks=True) - deprecations = [w for w in caught if issubclass(w.category, DeprecationWarning)] - assert deprecations, "expected a DeprecationWarning for enable_checks" - assert "enable_checks is deprecated" in str(deprecations[0].message) + assert normalize_check_level(enable_checks=False) == "none" + + def test_enable_checks_emits_deprecation_warning_when_requested(self): + """The shim warns only under ``warn_legacy=True``; no production caller + passes it, so the warning is opt-in for callers still bridging the + Boolean switch.""" + with pytest.warns(DeprecationWarning) as caught: + normalize_check_level(enable_checks=True, warn_legacy=True) + assert "enable_checks is deprecated" in str(caught[0].message) diff --git a/tests/context/test_context_bound.py b/tests/context/test_context_bound.py index 075afb0..5be480d 100644 --- a/tests/context/test_context_bound.py +++ b/tests/context/test_context_bound.py @@ -1,4 +1,11 @@ -"""Tests for :class:`spacecore._contextual.ContextBound`. +"""Tests for :class:`spacecore.contextual.ContextBound`. + +Book justification: the abstract base's contract is exercised through a minimal +Test-Specific Subclass rather than a production subclass, isolating the SUT from +concrete implementations (Meszaros, *xUnit Test Patterns*). Property delegation +and ``convert`` dispatch are pinned as a contract (Hunt & Thomas, *The Pragmatic +Programmer*, tip 37), and the ``_checks_at_least`` truth table is a Parameterized +Test to keep one reported concern per case (avoid Assertion Roulette). ``ContextBound`` is the abstract base of every object that lives in a SpaceCore ``Context`` — spaces, linear operators, functionals. The tests @@ -10,6 +17,8 @@ * ``ctx`` property returns the bound ``Context`` * ``ops`` property delegates to ``ctx.ops`` +* ``check_level`` is a property of the *object*, seeded at construction and + independent of the backend/dtype context it is bound to * ``convert(new_ctx)`` is idempotent for the same context and dispatches to ``_convert`` otherwise * the subclass ``_convert`` hook is invoked @@ -23,7 +32,7 @@ import pytest import spacecore as sc -from spacecore._contextual import ContextBound +from spacecore.contextual import ContextBound class _ToyBound(ContextBound): @@ -32,8 +41,12 @@ class _ToyBound(ContextBound): Records every ``_convert`` invocation so the test can confirm dispatch. """ - def __init__(self, ctx: sc.Context | str | None = None) -> None: - super().__init__(ctx) + def __init__( + self, + ctx: sc.Context | str | None = None, + check_level: sc.CheckLevel | None = None, + ) -> None: + super().__init__(ctx, check_level) self._convert_calls: list[sc.Context] = [] def _convert(self, new_ctx: sc.Context) -> Self: @@ -73,13 +86,29 @@ def test_dtype_property_delegates_to_ctx(self): bound = _ToyBound(ctx) assert bound.dtype == ctx.dtype - def test_check_level_property_delegates_to_ctx(self): - ctx = sc.Context(sc.NumpyOps(), check_level="cheap") - bound = _ToyBound(ctx) + def test_check_level_property_is_seeded_from_the_constructor(self): + ctx = sc.Context(sc.NumpyOps()) + bound = _ToyBound(ctx, check_level="cheap") assert bound.check_level == "cheap" + def test_check_level_defaults_to_the_ambient_level(self): + ctx = sc.Context(sc.NumpyOps()) + with sc.use_check_level("strict"): + bound = _ToyBound(ctx) + assert bound.check_level == "strict" + + def test_check_level_is_independent_of_the_context(self): + """Two objects on one context may validate at different strictness, and + they still share a math context.""" + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + cheap = _ToyBound(ctx, check_level="cheap") + strict = _ToyBound(ctx, check_level="strict") + assert cheap.check_level == "cheap" + assert strict.check_level == "strict" + assert cheap.same_math(strict) is True + def test_default_init_uses_active_default_context(self, preserve_default_context): - explicit = sc.Context(sc.NumpyOps(), dtype=np.float32, check_level="cheap") + explicit = sc.Context(sc.NumpyOps(), dtype=np.float32) sc.set_context(explicit) bound = _ToyBound() assert bound.ctx == explicit @@ -111,18 +140,17 @@ class TestChecksAtLeast: ("strict", "strict", True), ]) def test_truth_table(self, current, required, expected): - ctx = sc.Context(sc.NumpyOps(), check_level=current) - bound = _ToyBound(ctx) + ctx = sc.Context(sc.NumpyOps()) + bound = _ToyBound(ctx, check_level=current) assert bound._checks_at_least(required) is expected def test_enable_checks_property_is_legacy_view(self): """``ContextBound._enable_checks`` is the legacy bool view of ``check_level``: True for anything other than 'none'.""" + ctx = sc.Context(sc.NumpyOps()) for level in ("cheap", "standard", "strict"): - ctx = sc.Context(sc.NumpyOps(), check_level=level) - assert _ToyBound(ctx)._enable_checks is True - ctx = sc.Context(sc.NumpyOps(), check_level="none") - assert _ToyBound(ctx)._enable_checks is False + assert _ToyBound(ctx, check_level=level)._enable_checks is True + assert _ToyBound(ctx, check_level="none")._enable_checks is False # --------------------------------------------------------------------------- @@ -159,15 +187,25 @@ def test_convert_to_different_dtype_dispatches_to__convert(self): def test_convert_accepts_family_string(self): """``convert("numpy")`` resolves the string through ``normalize_context``. - Even though the resulting context is structurally a numpy context, - check_level differences make it distinct enough to dispatch. + The family string carries no dtype, so it resolves to the backend's + default (float64) — distinct from the float32 context below, hence a + real dispatch rather than the identity short-circuit. """ - ctx_strict = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="strict") - bound = _ToyBound(ctx_strict) + ctx_f32 = sc.Context(sc.NumpyOps(), dtype=np.float32) + bound = _ToyBound(ctx_f32) out = bound.convert("numpy") - # The default check_level is 'standard' so the contexts differ. assert out is not bound assert bound._convert_calls and bound._convert_calls[0].ops.family == "numpy" + assert out.dtype == sc.NumpyOps().sanitize_dtype(None) + + def test_convert_preserves_the_objects_check_level(self): + """``check_level`` is a property of the object, not of the context, so + it survives a conversion onto a different context.""" + ctx_a = sc.Context(sc.NumpyOps(), dtype=np.float32) + ctx_b = sc.Context(sc.NumpyOps(), dtype=np.float64) + bound = _ToyBound(ctx_a, check_level="strict") + out = bound.convert(ctx_b) + assert out.check_level == "strict" def test_convert_round_trip_returns_to_original_ctx(self): ctx_a = sc.Context(sc.NumpyOps(), dtype=np.float32) diff --git a/tests/context/test_context_contracts.py b/tests/context/test_context_contracts.py new file mode 100644 index 0000000..afc5ab5 --- /dev/null +++ b/tests/context/test_context_contracts.py @@ -0,0 +1,133 @@ +"""Contract tests for ``Context`` equality and its coarsening relations. + +Covers the value-object contracts of ``spacecore.contextual._context.Context``: +``__eq__``/``__hash__`` and the relations ``same_math`` (ops & dtype) and +``same_backend`` (ops). ``check_level`` is no longer part of the context — it is +a property of the bound object — so ``Context`` identity is purely mathematical +and ``same_math`` now coincides with ``__eq__``; only ``same_backend`` (which +drops the dtype) is a strictly coarser relation. + +Book justification: the relations are tested *as contracts* — reflexive / +symmetric / transitive — rather than by poking at fields (Design by Contract; +Hunt & Thomas, *The Pragmatic Programmer*, tip 37). The shared custom assertions +keep each test to one reported concern (Meszaros, *xUnit Test Patterns*: Custom +Assertion, avoid Assertion Roulette). +""" +from __future__ import annotations + +import dataclasses + +import numpy as np +import pytest + +import spacecore as sc + +from ._contracts import assert_equality_contract, assert_equivalence_relation + + +# A NumpyOps subclass with a distinct family, so cross-family cases need no +# optional backend installed. +class _OtherFamilyOps(sc.NumpyOps): + _family = "other_family" + + +def _ctx(dtype=np.float64, ops=None): + return sc.Context(ops or sc.NumpyOps(), dtype=dtype) + + +# --------------------------------------------------------------------------- +# __eq__ / __hash__ as a value-object contract +# --------------------------------------------------------------------------- +class TestEqualityContract: + def test_full_equality_contract(self): + assert_equality_contract( + equal=[_ctx(), _ctx()], + distinct=[ + _ctx(dtype=np.float32), # dtype differs + _ctx(ops=_OtherFamilyOps()), # backend family differs + ], + ) + + def test_frozen_is_immutable(self): + ctx = _ctx() + with pytest.raises(dataclasses.FrozenInstanceError): + ctx.dtype = np.float32 # type: ignore[misc] + + def test_usable_as_dict_key(self): + a, b = _ctx(), _ctx() + assert {a: "v"}[b] == "v" + + +# --------------------------------------------------------------------------- +# same_math (ops & dtype) — now coincides with __eq__ +# --------------------------------------------------------------------------- +class TestSameMathRelation: + def test_is_equivalence_relation(self): + assert_equivalence_relation( + lambda a, b: a.same_math(b), + equivalent=[_ctx(), _ctx()], + unrelated=[_ctx(dtype=np.float32), _ctx(ops=_OtherFamilyOps())], + ) + + def test_coincides_with_equality(self): + a, b = _ctx(dtype=np.float32), _ctx(dtype=np.float64) + assert a.same_math(b) is False and (a == b) is False + c, d = _ctx(), _ctx() + assert c.same_math(d) is True and (c == d) is True + + def test_non_context_is_false(self): + assert _ctx().same_math(object()) is False + assert _ctx().same_math(None) is False + + +# --------------------------------------------------------------------------- +# same_backend (ops only) — strictly coarser than same_math +# --------------------------------------------------------------------------- +class TestSameBackendRelation: + def test_is_equivalence_relation(self): + assert_equivalence_relation( + lambda a, b: a.same_backend(b), + equivalent=[ + _ctx(dtype=np.float64), + _ctx(dtype=np.float32), + _ctx(dtype=np.complex128), + ], + unrelated=[_ctx(ops=_OtherFamilyOps())], + ) + + def test_non_context_is_false(self): + assert _ctx().same_backend(object()) is False + assert _ctx().same_backend(None) is False + + +# --------------------------------------------------------------------------- +# The coarsening chain: __eq__ ⊆ same_math ⊆ same_backend +# --------------------------------------------------------------------------- +class TestCoarseningChain: + CASES = [ + ("identical", _ctx(), _ctx(), True, True, True), + ("dtype_differs", _ctx(dtype=np.float32), _ctx(dtype=np.float64), False, False, True), + ("family_differs", _ctx(ops=_OtherFamilyOps()), _ctx(), False, False, False), + ] + + @pytest.mark.parametrize( + "a,b,eq,sm,sb", + [(a, b, eq, sm, sb) for (_n, a, b, eq, sm, sb) in CASES], + ids=[n for (n, *_r) in CASES], + ) + def test_relation_values(self, a, b, eq, sm, sb): + assert (a == b) is eq + assert a.same_math(b) is sm + assert a.same_backend(b) is sb + + @pytest.mark.parametrize( + "a,b", + [(a, b) for (_n, a, b, *_r) in CASES], + ids=[n for (n, *_r) in CASES], + ) + def test_implications_hold(self, a, b): + # Coarsening: __eq__ ⟹ same_math ⟹ same_backend, for every case. + if a == b: + assert a.same_math(b) + if a.same_math(b): + assert a.same_backend(b) diff --git a/tests/context/test_context_resolution_policy.py b/tests/context/test_context_resolution_policy.py new file mode 100644 index 0000000..da7a5c4 --- /dev/null +++ b/tests/context/test_context_resolution_policy.py @@ -0,0 +1,102 @@ +"""Policy tests for context resolution: dtype promotion and surviving checks. + +Two properties of how operands are combined: + +1. **dtype promotion never narrows.** These operators run on ill-conditioned + problems (condition number growing like 1/ε), where the smallest scales sit + below float32's machine epsilon, so silently downcasting a float64 operand + to float32 would destroy precision the result depends on. Treating dtype / + numerical precision as part of correctness — not an incidental — is the + guidance in Irving et al., *Research Software Engineering with Python* + (numerical-correctness testing). + +2. **Surviving check-policy after combination should not depend on operand + order.** + +The dtype cases use a *Parameterized Test* with an independent literal expected +table (Meszaros, *xUnit Test Patterns*, ch. 27 Literal/Derived Value: do not +re-implement the production promotion rule inside the test). +""" +from __future__ import annotations + +import numpy as np +import pytest + +import spacecore as sc +from spacecore.contextual._contextual import Contextual + + +def _resolver(): + return Contextual() + + +def _space(dtype, check_level="standard"): + with sc.use_check_level(check_level): + return sc.DenseCoordinateSpace((2,), sc.Context(sc.NumpyOps(), dtype=dtype)) + + +# --------------------------------------------------------------------------- +# dtype promotion never narrows +# --------------------------------------------------------------------------- +class TestDtypePromotionNeverNarrows: + # (name, dtype_a, dtype_b, expected) — expected is an independent literal, + # not a call to the production promotion routine. + PAIRS = [ + ("f32_f32", np.float32, np.float32, np.float32), + ("f64_f64", np.float64, np.float64, np.float64), + ("f32_f64", np.float32, np.float64, np.float64), + ("f64_f32", np.float64, np.float32, np.float64), # order-independent + ("f32_c128", np.float32, np.complex128, np.complex128), + ] + + @pytest.mark.parametrize( + "da,db,expected", + [(da, db, exp) for (_n, da, db, exp) in PAIRS], + ids=[n for (n, *_r) in PAIRS], + ) + def test_promotes_to_expected(self, da, db, expected): + st = _resolver() + ctx = st.resolve_context_priority(None, np.zeros(2, da), np.zeros(2, db)) + assert ctx.dtype == np.dtype(expected) + + @pytest.mark.parametrize( + "da,db", + [(da, db) for (_n, da, db, _e) in PAIRS], + ids=[n for (n, *_r) in PAIRS], + ) + def test_result_is_never_narrower_than_inputs(self, da, db): + # The book property, independent of the exact promotion table: the + # resolved dtype's itemsize is at least each input's (no precision loss). + st = _resolver() + ctx = st.resolve_context_priority(None, np.zeros(2, da), np.zeros(2, db)) + widest = max(np.dtype(da).itemsize, np.dtype(db).itemsize) + assert np.dtype(ctx.dtype).itemsize >= widest + + def test_promotion_is_order_independent(self): + st = _resolver() + ab = st.resolve_context_priority(None, np.zeros(2, np.float32), np.zeros(2, np.float64)) + ba = st.resolve_context_priority(None, np.zeros(2, np.float64), np.zeros(2, np.float32)) + assert ab.dtype == ba.dtype + + +# --------------------------------------------------------------------------- +# Surviving check-policy is operand-order independent (critique §2.6) +# --------------------------------------------------------------------------- +class TestCheckLevelOrderIndependence: + def test_leaf_objects_take_check_level_from_ambient(self): + # check_level is a property of the bound object, seeded from the ambient + # default (or an explicit scope), not from the Context. + assert _space(np.float64, check_level="none").check_level == "none" + assert _space(np.float64, check_level="strict").check_level == "strict" + + def test_linop_algebra_surviving_check_level_is_order_independent(self): + # check_level is combined at the object level via the minimum rule + # (ContextBound), so A+B and B+A agree and equal the least-strict operand + # rather than whichever came first. + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + with sc.use_check_level("none"): + a = sc.DiagonalLinOp(np.ones(2), ctx=ctx) + with sc.use_check_level("strict"): + b = sc.DiagonalLinOp(np.ones(2), ctx=ctx) + assert a.check_level == "none" and b.check_level == "strict" + assert (a + b).check_level == (b + a).check_level == "none" diff --git a/tests/context/test_enable_checks.py b/tests/context/test_enable_checks.py index 6a2ed56..5b0c181 100644 --- a/tests/context/test_enable_checks.py +++ b/tests/context/test_enable_checks.py @@ -1,24 +1,50 @@ +"""Behavioral tests for the deprecated Boolean ``enable_checks`` policy shim. + +The Boolean switch no longer appears on any constructor: it survives only as +:func:`spacecore._check_policy.normalize_check_level`'s ``enable_checks`` +keyword, which maps ``True``/``False`` onto the ``"standard"``/``"none"`` +levels carried by context-bound objects. These tests drive spaces and linear +operators built at exactly those two mapped levels, so the legacy switch is +still pinned end-to-end: it must keep selecting an enforcing and a +non-enforcing policy respectively. + +Book justification: a deprecated, rarely-exercised path rots silently, so its +backward-compatibility behavior is pinned with explicit regression tests +(Ousterhout, *A Philosophy of Software Design*: rarely-run paths are +under-tested; Hunt & Thomas, *The Pragmatic Programmer*, tip 94: make preserved +behavior a lasting test). The rejection cases also check that invalid inputs +fail with meaningful feedback (Myers & Stylos, *Improving API Usability*: help +users recognize and recover from errors). +""" import numpy as np import pytest import spacecore as sc +from spacecore._check_policy import normalize_check_level from tests._helpers import has_jax, jax_real_dtype -def _checked_ctx(dtype=np.float64): - return sc.Context(sc.NumpyOps(), dtype=dtype, enable_checks=True) +CHECKED = normalize_check_level(enable_checks=True) +UNCHECKED = normalize_check_level(enable_checks=False) -def _unchecked_ctx(dtype=np.float64): - return sc.Context(sc.NumpyOps(), dtype=dtype, enable_checks=False) +def _ctx(dtype=np.float64): + return sc.Context(sc.NumpyOps(), dtype=dtype) + + +def test_enable_checks_maps_onto_an_enforcing_and_a_skipping_level(): + assert CHECKED == "standard" + assert UNCHECKED == "none" def test_enable_checks_accepts_valid_space_and_linop_inputs(): - ctx = _checked_ctx() - dom = sc.DenseCoordinateSpace((2,), ctx) - cod = sc.DenseCoordinateSpace((3,), ctx) - op = sc.DenseLinOp(ctx.asarray([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]), dom, cod, ctx) + ctx = _ctx() + dom = sc.DenseCoordinateSpace((2,), ctx, check_level=CHECKED) + cod = sc.DenseCoordinateSpace((3,), ctx, check_level=CHECKED) + op = sc.DenseLinOp( + ctx.asarray([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]), dom, cod, ctx, check_level=CHECKED + ) x = ctx.asarray([7.0, 8.0]) y = op.apply(x) @@ -29,16 +55,16 @@ def test_enable_checks_accepts_valid_space_and_linop_inputs(): def test_enable_checks_rejects_vector_shape_mismatch(): - ctx = _checked_ctx(np.float32) - space = sc.DenseCoordinateSpace((2,), ctx) + ctx = _ctx(np.float32) + space = sc.DenseCoordinateSpace((2,), ctx, check_level=CHECKED) with pytest.raises(TypeError, match=r"Expected shape \(2,\), got \(3,\)"): space.check_member(np.asarray([1.0, 2.0, 3.0], dtype=np.float32)) def test_enable_checks_rejects_vector_dtype_mismatch(): - ctx = _checked_ctx(np.float32) - space = sc.DenseCoordinateSpace((2,), ctx) + ctx = _ctx(np.float32) + space = sc.DenseCoordinateSpace((2,), ctx, check_level=CHECKED) with pytest.raises(TypeError, match=r"Expected dtype float32, got float64"): space.check_member(np.asarray([1.0, 2.0], dtype=np.float64)) @@ -46,27 +72,31 @@ def test_enable_checks_rejects_vector_dtype_mismatch(): @pytest.mark.skipif(not has_jax(), reason="jax is not installed") def test_enable_checks_rejects_cross_backend_dense_array(): - np_ctx = _checked_ctx(jax_real_dtype()) - jx_ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), enable_checks=True) - space = sc.DenseCoordinateSpace((2,), np_ctx) + np_ctx = _ctx(jax_real_dtype()) + jx_ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) + space = sc.DenseCoordinateSpace((2,), np_ctx, check_level=CHECKED) with pytest.raises(TypeError, match="Expected dense array for numpy"): space.check_member(jx_ctx.asarray([1.0, 2.0])) def test_enable_checks_rejects_non_hermitian_matrix(): - ctx = _checked_ctx() - space = sc.HermitianSpace(2, ctx=ctx) + ctx = _ctx() + space = sc.HermitianSpace(2, ctx=ctx, check_level=CHECKED) with pytest.raises(TypeError, match="not Hermitian"): space.check_member(ctx.asarray([[1.0, 2.0], [0.0, 1.0]])) def test_enable_checks_rejects_invalid_tree_structure(): - ctx = _checked_ctx() + ctx = _ctx() product = sc.TreeSpace.from_leaf_spaces( - (sc.DenseCoordinateSpace((2,), ctx), sc.DenseCoordinateSpace((3,), ctx)), + ( + sc.DenseCoordinateSpace((2,), ctx, check_level=CHECKED), + sc.DenseCoordinateSpace((3,), ctx, check_level=CHECKED), + ), ctx, + check_level=CHECKED, ) with pytest.raises(TypeError, match="structure mismatch"): @@ -80,14 +110,18 @@ def test_enable_checks_rejects_invalid_tree_structure(): def test_enable_checks_rejects_dense_linop_matrix_and_vector_dimensions(): - ctx = _checked_ctx() - dom = sc.DenseCoordinateSpace((2,), ctx) - cod = sc.DenseCoordinateSpace((3,), ctx) + ctx = _ctx() + dom = sc.DenseCoordinateSpace((2,), ctx, check_level=CHECKED) + cod = sc.DenseCoordinateSpace((3,), ctx, check_level=CHECKED) with pytest.raises(TypeError, match=r"Expected A\.shape == cod\.shape \+ dom\.shape"): - sc.DenseLinOp(ctx.asarray([[1.0, 2.0], [3.0, 4.0]]), dom, cod, ctx) + sc.DenseLinOp( + ctx.asarray([[1.0, 2.0], [3.0, 4.0]]), dom, cod, ctx, check_level=CHECKED + ) - op = sc.DenseLinOp(ctx.asarray([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]), dom, cod, ctx) + op = sc.DenseLinOp( + ctx.asarray([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]), dom, cod, ctx, check_level=CHECKED + ) with pytest.raises(TypeError, match=r"Expected shape \(2,\), got \(3,\)"): op.apply(ctx.asarray([1.0, 2.0, 3.0])) @@ -96,12 +130,14 @@ def test_enable_checks_rejects_dense_linop_matrix_and_vector_dimensions(): def test_enable_checks_rejects_tree_linop_domain_codomain_mismatch(): - ctx = _checked_ctx() - dom2 = sc.DenseCoordinateSpace((2,), ctx) - dom3 = sc.DenseCoordinateSpace((3,), ctx) - cod1 = sc.DenseCoordinateSpace((1,), ctx) - first = sc.DenseLinOp(ctx.asarray([[1.0, 2.0]]), dom2, cod1, ctx) - second = sc.DenseLinOp(ctx.asarray([[1.0, 2.0, 3.0]]), dom3, cod1, ctx) + ctx = _ctx() + dom2 = sc.DenseCoordinateSpace((2,), ctx, check_level=CHECKED) + dom3 = sc.DenseCoordinateSpace((3,), ctx, check_level=CHECKED) + cod1 = sc.DenseCoordinateSpace((1,), ctx, check_level=CHECKED) + first = sc.DenseLinOp(ctx.asarray([[1.0, 2.0]]), dom2, cod1, ctx, check_level=CHECKED) + second = sc.DenseLinOp( + ctx.asarray([[1.0, 2.0, 3.0]]), dom3, cod1, ctx, check_level=CHECKED + ) with pytest.raises( TypeError, match=r"Component op 1 must map dom -> cod\.leaf_spaces\[1\]" @@ -110,18 +146,18 @@ def test_enable_checks_rejects_tree_linop_domain_codomain_mismatch(): def test_enable_checks_rejects_invalid_conversion_target(): - ctx = _checked_ctx() - space = sc.DenseCoordinateSpace((2,), ctx) + ctx = _ctx() + space = sc.DenseCoordinateSpace((2,), ctx, check_level=CHECKED) with pytest.raises(TypeError, match="Expected Context, BackendFamily, str, or None"): space.convert(object()) def test_disabled_checks_skip_space_membership_validations(): - ctx = _unchecked_ctx() - vector = sc.DenseCoordinateSpace((2,), ctx) - hermitian = sc.HermitianSpace(2, ctx=ctx) - product = sc.TreeSpace.from_leaf_spaces((vector, vector), ctx) + ctx = _ctx() + vector = sc.DenseCoordinateSpace((2,), ctx, check_level=UNCHECKED) + hermitian = sc.HermitianSpace(2, ctx=ctx, check_level=UNCHECKED) + product = sc.TreeSpace.from_leaf_spaces((vector, vector), ctx, check_level=UNCHECKED) vector.check_member(np.asarray([1.0, 2.0, 3.0], dtype=np.float32)) hermitian.check_member(ctx.asarray([[1.0, 2.0], [0.0, 1.0]])) diff --git a/tests/context/test_policies_errors.py b/tests/context/test_policies_errors.py index 102d674..c9d020c 100644 --- a/tests/context/test_policies_errors.py +++ b/tests/context/test_policies_errors.py @@ -1,14 +1,26 @@ -"""Tests for the :mod:`spacecore._contextual._policies` error hierarchy. +"""Tests for the context/backend error hierarchy (:mod:`spacecore._errors`). + +The types are defined at the top level because both :mod:`spacecore.backend` +(the ops registry) and :mod:`spacecore.contextual` raise them; they remain +public API of ``spacecore.contextual``, which is how these tests import them. + + +Book justification: exception design deserves deliberate tests — each error is +raised at its real trigger, sits at the correct (base-catchable) place in the +hierarchy, and carries a useful message (Ousterhout, *A Philosophy of Software +Design*: exceptions add complexity and must be tested and kept a small, stable +set; Myers & Stylos, *Improving API Usability*: help users recognize, diagnose, +and recover from errors). Four exception types live in this module: -* :class:`spacecore._contextual.ContextError` — base, subclass of +* :class:`spacecore.contextual.ContextError` — base, subclass of ``RuntimeError``; -* :class:`spacecore._contextual.ContextInferenceError` — context cannot be +* :class:`spacecore.contextual.ContextInferenceError` — context cannot be inferred from input (typically: ambiguous backend match); -* :class:`spacecore._contextual.ContextConflictError` — contradictory +* :class:`spacecore.contextual.ContextConflictError` — contradictory registrations or contexts (typically: duplicate ``register_ops``); -* :class:`spacecore._contextual.UnknownBackendError` — a family name that +* :class:`spacecore.contextual.UnknownBackendError` — a family name that was never registered. Each test pins both the hierarchy (``isinstance`` relationship) and the @@ -22,7 +34,8 @@ import pytest import spacecore as sc -from spacecore._contextual import ( +from spacecore.backend import ops_registry +from spacecore.contextual import ( ContextConflictError, ContextError, ContextInferenceError, @@ -93,7 +106,6 @@ class _EphemeralOps(sc.NumpyOps): class TestContextConflictError: def test_duplicate_register_ops_raises(self): - from spacecore._contextual._state import _state family = "test_duplicate_register_ops_raises" cls = _make_ephemeral_backend(family) @@ -102,10 +114,9 @@ def test_duplicate_register_ops_raises(self): with pytest.raises(ContextConflictError, match="already registered"): sc.register_ops(cls) finally: - _state().available_ops.pop(family, None) + ops_registry.unregister(family) def test_duplicate_message_includes_family(self): - from spacecore._contextual._state import _state family = "test_duplicate_message_includes_family" cls = _make_ephemeral_backend(family) @@ -115,7 +126,7 @@ def test_duplicate_message_includes_family(self): sc.register_ops(cls) assert family in str(exc_info.value) finally: - _state().available_ops.pop(family, None) + ops_registry.unregister(family) # =========================================================================== @@ -130,7 +141,6 @@ def test_ambiguous_inference_when_two_backends_claim_an_object(self): both the real numpy and our impostor in the registry, inferring a backend for a plain numpy array becomes ambiguous. """ - from spacecore._contextual._state import _state family = "test_ambiguous_inference_two_claimants" @@ -139,13 +149,15 @@ class _ClaimsNdarrayOps(sc.NumpyOps): _ClaimsNdarrayOps.__name__ = f"_Ephemeral_{family}_Ops" + from spacecore.contextual._state import _state + x = np.asarray([1.0, 2.0, 3.0]) try: sc.register_ops(_ClaimsNdarrayOps) with pytest.raises(ContextInferenceError, match="(?i)ambiguous"): _state().infer_context(x) finally: - _state().available_ops.pop(family, None) + ops_registry.unregister(family) # =========================================================================== diff --git a/tests/context/test_state_free_functions.py b/tests/context/test_state_free_functions.py index 4f6da41..c92c891 100644 --- a/tests/context/test_state_free_functions.py +++ b/tests/context/test_state_free_functions.py @@ -1,4 +1,11 @@ -"""Tests for the free-function API in :mod:`spacecore._contextual._state`. +"""Tests for the free-function API in :mod:`spacecore.contextual._state`. + +Book justification: the public free-function API is tested as its own first +client — defaults that do the right thing and consistent behavior across input +types (Myers & Stylos, *Improving API Usability*). Because these functions +mutate a process-wide default context, they use a Fresh Fixture with guaranteed +teardown so tests stay independent and order-insensitive — the remedy for the +Erratic Test smell (Meszaros, *xUnit Test Patterns*). The functions covered: @@ -22,9 +29,10 @@ import spacecore as sc from spacecore.backend import BackendFamily -from spacecore._contextual import UnknownBackendError +from spacecore.contextual import UnknownBackendError from tests._helpers import has_cupy, has_jax, has_torch +from spacecore.backend import ops_registry _OPTIONAL_BACKEND_PROBES = {"jax": has_jax, "torch": has_torch, "cupy": has_cupy} @@ -43,7 +51,7 @@ def test_default_context_is_numpy(self): assert ctx.ops.family == "numpy" def test_set_context_with_context_object(self, preserve_default_context): - target = sc.Context(sc.NumpyOps(), dtype=np.float32, check_level="cheap") + target = sc.Context(sc.NumpyOps(), dtype=np.float32) sc.set_context(target) assert sc.get_context() == target @@ -141,7 +149,7 @@ def test_unknown_string_raises(self): # =========================================================================== class TestNormalizeContext: def test_none_returns_active_default(self, preserve_default_context): - explicit = sc.Context(sc.NumpyOps(), dtype=np.float32, check_level="cheap") + explicit = sc.Context(sc.NumpyOps(), dtype=np.float32) sc.set_context(explicit) out = sc.normalize_context(None) assert out == explicit @@ -160,9 +168,12 @@ def test_family_string_with_explicit_dtype(self): out = sc.normalize_context("numpy", dtype=np.float32) assert out.dtype == np.dtype(np.float32) - def test_family_string_with_check_level(self): - out = sc.normalize_context("numpy", check_level="strict") - assert out.check_level == "strict" + def test_family_string_context_leaves_policy_to_the_bound_object(self): + """A normalized context carries backend and dtype only; the validation + policy is seeded on the object bound to it.""" + out = sc.normalize_context("numpy") + assert not hasattr(out, "check_level") + assert sc.DenseCoordinateSpace((2,), out, check_level="strict").check_level == "strict" def test_rejects_unknown_type(self): with pytest.raises(TypeError): @@ -177,10 +188,6 @@ def test_warns_when_none_provided_with_dtype_override(self, preserve_default_con with pytest.warns(UserWarning, match="ignored"): sc.normalize_context(None, dtype=np.float32) - def test_rejects_both_check_level_and_enable_checks(self): - with pytest.raises(TypeError, match="either check_level or enable_checks"): - sc.normalize_context("numpy", check_level="strict", enable_checks=True) - # =========================================================================== # register_ops @@ -205,19 +212,17 @@ class _EphemeralOps(sc.NumpyOps): class TestRegisterOps: def test_register_adds_new_family(self): - from spacecore._contextual._state import _state family = "test_register_adds_new_family_one" cls = _make_ephemeral_backend(family) try: sc.register_ops(cls) - assert family in _state().available_ops - assert _state().available_ops[family] is cls + assert family in ops_registry + assert ops_registry.classes()[family] is cls finally: - _state().available_ops.pop(family, None) + ops_registry.unregister(family) def test_registered_family_is_usable_via_set_context(self, preserve_default_context): - from spacecore._contextual._state import _state family = "test_registered_family_is_usable" cls = _make_ephemeral_backend(family) @@ -226,11 +231,10 @@ def test_registered_family_is_usable_via_set_context(self, preserve_default_cont sc.set_context(family) assert sc.get_context().ops.family == family finally: - _state().available_ops.pop(family, None) + ops_registry.unregister(family) def test_duplicate_registration_raises_context_conflict_error(self): - from spacecore._contextual import ContextConflictError - from spacecore._contextual._state import _state + from spacecore.contextual import ContextConflictError family = "test_duplicate_registration_raises" cls = _make_ephemeral_backend(family) @@ -239,7 +243,7 @@ def test_duplicate_registration_raises_context_conflict_error(self): with pytest.raises(ContextConflictError, match="already registered"): sc.register_ops(cls) finally: - _state().available_ops.pop(family, None) + ops_registry.unregister(family) def test_rejects_non_class_argument(self): with pytest.raises(TypeError, match="Expected type"): @@ -253,7 +257,6 @@ class NotABackend: sc.register_ops(NotABackend) # type: ignore[arg-type] def test_returns_the_registered_class(self): - from spacecore._contextual._state import _state family = "test_returns_the_registered_class" cls = _make_ephemeral_backend(family) @@ -261,7 +264,7 @@ def test_returns_the_registered_class(self): returned = sc.register_ops(cls) assert returned is cls finally: - _state().available_ops.pop(family, None) + ops_registry.unregister(family) # =========================================================================== @@ -269,26 +272,26 @@ def test_returns_the_registered_class(self): # =========================================================================== class TestResolveContextPriority: def test_default_used_when_no_inputs(self, preserve_default_context): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float32, check_level="cheap") + ctx = sc.Context(sc.NumpyOps(), dtype=np.float32) sc.set_context(ctx) out = sc.resolve_context_priority(None) assert out == ctx def test_explicit_overrides_inferred(self, preserve_default_context): sc.set_context(sc.Context(sc.NumpyOps(), dtype=np.float16)) - inferred = sc.Context(sc.NumpyOps(), dtype=np.float32, check_level="cheap") - explicit = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") + inferred = sc.Context(sc.NumpyOps(), dtype=np.float32) + explicit = sc.Context(sc.NumpyOps(), dtype=np.float64) X = sc.DenseCoordinateSpace((2,), inferred) out = sc.resolve_context_priority(explicit, X) assert out == explicit def test_inferred_used_when_explicit_is_none(self, preserve_default_context): sc.set_context(sc.Context(sc.NumpyOps(), dtype=np.float16)) - inferred = sc.Context(sc.NumpyOps(), dtype=np.float32, check_level="cheap") + inferred = sc.Context(sc.NumpyOps(), dtype=np.float32) X = sc.DenseCoordinateSpace((2,), inferred) out = sc.resolve_context_priority(None, X) assert out.dtype == inferred.dtype - assert out.check_level == inferred.check_level + assert out.ops.family == inferred.ops.family def test_default_used_when_no_inferred(self, preserve_default_context): ctx = sc.Context(sc.NumpyOps(), dtype=np.float32) @@ -306,13 +309,14 @@ def test_compatible_inferred_dtypes_are_promoted(self, preserve_default_context) out = sc.resolve_context_priority(None, Xa, Xb) assert out.dtype == np.dtype(np.float64) - def test_minimum_check_level_among_inferred_contexts(self, preserve_default_context): - strict = sc.Context(sc.NumpyOps(), check_level="strict") - cheap = sc.Context(sc.NumpyOps(), check_level="cheap") - Xs = sc.DenseCoordinateSpace((2,), strict) - Xc = sc.DenseCoordinateSpace((3,), cheap) - out = sc.resolve_context_priority(None, Xs, Xc) - assert out.check_level == "cheap" + def test_minimum_check_level_among_bound_sources(self, preserve_default_context): + """The resolved *context* carries no policy; the object built from the + sources takes the least expensive level among them.""" + ctx = sc.Context(sc.NumpyOps()) + Xs = sc.DenseCoordinateSpace((2,), ctx, check_level="strict") + Xc = sc.DenseCoordinateSpace((3,), ctx, check_level="cheap") + op = sc.ZeroLinOp(Xs, Xc) + assert op.check_level == "cheap" def test_incompatible_inferred_contexts_raise(self): if not has_jax(): diff --git a/tests/functional/test_algebra.py b/tests/functional/test_algebra.py index ad33cd9..58f4cc2 100644 --- a/tests/functional/test_algebra.py +++ b/tests/functional/test_algebra.py @@ -1,4 +1,11 @@ -"""Functional algebra: scalar multiples and sums (0.4.2 W4, mirrors LinOp algebra).""" +"""Functional algebra: scalar multiples, sums, shifts, constants and products. + +Per-node behavior for the lazy combinators in ``spacecore.functional._algebra`` +and the operator overloads that build them. The *shared* container contract +(pytree round-trip, equality, repr, ``value_and_grad``) is specified once for +every functional in ``test_functional_contract.py``; the guard clauses live in +``test_functional_guards.py``. Mirrors ``linop/_algebra.py`` where they overlap. +""" from __future__ import annotations import numpy as np @@ -73,7 +80,10 @@ def test_negation(self, numpy_ctx): def test_non_scalar_returns_notimplemented(self, numpy_ctx): _, F, _ = _dense(numpy_ctx) assert F.__mul__("x") is NotImplemented - assert F.__mul__(F) is NotImplemented # a functional is not a scalar multiplier + # A Functional operand is *not* rejected — it builds a pointwise product + # (see TestProduct). Only operands that are neither scalar-like nor + # Functional defer to the reflected operation. + assert F.__mul__(object()) is NotImplemented def test_type_and_scalar_guards(self, numpy_ctx): _, F, _ = _dense(numpy_ctx) @@ -85,10 +95,10 @@ def test_type_and_scalar_guards(self, numpy_ctx): def test_complex_scalar_conjugates_gradient(self): # Riesz gradient of a*F is conj(a)*grad(F): the inner product conjugates # its first argument, so must recover a * . - ctx = sc.Context(sc.NumpyOps(), dtype=np.complex128, check_level="standard") - X = sc.DenseCoordinateSpace((3,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.complex128) + X = sc.DenseCoordinateSpace((3,), ctx, check_level="standard") c = ctx.asarray([1 + 1j, 2 - 0.5j, -1 + 0.3j]) - F = sc.InnerProductFunctional(c, X, ctx) + F = sc.InnerProductFunctional(c, X, ctx, check_level="standard") a = 2 + 3j x = ctx.asarray([0.5 - 1j, 1 + 0j, -2 + 0.5j]) h = ctx.asarray([1 + 0j, 0 + 1j, 0.5 - 0.5j]) @@ -276,6 +286,175 @@ def test_pytree_round_trip(self, numpy_ctx): assert cls.tree_unflatten(aux, children) == obj +# =========================================================================== +# ConstantFunctional (the scalar embedding) + factory zero collapse +# =========================================================================== +class TestConstant: + def test_value_is_the_constant_everywhere(self, numpy_ctx): + X, _, x = _dense(numpy_ctx) + C = sc.ConstantFunctional(X, 2.5, numpy_ctx) + np.testing.assert_allclose(to_numpy(C.value(x)), 2.5) + np.testing.assert_allclose( + to_numpy(C.value(numpy_ctx.asarray([9.0, 9.0, 9.0]))), 2.5 + ) + + def test_gradient_is_zero(self, numpy_ctx): + X, _, x = _dense(numpy_ctx) + C = sc.ConstantFunctional(X, 2.5, numpy_ctx) + np.testing.assert_allclose(to_numpy(C.grad(x)), 0.0) + v, g = C.value_and_grad(x) + np.testing.assert_allclose(to_numpy(v), 2.5) + np.testing.assert_allclose(to_numpy(g), 0.0) + + def test_factory_collapses_zero_to_zero_functional(self, numpy_ctx): + """The additive identity keeps a single representation.""" + X, _, _ = _dense(numpy_ctx) + assert isinstance(sc.make_constant_functional(X, 0.0, numpy_ctx), sc.ZeroFunctional) + assert isinstance( + sc.make_constant_functional(X, 2.5, numpy_ctx), sc.ConstantFunctional + ) + + def test_rejects_non_scalar(self, numpy_ctx): + X, _, _ = _dense(numpy_ctx) + with pytest.raises(TypeError, match="scalar-like"): + sc.ConstantFunctional(X, "nope", numpy_ctx) + + def test_equality_and_pytree_round_trip(self, numpy_ctx): + X, _, _ = _dense(numpy_ctx) + C = sc.ConstantFunctional(X, 2.5, numpy_ctx) + assert C == sc.ConstantFunctional(X, 2.5, numpy_ctx) + assert C != sc.ConstantFunctional(X, 3.5, numpy_ctx) + children, aux = C.tree_flatten() + assert sc.ConstantFunctional.tree_unflatten(aux, children) == C + + +# =========================================================================== +# ProductFunctional (pointwise product) + product-rule gradient +# =========================================================================== +class TestProduct: + def test_value_is_the_pointwise_product(self, numpy_ctx): + X, F, x = _dense(numpy_ctx) + G = F * F + assert isinstance(G, sc.ProductFunctional) + np.testing.assert_allclose( + to_numpy(G.value(x)), to_numpy(F.value(x)) ** 2 + ) + + def test_gradient_follows_the_product_rule(self, numpy_ctx): + """``grad(F·F) = 2 F(x) grad F(x)`` for a real-valued ``F``.""" + X, F, x = _dense(numpy_ctx) + expected = 2.0 * to_numpy(F.value(x)) * to_numpy(F.grad(x)) + np.testing.assert_allclose(to_numpy((F * F).grad(x)), expected) + + def test_value_and_grad_matches_separate_calls(self, numpy_ctx): + X, F, x = _dense(numpy_ctx) + P = F * F + v, g = P.value_and_grad(x) + np.testing.assert_allclose(to_numpy(v), to_numpy(P.value(x))) + np.testing.assert_allclose(to_numpy(g), to_numpy(P.grad(x))) + + def test_gradient_matches_finite_differences(self, numpy_ctx): + """Independent check of the product rule on two *different* factors.""" + X, F, x = _dense(numpy_ctx) + G = sc.InnerProductFunctional(numpy_ctx.asarray([1.0, 0.5, -2.0]), X, numpy_ctx) + P = F * G + direction = numpy_ctx.asarray([0.25, -0.5, 0.75]) + eps = 1e-6 + finite_difference = ( + P.value(X.axpy(eps, direction, x)) - P.value(X.axpy(-eps, direction, x)) + ) / (2.0 * eps) + np.testing.assert_allclose( + to_numpy(X.inner(P.grad(x), direction)), + to_numpy(finite_difference), + rtol=1e-7, + atol=1e-7, + ) + + def test_tree_domain_gradient_uses_domain_ops(self, numpy_ctx): + """On a pytree domain the product rule must go through ``X.scale``/``X.add``. + + A raw ``*``/``+`` on the gradients would fail here, where an element is a + tuple of leaves rather than an array. + """ + X, F, x = _tree(numpy_ctx) + np.testing.assert_allclose( + to_numpy(X.flatten((F * F).grad(x))), + 2.0 * to_numpy(F.value(x)) * to_numpy(X.flatten(F.grad(x))), + ) + + def test_constant_factor_folds_to_a_scaled_functional(self, numpy_ctx): + """Multiplying by a constant *is* scaling — the product node is not built.""" + X, F, x = _dense(numpy_ctx) + C = sc.ConstantFunctional(X, 3.0, numpy_ctx) + for product in (C * F, F * C): + assert isinstance(product, sc.ScaledFunctional) + np.testing.assert_allclose( + to_numpy(product.value(x)), 3.0 * to_numpy(F.value(x)) + ) + + def test_zero_factor_collapses_to_zero(self, numpy_ctx): + X, F, _ = _dense(numpy_ctx) + Z = sc.ZeroFunctional(X, numpy_ctx) + assert isinstance(F * Z, sc.ZeroFunctional) + assert isinstance(Z * F, sc.ZeroFunctional) + + def test_no_value_based_simplification(self, numpy_ctx): + """A functional that merely *happens* to be constant is not recognized. + + Canonicalization is structural: it reads node types, never values. + """ + X, F, _ = _dense(numpy_ctx) + constant_valued = sc.InnerProductFunctional( + numpy_ctx.asarray([0.0, 0.0, 0.0]), X, numpy_ctx + ) + assert isinstance(F * constant_valued, sc.ProductFunctional) + + def test_domain_mismatch_raises(self, numpy_ctx): + X, F, _ = _dense(numpy_ctx) + other = sc.ZeroFunctional(sc.DenseCoordinateSpace((2,), numpy_ctx), numpy_ctx) + with pytest.raises(ValueError, match="same domain"): + sc.make_functional_product(F, other) + + def test_type_guard(self, numpy_ctx): + _, F, _ = _dense(numpy_ctx) + with pytest.raises(TypeError, match="Functional"): + sc.ProductFunctional(F, "nope") + + def test_equality_is_structural_and_ordered(self, numpy_ctx): + X, F, _ = _dense(numpy_ctx) + G = sc.InnerProductFunctional(numpy_ctx.asarray([1.0, 0.5, -2.0]), X, numpy_ctx) + assert (F * G) == (F * G) + # Multiplication commutes, but this is expression-tree equality. + assert (F * G) != (G * F) + + def test_pytree_round_trip(self, numpy_ctx): + X, F, _ = _dense(numpy_ctx) + G = sc.InnerProductFunctional(numpy_ctx.asarray([1.0, 0.5, -2.0]), X, numpy_ctx) + P = F * G + children, aux = P.tree_flatten() + assert sc.ProductFunctional.tree_unflatten(aux, children) == P + + def test_complex_factors_conjugate_the_cofactor(self): + """Riesz convention: coefficients enter conjugated, as in ``ScaledFunctional``. + + With ``F = `` and ``G = ``, ```` must recover + ``G(x) + F(x)``. + """ + ctx = sc.Context(sc.NumpyOps(), dtype=np.complex128) + X = sc.DenseCoordinateSpace((3,), ctx, check_level="standard") + a = ctx.asarray([1 + 1j, 2 - 0.5j, -1 + 0.3j]) + b = ctx.asarray([0.5 - 1j, 1 + 0j, 0.25 + 0.75j]) + x = ctx.asarray([0.5 + 0.25j, -1.0 + 0.75j, 2.0 - 0.5j]) + h = ctx.asarray([1.0 + 0j, -0.5 + 0.25j, 0.75 - 1j]) + + F = sc.InnerProductFunctional(a, X, ctx) + G = sc.InnerProductFunctional(b, X, ctx) + expected = G.value(x) * X.inner(a, h) + F.value(x) * X.inner(b, h) + np.testing.assert_allclose( + to_numpy(X.inner((F * G).grad(x), h)), to_numpy(expected) + ) + + # =========================================================================== # End-to-end: a composed functional drives minimize_optax # =========================================================================== diff --git a/tests/functional/test_composed_functional.py b/tests/functional/test_composed_functional.py index b471723..36ef7bb 100644 --- a/tests/functional/test_composed_functional.py +++ b/tests/functional/test_composed_functional.py @@ -1,6 +1,6 @@ """Tests for :class:`spacecore.ComposedFunctional` and ``make_functional_composed``. -Checklist section 7, ``ComposedFunctional`` / ``make_functional_composed``: +Contract specified here: * ``make_functional_composed`` specializes by type: ``InnerProductFunctional ∘ A`` -> ``InnerProductFunctional``, @@ -183,6 +183,172 @@ def test_convert_preserves_value_across_dtype(self, numpy_f32_ctx, numpy_ctx): np.testing.assert_allclose(to_numpy(converted.value(x)), 52.0) +# =========================================================================== +# Chain rule: grad(F o A)(x) = A^#(grad F(A x)) +# =========================================================================== +def _weighted(ctx, weights): + return sc.DenseCoordinateSpace( + (3,), ctx, geometry=sc.WeightedInnerProduct(ctx.asarray(np.asarray(weights))) + ) + + +_M = np.array([[1.0, 2.0, 0.0], [0.0, 1.0, 3.0], [2.0, 0.0, 1.0]]) + + +class TestChainRuleGradient: + """``ComposedFunctional`` had no gradient at all; it raised ``NotImplementedError``. + + The load-bearing case is **different** non-Euclidean metrics on domain and + codomain: there the metric adjoint ``rapply`` differs from the coordinate + adjoint, so a Euclidean-only test would certify nothing. ``rapply`` *is* + ``A^#``, so no extra Riesz map belongs in the chain rule — applying one + would count the geometry twice. + """ + + def _setup(self, ctx, weighted): + if weighted: + X, Y = _weighted(ctx, [2.0, 5.0, 11.0]), _weighted(ctx, [3.0, 1.0, 7.0]) + else: + X = Y = sc.DenseCoordinateSpace((3,), ctx) + A = sc.DenseLinOp(ctx.asarray(_M), X, Y, ctx) + return X, Y, A, sc.SquaredL2NormFunctional(Y) + + @pytest.mark.parametrize("weighted", [False, True], ids=["euclidean", "weighted"]) + def test_matches_the_explicit_formula(self, numpy_ctx, weighted): + X, Y, A, F = self._setup(numpy_ctx, weighted) + G = sc.ComposedFunctional(F, A) + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + np.testing.assert_allclose( + to_numpy(G.grad(x)), to_numpy(A.rapply(F.grad(A.apply(x)))) + ) + + @pytest.mark.parametrize("weighted", [False, True], ids=["euclidean", "weighted"]) + def test_satisfies_the_riesz_defining_property(self, numpy_ctx, weighted): + """``_X == DG(x)[h]`` — the gradient's actual definition.""" + X, Y, A, F = self._setup(numpy_ctx, weighted) + G = sc.ComposedFunctional(F, A) + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + h = numpy_ctx.asarray([0.5, -1.0, 2.0]) + eps = 1e-6 + directional = ( + float(G.value(X.axpy(eps, h, x))) - float(G.value(X.axpy(-eps, h, x))) + ) / (2.0 * eps) + np.testing.assert_allclose( + to_numpy(X.inner(G.grad(x), h)), directional, rtol=1e-6, atol=1e-6 + ) + + def test_metric_case_rejects_the_coordinate_adjoint(self, numpy_ctx): + """Without this the weighted test above could pass a wrong implementation.""" + X, Y, A, F = self._setup(numpy_ctx, weighted=True) + G = sc.ComposedFunctional(F, A) + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + coordinate_answer = _M.T @ to_numpy(F.grad(A.apply(x))) + assert not np.allclose(to_numpy(G.grad(x)), coordinate_answer) + + def test_agrees_with_the_specialized_pullback(self, numpy_ctx): + """Cross-check against a path that was already correct. + + ``make_functional_composed`` rewrites `` o A`` to + ```` without ever building a ``ComposedFunctional``. Forcing + the generic node on the same operands must give the same gradient. + """ + X, Y = _weighted(numpy_ctx, [2.0, 5.0, 11.0]), _weighted(numpy_ctx, [3.0, 1.0, 7.0]) + A = sc.DenseLinOp(numpy_ctx.asarray(_M), X, Y, numpy_ctx) + c = numpy_ctx.asarray([1.0, 0.5, -2.0]) + F = sc.InnerProductFunctional(c, Y, numpy_ctx) + + specialized = make_functional_composed(F, A) + assert not isinstance(specialized, sc.ComposedFunctional) + generic = sc.ComposedFunctional(F, A) + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + np.testing.assert_allclose(to_numpy(generic.grad(x)), to_numpy(specialized.grad(x))) + + def test_value_and_grad_is_consistent(self, numpy_ctx): + X, Y, A, F = self._setup(numpy_ctx, weighted=True) + G = sc.ComposedFunctional(F, A) + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + value, gradient = G.value_and_grad(x) + np.testing.assert_allclose(to_numpy(value), to_numpy(G.value(x))) + np.testing.assert_allclose(to_numpy(gradient), to_numpy(G.grad(x))) + + def test_value_and_grad_applies_the_operator_once(self, numpy_ctx): + """The point of the fused path: the base default would apply ``A`` twice.""" + X, Y, A, F = self._setup(numpy_ctx, weighted=False) + calls = [] + counting = sc.MatrixFreeLinOp( + lambda v: calls.append("apply") or A.apply(v), + lambda v: A.rapply(v), + X, Y, numpy_ctx, + ) + sc.ComposedFunctional(F, counting).value_and_grad(numpy_ctx.asarray([1.0, 2.0, -1.0])) + assert calls == ["apply"] + + def test_batched_gradient(self, numpy_ctx): + X, Y, A, F = self._setup(numpy_ctx, weighted=True) + G = sc.ComposedFunctional(F, A) + xs = numpy_ctx.asarray([[1.0, 2.0, -1.0], [0.5, -1.0, 2.0]]) + np.testing.assert_allclose( + to_numpy(G.vgrad(xs)), + np.stack([to_numpy(G.grad(numpy_ctx.asarray(row))) for row in to_numpy(xs)]), + ) + + def test_propagates_missing_inner_gradient(self, numpy_ctx): + """A composition is differentiable only if its inner functional is.""" + X = sc.DenseCoordinateSpace((3,), numpy_ctx) + A = sc.DenseLinOp(numpy_ctx.asarray(_M), X, X, numpy_ctx) + G = sc.ComposedFunctional(_SumSquares(X, numpy_ctx), A) + with pytest.raises(NotImplementedError, match="grad"): + G.grad(numpy_ctx.asarray([1.0, 2.0, -1.0])) + + +class TestChainRuleComplex: + """Complex domains, where the conjugation in ``rapply`` has to be right.""" + + def _spaces(self): + ctx = sc.Context(sc.NumpyOps(), dtype=np.complex128) + X = _weighted(ctx, [2.0, 5.0, 11.0]) + Y = _weighted(ctx, [3.0, 1.0, 7.0]) + A = sc.DenseLinOp( + ctx.asarray(np.array([[1 + 1j, 2, 0], [0, 1, 3 - 2j], [2, 0, 1j]])), X, Y, ctx + ) + return ctx, X, Y, A + + def test_holomorphic_functional_satisfies_the_plain_identity(self): + """A complex-*linear* inner: `` == DG[h]`` exactly.""" + ctx, X, Y, A = self._spaces() + c = ctx.asarray([1 + 1j, 2 - 0.5j, -1 + 0.3j]) + G = sc.ComposedFunctional(sc.InnerProductFunctional(c, Y, ctx), A) + x = ctx.asarray([1 + 0j, 2 - 1j, -1 + 2j]) + h = ctx.asarray([0.5 + 1j, -1 + 0j, 2 + 0j]) + eps = 1e-6 + directional = ( + complex(G.value(X.axpy(eps, h, x))) - complex(G.value(X.axpy(-eps, h, x))) + ) / (2.0 * eps) + np.testing.assert_allclose( + complex(X.inner(G.grad(x), h)), directional, rtol=1e-6, atol=1e-6 + ) + # ...and equals the adjoint pull-back of the representer. + np.testing.assert_allclose(to_numpy(G.grad(x)), to_numpy(A.H.apply(c))) + + def test_real_valued_functional_uses_the_real_part_convention(self): + """``1/2||y||^2`` is real-valued on a complex space, hence not holomorphic. + + The library's convention there is ``Re == DF[h]`` (CR calculus); + that already holds for the inner functional, and composition preserves it. + """ + ctx, X, Y, A = self._spaces() + G = sc.ComposedFunctional(sc.SquaredL2NormFunctional(Y), A) + x = ctx.asarray([1 + 0j, 2 - 1j, -1 + 2j]) + h = ctx.asarray([0.5 + 1j, -1 + 0j, 2 + 0j]) + eps = 1e-6 + directional = ( + complex(G.value(X.axpy(eps, h, x))) - complex(G.value(X.axpy(-eps, h, x))) + ) / (2.0 * eps) + np.testing.assert_allclose( + complex(X.inner(G.grad(x), h)).real, directional.real, rtol=1e-6, atol=1e-6 + ) + + # =========================================================================== # Private element helpers # =========================================================================== diff --git a/tests/functional/test_functional_base.py b/tests/functional/test_functional_base.py index 6678745..10629db 100644 --- a/tests/functional/test_functional_base.py +++ b/tests/functional/test_functional_base.py @@ -1,6 +1,6 @@ """Tests for :class:`spacecore.Functional` — the scalar-map base. -Checklist section 7, ``Functional`` base: +Contract specified here: * ``domain`` property returns the context-converted domain. * ``value(x)`` / ``__call__(x)`` alias. @@ -21,6 +21,7 @@ import spacecore as sc +from spacecore._checks import checked_method from spacecore.functional._base import ( _check_scalar_shape, _leading_batch_size, @@ -268,12 +269,112 @@ def test_accepts_matching_shape(self): _check_scalar_shape(np.zeros((4,)), (4,)) # batch, no raise def test_rejects_mismatched_shape(self): - with pytest.raises(ValueError, match="Expected scalar batch output with shape"): + with pytest.raises(ValueError, match="Expected scalar output with shape"): _check_scalar_shape(np.zeros((2,)), ()) def test_treats_objects_without_shape_as_scalar(self): _check_scalar_shape(3.0, ()) # Python float has no ``shape`` -> () + def test_message_distinguishes_single_from_batched(self): + """The helper now guards every ``value``/``vvalue``, so wording matters.""" + with pytest.raises(ValueError, match="Expected scalar output"): + _check_scalar_shape(np.zeros((2,)), ()) + with pytest.raises(ValueError, match="Expected scalar batch output"): + _check_scalar_shape(np.zeros((2,)), (4,)) + + +# =========================================================================== +# Functional codomain contract: out_scalar / out_batched_scalar +# =========================================================================== +class _NonScalar(sc.Functional): + """Violates the ``F : X -> K`` contract by returning its input.""" + + def value(self, x, *args, **kwargs): + return x + + def grad(self, x, *args, **kwargs): + return self.domain.zeros() + + def tree_flatten(self): + return (), (self.domain, self.ctx) + + @classmethod + def tree_unflatten(cls, aux, children): + domain, ctx = aux + return cls(domain, ctx) + + def _convert(self, new_ctx): + return _NonScalar(self.domain.convert(new_ctx), new_ctx) + + +class TestScalarOutputContract: + """``out_scalar`` is the codomain check ``out_space`` cannot express. + + A ``Functional``'s codomain is the scalar field, reported only as a string + via ``domain.field`` — there is no ``Space`` object to bind ``out_space`` to, + which is why the output side went unchecked while every input was validated. + """ + + def _domain(self, ctx, check_level=None): + return sc.DenseCoordinateSpace((2,), ctx, check_level=check_level) + + @pytest.mark.parametrize("build, label", [ + (lambda F: 2.0 * F, "scaled"), + (lambda F: F + F, "sum"), + (lambda F: F + 1.0, "shifted"), + (lambda F: F * F, "product"), + ]) + def test_algebra_nodes_reject_non_scalar_operand_output(self, numpy_ctx, build, label): + X = self._domain(numpy_ctx) + node = build(_NonScalar(X, numpy_ctx)) + with pytest.raises(ValueError, match="Expected scalar output"): + node.value(numpy_ctx.asarray([1.0, 2.0])) + + @pytest.mark.parametrize("level, raises", [ + ("none", False), + ("cheap", False), + ("standard", True), + ("strict", True), + ]) + def test_runs_at_standard_and_above(self, numpy_ctx, level, raises): + """Matches the level the hand-written checks it replaces already used.""" + X = self._domain(numpy_ctx, check_level=level) + node = sc.ScaledFunctional(2.0, _NonScalar(X, numpy_ctx, check_level=level), + check_level=level) + x = numpy_ctx.asarray([1.0, 2.0]) + if raises: + with pytest.raises(ValueError, match="Expected scalar output"): + node.value(x) + else: + node.value(x) # no raise + + def test_conforming_functionals_are_unaffected(self, numpy_ctx): + X = self._domain(numpy_ctx) + F = sc.InnerProductFunctional(numpy_ctx.asarray([1.0, 1.0]), X, numpy_ctx) + assert tuple(np.shape(F.value(numpy_ctx.asarray([1.0, 2.0])))) == () + assert F.vvalue(numpy_ctx.asarray([[1.0, 2.0], [3.0, 4.0]])).shape == (2,) + + def test_batched_output_must_be_one_scalar_per_element(self, numpy_ctx): + X = self._domain(numpy_ctx) + F = sc.MatrixFreeLinearFunctional( + lambda x: X.inner(numpy_ctx.asarray([1.0, 1.0]), x), X, numpy_ctx, + vvalue=lambda xs: xs, # returns (N, 2), not (N,) + ) + with pytest.raises(ValueError, match="scalar batch output"): + F.vvalue(numpy_ctx.asarray([[1.0, 2.0], [3.0, 4.0]])) + + +class TestCheckedMethodScalarFlags: + """Guards on the decorator itself.""" + + def test_flags_are_mutually_exclusive(self): + with pytest.raises(TypeError): + checked_method(in_space="domain", out_scalar=True, out_batched_scalar=True) + + def test_batched_scalar_requires_in_space(self): + with pytest.raises(TypeError, match="requires in_space"): + checked_method(out_batched_scalar=True) + # =========================================================================== # Private helper: _leading_batch_size diff --git a/tests/functional/test_generated_functionals.py b/tests/functional/test_generated_functionals.py index a0dec16..1a6a458 100644 --- a/tests/functional/test_generated_functionals.py +++ b/tests/functional/test_generated_functionals.py @@ -211,27 +211,33 @@ def test_functional_generator_records_check_level_and_required_reference_fields( def test_none_skips_optional_functional_input_membership_checks(): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - domain = sc.DenseCoordinateSpace((2,), ctx) - functional = sc.MatrixFreeLinearFunctional(lambda _x: ctx.asarray(3.0), domain, ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + domain = sc.DenseCoordinateSpace((2,), ctx, check_level="none") + functional = sc.MatrixFreeLinearFunctional( + lambda _x: ctx.asarray(3.0), domain, ctx, check_level="none" + ) np.testing.assert_allclose(functional.value(ctx.asarray([1.0, 2.0, 3.0])), 3.0) @pytest.mark.parametrize("check_level", check_level_params(("cheap", "standard", "strict"))) def test_checked_functional_levels_reject_wrong_input_shape(check_level): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level=check_level) - domain = sc.DenseCoordinateSpace((2,), ctx) - functional = sc.MatrixFreeLinearFunctional(lambda _x: ctx.asarray(3.0), domain, ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + domain = sc.DenseCoordinateSpace((2,), ctx, check_level=check_level) + functional = sc.MatrixFreeLinearFunctional( + lambda _x: ctx.asarray(3.0), domain, ctx, check_level=check_level + ) with pytest.raises(TypeError, match="Expected shape"): functional.value(ctx.asarray([1.0, 2.0, 3.0])) def test_cheap_functional_checks_reject_field_and_dtype_mismatch(): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="cheap") - domain = sc.DenseCoordinateSpace((2,), ctx) - functional = sc.InnerProductFunctional(ctx.asarray([1.0, 2.0]), domain, ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + domain = sc.DenseCoordinateSpace((2,), ctx, check_level="cheap") + functional = sc.InnerProductFunctional( + ctx.asarray([1.0, 2.0]), domain, ctx, check_level="cheap" + ) with pytest.raises(TypeError, match="real scalar field"): functional.value(np.asarray([1.0 + 1.0j, 2.0], dtype=np.complex128)) @@ -241,26 +247,29 @@ def test_cheap_functional_checks_reject_field_and_dtype_mismatch(): @pytest.mark.parametrize("check_level", check_level_params(("standard", "strict"))) def test_standard_functional_checks_reject_nonscalar_output(check_level): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level=check_level) - domain = sc.DenseCoordinateSpace((2,), ctx) - functional = sc.MatrixFreeLinearFunctional(lambda x: x, domain, ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + domain = sc.DenseCoordinateSpace((2,), ctx, check_level=check_level) + functional = sc.MatrixFreeLinearFunctional( + lambda x: x, domain, ctx, check_level=check_level + ) - with pytest.raises(ValueError, match="Expected scalar batch output"): + with pytest.raises(ValueError, match="Expected scalar output"): functional.value(ctx.asarray([1.0, 2.0])) @pytest.mark.parametrize("check_level", check_level_params()) def test_pullback_domain_codomain_mismatch_is_always_rejected(check_level): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level=check_level) - domain = sc.DenseCoordinateSpace((2,), ctx) - wrong_codomain = sc.DenseCoordinateSpace((3,), ctx) - functional = sc.InnerProductFunctional(ctx.asarray([1.0, 2.0]), domain, ctx) - operator = sc.DenseLinOp( - ctx.asarray([[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]), - domain, - wrong_codomain, - ctx, - ) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + with sc.use_check_level(check_level): + domain = sc.DenseCoordinateSpace((2,), ctx) + wrong_codomain = sc.DenseCoordinateSpace((3,), ctx) + functional = sc.InnerProductFunctional(ctx.asarray([1.0, 2.0]), domain, ctx) + operator = sc.DenseLinOp( + ctx.asarray([[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]), + domain, + wrong_codomain, + ctx, + ) with pytest.raises(ValueError, match="A.codomain == F.domain"): functional.compose(operator) diff --git a/tests/functional/test_inner_product_functional.py b/tests/functional/test_inner_product_functional.py index c93b451..4f4138e 100644 --- a/tests/functional/test_inner_product_functional.py +++ b/tests/functional/test_inner_product_functional.py @@ -1,6 +1,6 @@ """Tests for :class:`spacecore.InnerProductFunctional`. -Checklist section 7, ``InnerProductFunctional``: +Contract specified here: * ``value(x) == domain.inner(representer, x)`` (Euclidean and weighted). * ``representer`` property returns the stored, context-converted element. diff --git a/tests/functional/test_linear_functional.py b/tests/functional/test_linear_functional.py index 17f9bd2..2a81341 100644 --- a/tests/functional/test_linear_functional.py +++ b/tests/functional/test_linear_functional.py @@ -1,6 +1,6 @@ """Tests for :class:`spacecore.LinearFunctional` — the linear-map base. -Checklist section 7, ``LinearFunctional``: +Contract specified here: * ``grad(x)`` returns the *constant* Riesz representer, independent of ``x``. * ``vgrad(xs)`` broadcasts that constant representer across the batch axis. diff --git a/tests/functional/test_linop_quadratic_form.py b/tests/functional/test_linop_quadratic_form.py index 58e4927..a46762b 100644 --- a/tests/functional/test_linop_quadratic_form.py +++ b/tests/functional/test_linop_quadratic_form.py @@ -1,6 +1,6 @@ """Tests for :class:`spacecore.LinOpQuadraticForm`. -Checklist section 7, ``LinOpQuadraticForm``: +Contract specified here: * Construction guards: ``Q`` must be a square ``LinOp``, ``linear`` a ``LinearFunctional`` on ``Q.domain``, ``a`` scalar, ``Q`` Hermitian. @@ -71,11 +71,12 @@ def test_rejects_non_hermitian_dense_operator(self, numpy_ctx): sc.LinOpQuadraticForm(Q, ctx=numpy_ctx) def test_rejects_nonscalar_constant(self): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - space = sc.DenseCoordinateSpace((2,), ctx) - Q = sc.IdentityLinOp(space, ctx) - with pytest.raises(ValueError, match="scalar batch"): - sc.LinOpQuadraticForm(Q, a=ctx.asarray([0.0, 0.0]), ctx=ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + with sc.use_check_level("none"): + space = sc.DenseCoordinateSpace((2,), ctx) + Q = sc.IdentityLinOp(space, ctx) + with pytest.raises(ValueError, match="scalar output"): + sc.LinOpQuadraticForm(Q, a=ctx.asarray([0.0, 0.0]), ctx=ctx) def test_explicit_context_overrides_inferred(self, numpy_f32_ctx, numpy_ctx): space = sc.DenseCoordinateSpace((2,), numpy_f32_ctx) diff --git a/tests/functional/test_matrix_free_linear_functional.py b/tests/functional/test_matrix_free_linear_functional.py index 0414ca6..cf1f954 100644 --- a/tests/functional/test_matrix_free_linear_functional.py +++ b/tests/functional/test_matrix_free_linear_functional.py @@ -1,6 +1,6 @@ """Tests for :class:`spacecore.MatrixFreeLinearFunctional`. -Checklist section 7, ``MatrixFreeLinearFunctional``: +Contract specified here: * The supplied ``value`` callable is used verbatim. * ``representer`` raises ``NotImplementedError`` (no stored dual vector). @@ -61,13 +61,13 @@ def test_value_enforces_scalar_output_under_standard_checks(self, numpy_ctx): space = sc.DenseCoordinateSpace((2,), numpy_ctx) # Callable returns a vector, not a scalar. f = sc.MatrixFreeLinearFunctional(lambda x: x, space, numpy_ctx) - with pytest.raises(ValueError, match="Expected scalar batch output"): + with pytest.raises(ValueError, match="Expected scalar output"): f.value(numpy_ctx.asarray([1.0, 2.0])) def test_value_skips_scalar_check_at_none_level(self): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - space = sc.DenseCoordinateSpace((2,), ctx) - f = sc.MatrixFreeLinearFunctional(lambda x: x, space, ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + space = sc.DenseCoordinateSpace((2,), ctx, check_level="none") + f = sc.MatrixFreeLinearFunctional(lambda x: x, space, ctx, check_level="none") out = f.value(ctx.asarray([1.0, 2.0])) # no raise np.testing.assert_allclose(to_numpy(out), [1.0, 2.0]) @@ -104,10 +104,12 @@ def test_falls_back_to_base_vmap_without_callable(self, numpy_ctx): def test_python_loop_fallback_warns_once_on_numpy(self): _VMAP_FALLBACK_WARNED.clear() - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - space = sc.DenseCoordinateSpace((2,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + space = sc.DenseCoordinateSpace((2,), ctx, check_level="none") c = ctx.asarray([1.0, -2.0]) - f = sc.MatrixFreeLinearFunctional(lambda y: space.inner(c, y), space, ctx) + f = sc.MatrixFreeLinearFunctional( + lambda y: space.inner(c, y), space, ctx, check_level="none" + ) xs = ctx.asarray(np.arange(80.0).reshape(40, 2)) with pytest.warns(RuntimeWarning, match="falls back to a Python loop"): diff --git a/tests/functional/test_metric_gradient.py b/tests/functional/test_metric_gradient.py index cc7c3b7..5b5f9c2 100644 --- a/tests/functional/test_metric_gradient.py +++ b/tests/functional/test_metric_gradient.py @@ -1,4 +1,5 @@ import warnings +from contextlib import contextmanager import numpy as np import pytest @@ -8,14 +9,15 @@ from tests._helpers import has_jax, jax_real_dtype, to_numpy +@contextmanager def _numpy_context(): - - return sc.Context(sc.NumpyOps(), dtype=np.float64) + yield sc.Context(sc.NumpyOps(), dtype=np.float64) +@contextmanager def _jax_context(): - - return sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") + with sc.use_check_level("none"): + yield sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) def _contexts(): @@ -72,25 +74,27 @@ def _assert_gradient_identity(functional, x, v, eps, atol): @pytest.mark.parametrize("ctx_factory,eps,atol", list(_contexts())) def test_linop_quadratic_gradient_identity_on_weighted_space(ctx_factory, eps, atol): - ctx = ctx_factory() - space = _weighted_space(ctx) - Q = sc.DenseLinOp(_self_adjoint_metric_matrix(ctx), space, space, ctx) - c = ctx.asarray([0.25, -1.5, 2.0]) - functional = sc.LinOpQuadraticForm(Q, sc.InnerProductFunctional(c, space, ctx), 1.25, ctx) - x = ctx.asarray([0.5, -1.0, 2.0]) - v = ctx.asarray([1.25, 0.75, -0.5]) + with ctx_factory() as ctx: + space = _weighted_space(ctx) + Q = sc.DenseLinOp(_self_adjoint_metric_matrix(ctx), space, space, ctx) + c = ctx.asarray([0.25, -1.5, 2.0]) + functional = sc.LinOpQuadraticForm( + Q, sc.InnerProductFunctional(c, space, ctx), 1.25, ctx + ) + x = ctx.asarray([0.5, -1.0, 2.0]) + v = ctx.asarray([1.25, 0.75, -0.5]) _assert_gradient_identity(functional, x, v, eps, atol) @pytest.mark.parametrize("ctx_factory,eps,atol", list(_contexts())) def test_inner_product_functional_gradient_identity_on_weighted_space(ctx_factory, eps, atol): - ctx = ctx_factory() - space = _weighted_space(ctx) - c = ctx.asarray([0.25, -1.5, 2.0]) - functional = sc.InnerProductFunctional(c, space, ctx) - x = ctx.asarray([0.5, -1.0, 2.0]) - v = ctx.asarray([1.25, 0.75, -0.5]) + with ctx_factory() as ctx: + space = _weighted_space(ctx) + c = ctx.asarray([0.25, -1.5, 2.0]) + functional = sc.InnerProductFunctional(c, space, ctx) + x = ctx.asarray([0.5, -1.0, 2.0]) + v = ctx.asarray([1.25, 0.75, -0.5]) _assert_gradient_identity(functional, x, v, eps, atol) np.testing.assert_allclose(to_numpy(functional.grad(x)), to_numpy(c)) @@ -110,16 +114,16 @@ def test_euclidean_quadratic_gradient_behavior_is_unchanged(): @pytest.mark.parametrize("ctx_factory,eps,atol", list(_contexts())) def test_inner_product_functional_compose_uses_metric_adjoint_pullback(ctx_factory, eps, atol): - ctx = ctx_factory() - domain = _weighted_vector_space(ctx, [2.0, 5.0]) - codomain = _weighted_vector_space(ctx, [3.0, 7.0, 11.0]) - matrix = ctx.asarray([[1.0, -2.0], [0.5, 3.0], [4.0, -1.0]]) - A = sc.DenseLinOp(matrix, domain, codomain, ctx) - c = ctx.asarray([0.25, -1.5, 2.0]) - functional = sc.InnerProductFunctional(c, codomain, ctx) - composed = functional.compose(A) - x = ctx.asarray([0.5, -1.0]) - v = ctx.asarray([1.25, 0.75]) + with ctx_factory() as ctx: + domain = _weighted_vector_space(ctx, [2.0, 5.0]) + codomain = _weighted_vector_space(ctx, [3.0, 7.0, 11.0]) + matrix = ctx.asarray([[1.0, -2.0], [0.5, 3.0], [4.0, -1.0]]) + A = sc.DenseLinOp(matrix, domain, codomain, ctx) + c = ctx.asarray([0.25, -1.5, 2.0]) + functional = sc.InnerProductFunctional(c, codomain, ctx) + composed = functional.compose(A) + x = ctx.asarray([0.5, -1.0]) + v = ctx.asarray([1.25, 0.75]) np.testing.assert_allclose( to_numpy(composed.value(x)), @@ -141,17 +145,17 @@ def test_inner_product_functional_compose_uses_metric_adjoint_pullback(ctx_facto def test_inner_product_functional_vectorized_batches_match_elementwise( ctx_factory, eps, atol, weighted ): - ctx = ctx_factory() - space = _weighted_space(ctx) if weighted else sc.DenseCoordinateSpace((3,), ctx) - c = ctx.asarray([0.25, -1.5, 2.0]) - functional = sc.InnerProductFunctional(c, space, ctx) - xs = ctx.asarray( - [ - [0.5, -1.0, 2.0], - [1.25, 0.75, -0.5], - [-2.0, 0.25, 1.5], - ] - ) + with ctx_factory() as ctx: + space = _weighted_space(ctx) if weighted else sc.DenseCoordinateSpace((3,), ctx) + c = ctx.asarray([0.25, -1.5, 2.0]) + functional = sc.InnerProductFunctional(c, space, ctx) + xs = ctx.asarray( + [ + [0.5, -1.0, 2.0], + [1.25, 0.75, -0.5], + [-2.0, 0.25, 1.5], + ] + ) expected_values = functional.ops.stack(tuple(functional.value(x) for x in xs), axis=0) expected_grads = functional.ops.stack(tuple(functional.grad(x) for x in xs), axis=0) @@ -169,29 +173,29 @@ def test_inner_product_functional_vectorized_batches_match_elementwise( def test_linop_quadratic_form_vectorized_batches_match_elementwise( ctx_factory, eps, atol, weighted ): - ctx = ctx_factory() - space = _weighted_space(ctx) if weighted else sc.DenseCoordinateSpace((3,), ctx) - matrix = ( - _self_adjoint_metric_matrix(ctx) - if weighted - else ctx.asarray( + with ctx_factory() as ctx: + space = _weighted_space(ctx) if weighted else sc.DenseCoordinateSpace((3,), ctx) + matrix = ( + _self_adjoint_metric_matrix(ctx) + if weighted + else ctx.asarray( + [ + [4.0, 1.0, -0.5], + [1.0, 6.0, 2.0], + [-0.5, 2.0, 3.0], + ] + ) + ) + Q = sc.DenseLinOp(matrix, space, space, ctx) + linear = sc.InnerProductFunctional(ctx.asarray([0.25, -1.5, 2.0]), space, ctx) + functional = sc.LinOpQuadraticForm(Q, linear, 1.25, ctx) + xs = ctx.asarray( [ - [4.0, 1.0, -0.5], - [1.0, 6.0, 2.0], - [-0.5, 2.0, 3.0], + [0.5, -1.0, 2.0], + [1.25, 0.75, -0.5], + [-2.0, 0.25, 1.5], ] ) - ) - Q = sc.DenseLinOp(matrix, space, space, ctx) - linear = sc.InnerProductFunctional(ctx.asarray([0.25, -1.5, 2.0]), space, ctx) - functional = sc.LinOpQuadraticForm(Q, linear, 1.25, ctx) - xs = ctx.asarray( - [ - [0.5, -1.0, 2.0], - [1.25, 0.75, -0.5], - [-2.0, 0.25, 1.5], - ] - ) expected_values = functional.ops.stack(tuple(functional.value(x) for x in xs), axis=0) expected_grads = functional.ops.stack(tuple(functional.grad(x) for x in xs), axis=0) @@ -206,10 +210,12 @@ def test_linop_quadratic_form_vectorized_batches_match_elementwise( def test_matrix_free_functional_vvalue_python_loop_warns_once_on_numpy(): _functional_base._VMAP_FALLBACK_WARNED.clear() - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - space = sc.DenseCoordinateSpace((2,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + space = sc.DenseCoordinateSpace((2,), ctx, check_level="none") c = ctx.asarray([1.0, -2.0]) - functional = sc.MatrixFreeLinearFunctional(lambda x: space.inner(c, x), space, ctx) + functional = sc.MatrixFreeLinearFunctional( + lambda x: space.inner(c, x), space, ctx, check_level="none" + ) xs = ctx.asarray(np.arange(80.0).reshape(40, 2)) with pytest.warns(RuntimeWarning, match="falls back to a Python loop"): @@ -221,10 +227,10 @@ def test_matrix_free_functional_vvalue_python_loop_warns_once_on_numpy(): def test_vectorized_functionals_do_not_warn_on_numpy(): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - space = sc.DenseCoordinateSpace((2,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + space = sc.DenseCoordinateSpace((2,), ctx, check_level="none") c = ctx.asarray([1.0, -2.0]) - functional = sc.InnerProductFunctional(c, space, ctx) + functional = sc.InnerProductFunctional(c, space, ctx, check_level="none") xs = ctx.asarray(np.arange(80.0).reshape(40, 2)) with warnings.catch_warnings(record=True) as caught: @@ -236,10 +242,12 @@ def test_vectorized_functionals_do_not_warn_on_numpy(): @pytest.mark.skipif(not has_jax(), reason="jax is not installed") def test_matrix_free_functional_vvalue_does_not_warn_on_native_vmap_backend(): - ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") - space = sc.DenseCoordinateSpace((2,), ctx) + ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) + space = sc.DenseCoordinateSpace((2,), ctx, check_level="none") c = ctx.asarray([1.0, -2.0]) - functional = sc.MatrixFreeLinearFunctional(lambda x: space.inner(c, x), space, ctx) + functional = sc.MatrixFreeLinearFunctional( + lambda x: space.inner(c, x), space, ctx, check_level="none" + ) xs = ctx.asarray(np.arange(80.0).reshape(40, 2)) with warnings.catch_warnings(record=True) as caught: diff --git a/tests/functional/test_quadratic_form.py b/tests/functional/test_quadratic_form.py index 5c5332b..0c01d4f 100644 --- a/tests/functional/test_quadratic_form.py +++ b/tests/functional/test_quadratic_form.py @@ -1,6 +1,6 @@ """Tests for :class:`spacecore.QuadraticForm` — the quadratic-objective base. -Checklist section 7, ``QuadraticForm`` base: +Contract specified here: * ``grad`` / ``hess_apply`` raise ``NotImplementedError`` until a subclass provides them. diff --git a/tests/functional/tools/test_proximal.py b/tests/functional/tools/test_proximal.py index 1612b97..c639b56 100644 --- a/tests/functional/tools/test_proximal.py +++ b/tests/functional/tools/test_proximal.py @@ -198,6 +198,15 @@ def test_project_nonneg(self, numpy_ctx): out = sc.project_nonneg(v, X) np.testing.assert_allclose(to_numpy(out), np.maximum(to_numpy(v), 0.0)) + def test_wrappers_reject_negative_step_nonneg(self, numpy_ctx): + """The step guard runs before the constrained branch too.""" + X = sc.DenseCoordinateSpace((3,), numpy_ctx) + v = numpy_ctx.asarray([1.0, 2.0, 3.0]) + with pytest.raises(ValueError): + sc.prox_l1(v, -1.0, X, nonneg=True) + with pytest.raises(ValueError): + sc.prox_l2sq(v, -1.0, X, nonneg=True) + def test_wrappers_reject_negative_step(self, numpy_ctx): X = sc.DenseCoordinateSpace((3,), numpy_ctx) v = numpy_ctx.asarray([1.0, 2.0, 3.0]) @@ -226,6 +235,136 @@ def test_wrong_shape_v_names_the_caller_argument(self, numpy_ctx): # --------------------------------------------------------------------------- # S04 regression: metric-correct proximal gradient on a weighted space # --------------------------------------------------------------------------- +class TestNonnegativeWrappers: + """``prox_l1`` / ``prox_l2sq`` with ``nonneg=True``. + + Both delegate the constraint to ``generalized_shrinkage``, which already + implemented it — the wrappers simply did not expose the flag. Each is + verified against its **defining constrained minimization** on a weighted + metric, not against a twin formula. + """ + + W = np.array([2.0, 5.0, 11.0]) + + def _space(self, ctx): + return _weighted_space(ctx, self.W) + + def _assert_is_constrained_minimizer(self, out, objective, seed): + """Local minimality on the feasible set, by random feasible perturbation.""" + base = objective(out) + rng = np.random.default_rng(seed) + for _ in range(200): + perturbed = np.maximum(out + 1e-3 * rng.standard_normal(out.size), 0.0) + assert objective(perturbed) >= base - 1e-9 + + def test_prox_l1_nonneg_is_the_constrained_minimizer(self, numpy_ctx): + X = self._space(numpy_ctx) + v = numpy_ctx.asarray([-3.0, 0.5, 2.0]) + t = 1.3 + out = to_numpy(sc.prox_l1(v, t, X, nonneg=True)) + assert np.all(out >= 0.0) + + def objective(x): + diff = x - to_numpy(v) + return 0.5 * np.sum(self.W * diff * diff) + t * np.sum(np.abs(x)) + + self._assert_is_constrained_minimizer(out, objective, seed=1) + + def test_prox_l2sq_nonneg_is_the_constrained_minimizer(self, numpy_ctx): + X = self._space(numpy_ctx) + v = numpy_ctx.asarray([-3.0, 0.5, 2.0]) + t = 1.3 + out = to_numpy(sc.prox_l2sq(v, t, X, nonneg=True)) + assert np.all(out >= 0.0) + + def objective(x): + diff = x - to_numpy(v) + return 0.5 * np.sum(self.W * diff * diff) + t * 0.5 * np.sum(self.W * x * x) + + self._assert_is_constrained_minimizer(out, objective, seed=2) + + def test_prox_l1_nonneg_uses_the_metric_aware_one_sided_threshold(self, numpy_ctx): + """On the orthant ``||x||_1`` is linear, giving ``max(v - t/w, 0)``.""" + X = self._space(numpy_ctx) + v = numpy_ctx.asarray([-3.0, 0.5, 2.0]) + t = 1.3 + np.testing.assert_allclose( + to_numpy(sc.prox_l1(v, t, X, nonneg=True)), + np.maximum(to_numpy(v) - t / self.W, 0.0), + ) + + def test_prox_l2sq_nonneg_is_clipped_shrinkage(self, numpy_ctx): + X = self._space(numpy_ctx) + v = numpy_ctx.asarray([-3.0, 0.5, 2.0]) + t = 1.3 + np.testing.assert_allclose( + to_numpy(sc.prox_l2sq(v, t, X, nonneg=True)), + np.maximum(to_numpy(v), 0.0) / (1.0 + t), + ) + + @pytest.mark.parametrize("prox", [sc.prox_l1, sc.prox_l2sq], ids=["l1", "l2sq"]) + def test_clipping_the_unconstrained_prox_agrees(self, numpy_ctx, prox): + """Pins an identity the docstrings claim; it is not general to prox operators. + + For these two the constrained minimizer coincides with clipping the + unconstrained one at zero — for ``l1`` because ``v <= 0`` maps to ``0`` + either way and ``v > 0`` already reduces to the one-sided form. Asserted + rather than assumed, so a future change to either branch is caught. + """ + X = self._space(numpy_ctx) + rng = np.random.default_rng(3) + for _ in range(50): + v = numpy_ctx.asarray(rng.normal(scale=3.0, size=3)) + t = float(rng.uniform(0.0, 4.0)) + np.testing.assert_allclose( + np.maximum(to_numpy(prox(v, t, X)), 0.0), + to_numpy(prox(v, t, X, nonneg=True)), + ) + + @pytest.mark.parametrize("prox", [sc.prox_l1, sc.prox_l2sq], ids=["l1", "l2sq"]) + def test_zero_step_degenerates_to_projection(self, numpy_ctx, prox): + X = self._space(numpy_ctx) + v = numpy_ctx.asarray([-3.0, 0.5, 2.0]) + np.testing.assert_allclose( + to_numpy(prox(v, 0.0, X, nonneg=True)), to_numpy(sc.project_nonneg(v, X)) + ) + + @pytest.mark.parametrize("prox", [sc.prox_l1, sc.prox_l2sq], ids=["l1", "l2sq"]) + def test_default_is_unconstrained(self, numpy_ctx, prox): + """``nonneg`` defaults to ``False``; existing callers are unaffected.""" + X = self._space(numpy_ctx) + v = numpy_ctx.asarray([-3.0, 0.5, 2.0]) + np.testing.assert_allclose( + to_numpy(prox(v, 1.3, X)), to_numpy(prox(v, 1.3, X, nonneg=False)) + ) + assert np.any(to_numpy(prox(v, 1.3, X)) < 0.0) # genuinely unconstrained + + @pytest.mark.parametrize("prox", [sc.prox_l1, sc.prox_l2sq], ids=["l1", "l2sq"]) + def test_nonneg_is_keyword_only(self, numpy_ctx, prox): + """Guards the signature: a positional fourth argument must not bind it.""" + X = self._space(numpy_ctx) + v = numpy_ctx.asarray([-3.0, 0.5, 2.0]) + with pytest.raises(TypeError): + prox(v, 1.3, X, True) + + @pytest.mark.parametrize("prox", [sc.prox_l1, sc.prox_l2sq], ids=["l1", "l2sq"]) + def test_nonneg_rejects_complex_spaces(self, numpy_complex_ctx, prox): + """``x >= 0`` is undefined on a complex space; the primitive's guard applies.""" + X = sc.DenseCoordinateSpace((3,), numpy_complex_ctx) + v = numpy_complex_ctx.asarray([1 + 1j, -2 + 0j, 0.5 - 1j]) + with pytest.raises(ValueError, match="real spaces only"): + prox(v, 1.3, X, nonneg=True) + + @pytest.mark.parametrize("prox", [sc.prox_l1, sc.prox_l2sq], ids=["l1", "l2sq"]) + def test_nonneg_still_rejects_a_non_diagonal_metric(self, numpy_ctx, prox): + """The separability refusal is not bypassed by the constrained branch.""" + matrix = numpy_ctx.asarray([[2.0, 0.3, 0.0], [0.3, 2.0, 0.1], [0.0, 0.1, 2.0]]) + X = sc.DenseCoordinateSpace((3,), numpy_ctx, geometry=_FullMetric(matrix)) + v = numpy_ctx.asarray([-3.0, 0.5, 2.0]) + with pytest.raises(ValueError, match="not diagonal"): + prox(v, 1.3, X, nonneg=True) + + class TestMetricTrap: def _setup(self, ctx): w = np.array([2.0, 5.0, 11.0]) diff --git a/tests/functional/tools/test_spectral.py b/tests/functional/tools/test_spectral.py index 8c9c2c5..eb50119 100644 --- a/tests/functional/tools/test_spectral.py +++ b/tests/functional/tools/test_spectral.py @@ -19,20 +19,30 @@ def _hermitian(ctx, matrix): _M = [[3.0, 1.0, 0.0], [1.0, 2.0, -1.0], [0.0, -1.0, 4.0]] +def _schatten(X, p): + """Schatten ``p``-norm as the spectral lift of the coordinate ``p``-norm. + + There is no dedicated Schatten class: ``SpectralFunctional`` lifts any + symmetric coordinate functional, so the Schatten norm is the coordinate + ``p``-norm composed with the spectrum. + """ + return sc.spectralize(X, lambda s: sc.LpNormFunctional(s, p)) + + class TestSpectralValue: @pytest.mark.parametrize("p", [1.0, 1.5, 2.0, 3.0]) def test_value_is_schatten_p_norm(self, numpy_ctx, p): X = sc.HermitianSpace(3, ctx=numpy_ctx) A, m = _hermitian(numpy_ctx, _M) evals = np.linalg.eigvalsh(m) - f = sc.SpectralLpNormFunctional(X, p) + f = _schatten(X, p) expected = np.sum(np.abs(evals) ** p) ** (1.0 / p) np.testing.assert_allclose(to_numpy(f.value(A)), expected) def test_schatten_2_is_frobenius_norm(self, numpy_ctx): X = sc.HermitianSpace(3, ctx=numpy_ctx) A, m = _hermitian(numpy_ctx, _M) - f = sc.SpectralLpNormFunctional(X, 2.0) + f = _schatten(X, 2.0) np.testing.assert_allclose(to_numpy(f.value(A)), np.linalg.norm(m, "fro")) def test_nuclear_norm_is_sum_of_singular_values(self, numpy_ctx): @@ -44,15 +54,15 @@ def test_nuclear_norm_is_sum_of_singular_values(self, numpy_ctx): def test_nuclear_norm_is_p1_spectral(self, numpy_ctx): X = sc.HermitianSpace(2, ctx=numpy_ctx) f = sc.NuclearNormFunctional(X) - assert isinstance(f, sc.SpectralLpNormFunctional) - assert f.p == 1.0 + assert isinstance(f, sc.SpectralFunctional) + assert f.base.p == 1.0 class TestSpectralGradient: def test_schatten_2_gradient_is_normalized_matrix(self, numpy_ctx): X = sc.HermitianSpace(3, ctx=numpy_ctx) A, m = _hermitian(numpy_ctx, _M) - f = sc.SpectralLpNormFunctional(X, 2.0) + f = _schatten(X, 2.0) np.testing.assert_allclose( to_numpy(f.grad(A)), m / np.linalg.norm(m, "fro"), atol=1e-12 ) @@ -67,14 +77,14 @@ def test_nuclear_gradient_is_matrix_sign(self, numpy_ctx): def test_gradient_is_hermitian(self, numpy_ctx): X = sc.HermitianSpace(3, ctx=numpy_ctx) A, _ = _hermitian(numpy_ctx, _M) - g = to_numpy(sc.SpectralLpNormFunctional(X, 3.0).grad(A)) + g = to_numpy(_schatten(X, 3.0).grad(A)) np.testing.assert_allclose(g, g.T, atol=1e-12) @pytest.mark.parametrize("p", [1.5, 2.0, 3.0]) def test_gradient_satisfies_directional_derivative_identity(self, numpy_ctx, p): X = sc.HermitianSpace(3, ctx=numpy_ctx) A, m = _hermitian(numpy_ctx, _M) - f = sc.SpectralLpNormFunctional(X, p) + f = _schatten(X, p) d = np.array([[0.2, 0.5, -0.1], [0.5, -0.3, 0.4], [-0.1, 0.4, 0.7]]) d = 0.5 * (d + d.T) eps = 1e-6 @@ -87,7 +97,7 @@ def test_gradient_satisfies_directional_derivative_identity(self, numpy_ctx, p): def test_gradient_at_zero_is_zero(self, numpy_ctx): X = sc.HermitianSpace(3, ctx=numpy_ctx) for p in (1.0, 2.0, 3.0): - g = to_numpy(sc.SpectralLpNormFunctional(X, p).grad(X.zeros())) + g = to_numpy(_schatten(X, p).grad(X.zeros())) assert np.all(np.isfinite(g)) np.testing.assert_allclose(g, 0.0) @@ -98,7 +108,7 @@ def test_reduces_to_coordinate_lp_on_elementwise_jordan(self, numpy_ctx): # spectral p-norm coincides with the coordinate p-norm. J = sc.EuclideanElementwiseJordanSpace((4,), numpy_ctx) v = numpy_ctx.asarray([1.0, -2.0, 0.5, 3.0]) - spectral = sc.SpectralLpNormFunctional(J, 1.5) + spectral = _schatten(J, 1.5) coordinate = sc.LpNormFunctional(J, 1.5) np.testing.assert_allclose(to_numpy(spectral.value(v)), to_numpy(coordinate.value(v))) np.testing.assert_allclose(to_numpy(spectral.grad(v)), to_numpy(coordinate.grad(v))) @@ -106,19 +116,19 @@ def test_reduces_to_coordinate_lp_on_elementwise_jordan(self, numpy_ctx): def test_rejects_non_jordan_domain(self, numpy_ctx): plain = sc.DenseCoordinateSpace((3,), numpy_ctx) with pytest.raises(TypeError, match="Jordan"): - sc.SpectralLpNormFunctional(plain, 2.0) + _schatten(plain, 2.0) def test_rejects_p_below_one(self, numpy_ctx): X = sc.HermitianSpace(2, ctx=numpy_ctx) with pytest.raises(ValueError): - sc.SpectralLpNormFunctional(X, 0.5) + _schatten(X, 0.5) def test_convert_preserves_p_and_value(self, numpy_ctx, numpy_f32_ctx): X = sc.HermitianSpace(3, ctx=numpy_ctx) A, m = _hermitian(numpy_ctx, _M) - f = sc.SpectralLpNormFunctional(X, 2.0) + f = _schatten(X, 2.0) g = f.convert(numpy_f32_ctx) - assert g.p == 2.0 and g.ctx == numpy_f32_ctx + assert g.base.p == 2.0 and g.ctx == numpy_f32_ctx a32 = numpy_f32_ctx.asarray(m.astype(np.float32)) np.testing.assert_allclose( to_numpy(g.value(a32)), np.linalg.norm(m, "fro"), rtol=2e-5 @@ -131,7 +141,7 @@ def test_complex_hermitian_value_and_hermitian_gradient(self, numpy_complex_ctx) m = np.array([[2.0, 1.0 + 1.0j], [1.0 - 1.0j, 3.0]], dtype=np.complex128) A = numpy_complex_ctx.asarray(m) evals = np.linalg.eigvalsh(m) - f = sc.SpectralLpNormFunctional(X, 1.0) + f = _schatten(X, 1.0) np.testing.assert_allclose(to_numpy(f.value(A)), np.sum(np.abs(evals))) g = to_numpy(f.grad(A)) np.testing.assert_allclose(g, g.conj().T, atol=1e-12) diff --git a/tests/generators/_arrays.py b/tests/generators/_arrays.py index 336cf84..40c4d89 100644 --- a/tests/generators/_arrays.py +++ b/tests/generators/_arrays.py @@ -5,7 +5,7 @@ import numpy as np -from spacecore.backend import Context +from spacecore import Context from spacecore.types import DenseArray from ._protocol import GeneratedCase diff --git a/tests/generators/_contexts.py b/tests/generators/_contexts.py index f5a02c4..195211c 100644 --- a/tests/generators/_contexts.py +++ b/tests/generators/_contexts.py @@ -6,7 +6,7 @@ import numpy as np import pytest -from spacecore.backend import Context, NumpyOps +from spacecore import Context, NumpyOps from ._protocol import GeneratedCase @@ -18,10 +18,8 @@ def _available_case( backend: str, dtype: Any, ops_factory: Callable[[], Any], - *, - check_level: str, ) -> ContextCase: - ctx = Context(ops_factory(), dtype=dtype, check_level=check_level) + ctx = Context(ops_factory(), dtype=dtype) field = "complex" if ctx.ops.is_complex_dtype(ctx.dtype) else "real" return GeneratedCase( obj=ctx, @@ -53,7 +51,7 @@ def _dtype_name(dtype: Any) -> str: return text.replace(".", "-") -def _optional_backend_cases(backend: str, check_level: str) -> tuple[ContextCase, ...]: +def _optional_backend_cases(backend: str) -> tuple[ContextCase, ...]: import spacecore.backend as backend_module class_name = {"jax": "JaxOps", "torch": "TorchOps", "cupy": "CuPyOps"}[backend] @@ -72,7 +70,7 @@ def _optional_backend_cases(backend: str, check_level: str) -> tuple[ContextCase cases: list[ContextCase] = [] try: for dtype in dtypes: - case = _available_case(backend, dtype, ops_type, check_level=check_level) + case = _available_case(backend, dtype, ops_type) assert case.obj is not None case.obj.asarray(np.zeros((1,), dtype=dtype)) cases.append(case) @@ -85,16 +83,15 @@ def context_cases( *, include_optional: bool = True, include_unavailable: bool = True, - check_level: str = "standard", ) -> tuple[ContextCase, ...]: """Generate supported backend/dtype contexts with explicit optional-backend skips.""" cases: list[ContextCase] = [ - _available_case("numpy", np.float64, NumpyOps, check_level=check_level), - _available_case("numpy", np.complex128, NumpyOps, check_level=check_level), + _available_case("numpy", np.float64, NumpyOps), + _available_case("numpy", np.complex128, NumpyOps), ] if include_optional: for backend in ("jax", "torch", "cupy"): - optional = _optional_backend_cases(backend, check_level) + optional = _optional_backend_cases(backend) if include_unavailable: cases.extend(optional) else: diff --git a/tests/generators/_hermitian.py b/tests/generators/_hermitian.py index 507607c..5d8120b 100644 --- a/tests/generators/_hermitian.py +++ b/tests/generators/_hermitian.py @@ -4,7 +4,7 @@ import numpy as np -from spacecore.backend import Context +from spacecore import Context from spacecore.types import DenseArray from ._arrays import _field, _numpy_dtype, _shape_id diff --git a/tests/generators/_metrics.py b/tests/generators/_metrics.py index 0d1b23a..083beac 100644 --- a/tests/generators/_metrics.py +++ b/tests/generators/_metrics.py @@ -2,7 +2,7 @@ import numpy as np -from spacecore.backend import Context +from spacecore import Context from spacecore.types import DenseArray from ._arrays import _field, _numpy_dtype diff --git a/tests/generators/_trees.py b/tests/generators/_trees.py index 20faf9d..a003444 100644 --- a/tests/generators/_trees.py +++ b/tests/generators/_trees.py @@ -5,7 +5,7 @@ import numpy as np -from spacecore.backend import Context +from spacecore import Context from spacecore.space import DenseCoordinateSpace, TreeSpace from ._arrays import dense_array_case diff --git a/tests/generators/functionals.py b/tests/generators/functionals.py index de32e9a..bdd976b 100644 --- a/tests/generators/functionals.py +++ b/tests/generators/functionals.py @@ -14,8 +14,10 @@ NUMPY_FUNCTIONAL_DTYPES = (np.float64, np.complex128) -def _context(dtype: Any, check_level: sc.CheckLevel | str) -> sc.Context: - return sc.Context(sc.NumpyOps(), dtype=dtype, check_level=check_level) +def _context(dtype: Any, check_level: sc.CheckLevel | str | None = None) -> sc.Context: + # check_level applies to the constructed objects via an ambient + # ``use_check_level`` scope in the case builders, not the Context. + return sc.Context(sc.NumpyOps(), dtype=dtype) def _target_dtype(dtype: Any) -> np.dtype[Any]: @@ -244,7 +246,7 @@ def _battery_cases( def _spectral_case(dtype: Any, check_level: sc.CheckLevel | str) -> FunctionalCase: - """Generated case for the spectral (Schatten) p-norm on a Hermitian space. + """Generated case for the Schatten p-norm as a spectral lift of the coordinate p-norm. Uses ``p = 2`` so the value is the Frobenius norm and the gradient is ``X / ||X||_F`` -- both computable without an eigendecomposition, giving an @@ -259,7 +261,7 @@ def _spectral_case(dtype: Any, check_level: sc.CheckLevel | str) -> FunctionalCa frobenius = float(np.linalg.norm(m, "fro")) target_ctx = _context(_target_dtype(dtype), check_level) return FunctionalCase( - obj=sc.SpectralLpNormFunctional(domain, 2.0, ctx), + obj=sc.spectralize(domain, lambda t: sc.LpNormFunctional(t, 2.0), ctx), reference={ "kind": "spectral-lp-norm", "domain": domain, @@ -428,7 +430,7 @@ def _algebra_case( gradient: np.ndarray, id_stub: str, ) -> FunctionalCase: - """Build a euclidean case for a lazy functional-algebra node (W4). + """Build a euclidean case for a lazy functional-algebra node. ``build(base, domain, ctx)`` returns the algebra functional wrapping the linear ``base`` (whose Riesz gradient is its representer). Euclidean geometry @@ -459,9 +461,14 @@ def _algebra_case( def _algebra_cases(dtype: Any, check_level: sc.CheckLevel | str) -> tuple[FunctionalCase, ...]: - """Generated cases for the W4 functional-algebra nodes (Scaled/Sum/Shifted/Zero).""" + """Generated cases for the functional-algebra nodes. + + Covers Scaled/Sum/Shifted/Zero plus the multiplicative nodes + Constant/Product. + """ c2 = np.asarray([0.5, -0.25, 1.0], dtype=dtype) offset = np.asarray(0.75, dtype=dtype) + constant = np.asarray(1.25, dtype=dtype) return ( _algebra_case( dtype, check_level, kind="scaled-functional", id_stub="scaled-functional", @@ -491,6 +498,87 @@ def _algebra_cases(dtype: Any, check_level: sc.CheckLevel | str) -> tuple[Functi value=lambda x, c: np.asarray(0.0, dtype=dtype), gradient=lambda x, c: np.zeros(3, dtype=dtype), ), + # Also kind "zero": a constant functional has constant value, so the + # directional-derivative law has nothing to compare (both sides vanish). + _algebra_case( + dtype, check_level, kind="zero", id_stub="constant-functional", + build=lambda base, domain, ctx: sc.ConstantFunctional(domain, constant, ctx), + value=lambda x, c: constant, + gradient=lambda x, c: np.zeros(3, dtype=dtype), + ), + # The product of two linear functionals: value ·, and by the + # product rule the Riesz gradient is conj()·c + conj()·c2. + # The real case is checked against finite differences by the + # directional-derivative law, which is an independent test of the rule. + _algebra_case( + dtype, check_level, kind="product-functional", id_stub="product-functional", + build=lambda base, domain, ctx: sc.ProductFunctional( + base, sc.InnerProductFunctional(ctx.asarray(c2), domain, ctx) + ), + value=lambda x, c: np.vdot(c, x) * np.vdot(c2, x), + gradient=lambda x, c: ( + np.conj(np.vdot(c2, x)) * c + np.conj(np.vdot(c, x)) * c2 + ), + ), + ) + + +def _spectral_wrapper_case(dtype, check_level): + """Generated case for the general spectral lift ``SpectralFunctional``. + + Uses ``SquaredL2NormFunctional`` on the eigenvalue space, so + ``F(X) = 1/2 sum lambda_i^2 = 1/2 ||X||_F^2`` with gradient ``X`` — both + available in closed form without an eigendecomposition, giving an + independent reference (the wrapper is checked against the hand-written + ``NuclearNormFunctional`` in ``tests/functional/tools/test_spectral.py``). + """ + ctx = _context(dtype, check_level) + domain = sc.HermitianSpace(2, ctx=ctx) + m = np.asarray([[2.0, 0.5], [0.5, 3.0]], dtype=dtype) + target_ctx = _context(_target_dtype(dtype), check_level) + return FunctionalCase( + obj=sc.spectralize(domain, sc.SquaredL2NormFunctional, ctx), + reference={ + "kind": "spectral-wrapper", + "domain": domain, + "x": ctx.asarray(m), + "value": 0.5 * float(np.sum(m * m)), + "gradient": ctx.asarray(m), + "target_ctx": target_ctx, + "check_level": check_level, + }, + capabilities=frozenset({"gradient", "conversion", "euclidean", "spectral"}), + id=f"spectral-wrapper-squared-l2-{np.dtype(dtype).name}-checks-{check_level}", + ) + + +def _realified_case(dtype, check_level): + """Generated case for ``RealifiedFunctional`` over stacked real coordinates. + + ``F(v) = 1/2 ||v||^2`` on a complex space has metric gradient ``v``, so the + realified gradient is the stacked real vector itself — the identity map, + which is a reference no part of the implementation shares. + """ + ctx = _context(dtype, check_level) + base_domain = sc.DenseCoordinateSpace((2,), ctx=ctx) + base = sc.SquaredL2NormFunctional(base_domain) + obj = sc.RealifiedFunctional(base) + v = np.asarray([3.0 + 4.0j, 1.0 - 2.0j], dtype=dtype) + w = np.concatenate([v.real, v.imag]) + target_ctx = _context(np.dtype(_target_dtype(dtype)).type(0).real.dtype, check_level) + return FunctionalCase( + obj=obj, + reference={ + "kind": "realified", + "domain": obj.domain, + "x": obj.domain.ctx.asarray(w), + "value": 0.5 * float(np.sum(np.abs(v) ** 2)), + "gradient": obj.domain.ctx.asarray(w), + "target_ctx": target_ctx, + "check_level": check_level, + }, + capabilities=frozenset({"gradient", "conversion", "euclidean"}), + id=f"realified-squared-l2-{np.dtype(dtype).name}-checks-{check_level}", ) @@ -502,20 +590,24 @@ def functional_cases( """Generate deterministic scalar-functional cases with direct references.""" cases = [] for check_level in check_levels: - for dtype in dtypes: - for weighted in (False, True): - for kind in ("zero", "linear", "quadratic"): - cases.append( - _dense_case(dtype, check_level, kind=kind, weighted=weighted) - ) - cases.append(_composed_case(dtype, check_level)) - cases.append(_explicit_composed_case(dtype, check_level)) - cases.append(_matrix_free_linear_case(dtype, check_level)) - cases.append(_tree_case(dtype, check_level)) - cases.extend(_algebra_cases(dtype, check_level)) - # ADR-019 battery functionals are real-coordinate objectives; generate - # them once per check level for float64 when that dtype is requested. - if any(np.dtype(d) == np.dtype(np.float64) for d in dtypes): - cases.extend(_battery_cases(np.float64, check_level)) - cases.append(_spectral_case(np.float64, check_level)) + with sc.use_check_level(check_level): + for dtype in dtypes: + for weighted in (False, True): + for kind in ("zero", "linear", "quadratic"): + cases.append( + _dense_case(dtype, check_level, kind=kind, weighted=weighted) + ) + cases.append(_composed_case(dtype, check_level)) + cases.append(_explicit_composed_case(dtype, check_level)) + cases.append(_matrix_free_linear_case(dtype, check_level)) + cases.append(_tree_case(dtype, check_level)) + cases.extend(_algebra_cases(dtype, check_level)) + # ADR-019 battery functionals are real-coordinate objectives; generate + # them once per check level for float64 when that dtype is requested. + if any(np.dtype(d) == np.dtype(np.complex128) for d in dtypes): + cases.append(_realified_case(np.complex128, check_level)) + if any(np.dtype(d) == np.dtype(np.float64) for d in dtypes): + cases.extend(_battery_cases(np.float64, check_level)) + cases.append(_spectral_case(np.float64, check_level)) + cases.append(_spectral_wrapper_case(np.float64, check_level)) return tuple(cases) diff --git a/tests/generators/linops.py b/tests/generators/linops.py index 484b58c..40a600f 100644 --- a/tests/generators/linops.py +++ b/tests/generators/linops.py @@ -14,8 +14,11 @@ NUMPY_LINOP_DTYPES = (np.float64, np.complex128) -def _context(dtype: Any, check_level: sc.CheckLevel | str) -> sc.Context: - return sc.Context(sc.NumpyOps(), dtype=dtype, check_level=check_level) +def _context(dtype: Any, check_level: sc.CheckLevel | str | None = None) -> sc.Context: + # check_level is a property of the bound object; it is applied to the + # constructed cases via an ambient ``use_check_level`` scope in the builders, + # not carried on the Context. + return sc.Context(sc.NumpyOps(), dtype=dtype) def _numpy_dtype(ctx: sc.Context) -> np.dtype[Any]: @@ -30,7 +33,7 @@ def _target_context(ctx: sc.Context) -> sc.Context | None: return None dtype = np.dtype(ctx.dtype) target = np.complex64 if np.issubdtype(dtype, np.complexfloating) else np.float32 - return sc.Context(sc.NumpyOps(), dtype=target, check_level=ctx.check_level) + return sc.Context(sc.NumpyOps(), dtype=target) def _array(ctx: sc.Context, values: Any) -> Any: @@ -530,20 +533,21 @@ def linop_cases( """Generate all concrete public LinOp families with direct references.""" cases = [] for check_level in check_levels: - for dtype in dtypes: - ctx = _context(dtype, check_level) - cases.extend( - ( - dense_linop_case(ctx), - sparse_linop_case(ctx), - diagonal_linop_case(ctx), - matrix_free_linop_case(ctx), + with sc.use_check_level(check_level): + for dtype in dtypes: + ctx = _context(dtype) + cases.extend( + ( + dense_linop_case(ctx), + sparse_linop_case(ctx), + diagonal_linop_case(ctx), + matrix_free_linop_case(ctx), + ) ) - ) - cases.extend(_coordinate_algebra_cases(ctx)) - cases.extend(tree_linop_cases(ctx)) - if include_weighted: - cases.append(dense_linop_case(ctx, weighted=True)) + cases.extend(_coordinate_algebra_cases(ctx)) + cases.extend(tree_linop_cases(ctx)) + if include_weighted: + cases.append(dense_linop_case(ctx, weighted=True)) return tuple(cases) diff --git a/tests/generators/spaces.py b/tests/generators/spaces.py index 2fd8c1d..39cf62e 100644 --- a/tests/generators/spaces.py +++ b/tests/generators/spaces.py @@ -34,8 +34,10 @@ def _target_dtype(dtype: Any) -> np.dtype[Any]: return mapping[dtype] -def _context(dtype: Any, check_level: sc.CheckLevel | str) -> sc.Context: - return sc.Context(sc.NumpyOps(), dtype=dtype, check_level=check_level) +def _context(dtype: Any, check_level: sc.CheckLevel | str | None = None) -> sc.Context: + # check_level is applied to constructed spaces via ``use_check_level``; it is + # not carried on the Context. + return sc.Context(sc.NumpyOps(), dtype=dtype) def _scalar(ctx: sc.Context, real: float, imag: float = 0.0) -> Any: @@ -124,7 +126,8 @@ def dense_coordinate_space_cases( ctx = _context(dtype, check_level) for shape_like in shapes: shape = tuple(int(dimension) for dimension in shape_like) - space = sc.DenseCoordinateSpace(shape, ctx) + with sc.use_check_level(check_level): + space = sc.DenseCoordinateSpace(shape, ctx) reference = _dense_reference(space, rng=rng) cases.append( GeneratedCase( @@ -157,7 +160,8 @@ def dense_vector_space_cases( for dtype in dtypes: ctx = _context(dtype, check_level) for size in sizes: - space = sc.DenseVectorSpace((int(size),), ctx) + with sc.use_check_level(check_level): + space = sc.DenseVectorSpace((int(size),), ctx) cases.append( GeneratedCase( obj=space, diff --git a/tests/generators/test_registry_completeness.py b/tests/generators/test_registry_completeness.py index bfa0885..3bdc295 100644 --- a/tests/generators/test_registry_completeness.py +++ b/tests/generators/test_registry_completeness.py @@ -7,7 +7,7 @@ ``type(case.obj)`` for at least one generated case drawn from the case-factory registry. -Checklist section 11: +Contract specified here: - Every concrete public ``LinOp`` subclass is produced by ``linop_cases()``. - Every concrete public ``Space`` subclass is produced by one of the space case factories (dense, vector, inner-product, tree, jordan, mixed). diff --git a/tests/integration/test_imports.py b/tests/integration/test_imports.py index 2a89bc8..586ba9d 100644 --- a/tests/integration/test_imports.py +++ b/tests/integration/test_imports.py @@ -29,6 +29,6 @@ def test_subpackages_import(): "spacecore.backend", "spacecore.space", "spacecore.linop", - "spacecore._contextual", + "spacecore.contextual", ]: assert importlib.import_module(name) is not None diff --git a/tests/integration/test_public_api.py b/tests/integration/test_public_api.py index ce237db..12ca29e 100644 --- a/tests/integration/test_public_api.py +++ b/tests/integration/test_public_api.py @@ -95,7 +95,11 @@ def test_expected_names_are_exported(): "expm_multiply", } if has_jax(): - expected |= {"JaxOps", "jax_pytree_class"} + # ``jax_pytree_class`` was removed from the public API: pytree + # registration is no longer a JAX-named decorator but an internal, + # backend-neutral registry (``spacecore.backend._container``) that each backend + # installs its own tree protocol into. + expected |= {"JaxOps"} if has_cupy(): expected |= {"CuPyOps"} if has_torch(): @@ -110,9 +114,9 @@ def test_top_level_objects_match_source_modules(): linop = importlib.import_module("spacecore.linop") functional = importlib.import_module("spacecore.functional") linalg = importlib.import_module("spacecore.linalg") - contextual = importlib.import_module("spacecore._contextual") + contextual = importlib.import_module("spacecore.contextual") - assert sc.Context is backend.Context + assert sc.Context is contextual.Context assert sc.NumpyOps is backend.NumpyOps if has_cupy(): assert sc.CuPyOps is backend.CuPyOps @@ -145,4 +149,4 @@ def test_package_version_matches_project_metadata(): assert metadata["tool"]["setuptools"]["dynamic"]["version"]["attr"] == ( "spacecore._version.__version__" ) - assert sc.__version__ == "0.4.2" + assert sc.__version__ == "0.4.3" diff --git a/tests/integration/test_smoke_torch.py b/tests/integration/test_smoke_torch.py index 9bf6014..7bd7174 100644 --- a/tests/integration/test_smoke_torch.py +++ b/tests/integration/test_smoke_torch.py @@ -11,27 +11,30 @@ def test_torch_vector_hermitian_product_and_linop_smoke(): sc = importlib.import_module("spacecore") dt = torch_real_dtype() - ctx = sc.Context(sc.TorchOps(), dtype=dt, check_level="standard") - - X = sc.DenseCoordinateSpace((2,), ctx) - x = ctx.asarray([1.0, 2.0]) - assert np.allclose(to_numpy(X.inner(x, x)), 5.0) - - H = sc.HermitianSpace(2, atol=1e-6, rtol=1e-6, ctx=ctx) - h = ctx.asarray([[2.0, 1.0], [1.0, 2.0]]) - evals, evecs = H.spectral_decompose(h) - assert np.allclose(to_numpy(evecs @ ctx.ops.diag(evals) @ evecs.T.conj()), to_numpy(h)) - - P = sc.TreeSpace.from_leaf_spaces((X, sc.DenseCoordinateSpace((3,), ctx)), ctx) - p = (x, ctx.asarray([3.0, 4.0, 5.0])) - flat = P.flatten(p) - assert flat.shape == (5,) - assert all(ctx.ops.is_dense(part) for part in P.unflatten(flat)) - - Y = sc.DenseCoordinateSpace((3,), ctx) - A = ctx.asarray([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]) - op = sc.DenseLinOp(A, X, Y, ctx) - assert np.allclose(to_numpy(op.apply(x)), np.array([5.0, 11.0, 17.0])) + ctx = sc.Context(sc.TorchOps(), dtype=dt) + + with sc.use_check_level("standard"): + X = sc.DenseCoordinateSpace((2,), ctx) + x = ctx.asarray([1.0, 2.0]) + assert np.allclose(to_numpy(X.inner(x, x)), 5.0) + + H = sc.HermitianSpace(2, atol=1e-6, rtol=1e-6, ctx=ctx) + h = ctx.asarray([[2.0, 1.0], [1.0, 2.0]]) + evals, evecs = H.spectral_decompose(h) + assert np.allclose( + to_numpy(evecs @ ctx.ops.diag(evals) @ evecs.T.conj()), to_numpy(h) + ) + + P = sc.TreeSpace.from_leaf_spaces((X, sc.DenseCoordinateSpace((3,), ctx)), ctx) + p = (x, ctx.asarray([3.0, 4.0, 5.0])) + flat = P.flatten(p) + assert flat.shape == (5,) + assert all(ctx.ops.is_dense(part) for part in P.unflatten(flat)) + + Y = sc.DenseCoordinateSpace((3,), ctx) + A = ctx.asarray([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]) + op = sc.DenseLinOp(A, X, Y, ctx) + assert np.allclose(to_numpy(op.apply(x)), np.array([5.0, 11.0, 17.0])) def test_torch_from_operators_joins_dtypes_through_backend(): diff --git a/tests/kernels/test_kernel_dispatch.py b/tests/kernels/test_kernel_dispatch.py index 0adcb7c..3c8cda0 100644 --- a/tests/kernels/test_kernel_dispatch.py +++ b/tests/kernels/test_kernel_dispatch.py @@ -70,7 +70,8 @@ def patched_registry(monkeypatch): @pytest.fixture def strict_ctx(): - return sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="strict") + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + return sc.DenseCoordinateSpace((2,), ctx, check_level="strict") # --------------------------------------------------------------------------- @@ -323,6 +324,53 @@ def test_strict_routes_through_verify(self, patched_registry, strict_ctx): K.dispatch("k", 1, generic=lambda *a: ("good", a), ctx=strict_ctx) +class TestCallSitesPassTheBoundOperand: + """Wired call sites must hand the dispatcher the context-*bound* operand. + + ``_is_strict`` reads ``check_level``, which lives on the bound object and not + on :class:`Context`. A call site passing ``self.ctx`` therefore disables the + ADR-014 strict rule *silently*: no error, no warning, optimized kernels just + stop being cross-checked exactly when the strictest policy was requested. + The unit tests above cannot catch that — they call ``effective_mode`` + directly with a correct argument — so this pins the production wiring. + """ + + def test_composed_apply_passes_an_object_carrying_check_level(self, monkeypatch): + # Patch at the use site: algebra.py binds the name at import time, so + # patching the package attribute would not reach it. + import spacecore.kernels.core.algebra as core_algebra + + seen = [] + monkeypatch.setattr( + core_algebra, + "should_consult_dispatch", + lambda ctx=None: (seen.append(ctx), False)[1], + ) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + with sc.use_check_level("strict"): + X = sc.DenseCoordinateSpace((3,), ctx) + A = sc.DenseLinOp(np.eye(3), X, X, ctx) + (A @ A).apply(np.ones(3)) + + assert seen, "composed apply never consulted the dispatcher" + assert all(getattr(o, "check_level", None) == "strict" for o in seen) + + def test_strict_operator_reaches_verify_mode_end_to_end(self): + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + with sc.use_check_level("strict"): + X = sc.DenseCoordinateSpace((3,), ctx) + strict_op = sc.DenseLinOp(np.eye(3), X, X, ctx) + with sc.use_check_level("standard"): + Y = sc.DenseCoordinateSpace((3,), ctx) + plain_op = sc.DenseLinOp(np.eye(3), Y, Y, ctx) + + K.set_dispatch_mode("off") + assert K.effective_mode(strict_op) == "verify" + assert K.effective_mode(plain_op) == "off" + # The context alone must NOT satisfy the strict rule — that is the bug. + assert K.effective_mode(strict_op.ctx) == "off" + + # --------------------------------------------------------------------------- # Memory gate # --------------------------------------------------------------------------- diff --git a/tests/linalg/_helpers.py b/tests/linalg/_helpers.py index a90ed15..12c954f 100644 --- a/tests/linalg/_helpers.py +++ b/tests/linalg/_helpers.py @@ -28,8 +28,16 @@ def ops_for_backend(name: str): def make_ctx(backend_name: str = "numpy", dtype=np.float64, check_level: str = "none"): - """Build a solver context. Checks default to ``none`` (the solver hot path).""" - return sc.Context(ops_for_backend(backend_name), dtype=dtype, check_level=check_level) + """Build a solver context and arm ``check_level`` for the objects built from it. + + ``check_level`` lives on the bound object rather than on the ``Context``, so the + requested level is installed as the ambient default that seeds the spaces and + operators the caller constructs next. Checks default to ``none`` (the solver hot + path); the autouse fixture in ``tests/conftest.py`` restores the previous ambient + level after each test. + """ + sc.set_check_level(check_level) + return sc.Context(ops_for_backend(backend_name), dtype=dtype) def backend_params(*, cupy: bool = True): diff --git a/tests/linalg/test_core_resolution.py b/tests/linalg/test_core_resolution.py index 549c67a..09a7ae5 100644 --- a/tests/linalg/test_core_resolution.py +++ b/tests/linalg/test_core_resolution.py @@ -15,7 +15,8 @@ def make_ctx(): - return sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="standard") + sc.set_check_level("standard") + return sc.Context(sc.NumpyOps(), dtype=np.float64) class _WeightedSpace(sc.DenseCoordinateSpace): diff --git a/tests/linalg/test_generated_solver_matrix.py b/tests/linalg/test_generated_solver_matrix.py index f9335fb..ba19ab6 100644 --- a/tests/linalg/test_generated_solver_matrix.py +++ b/tests/linalg/test_generated_solver_matrix.py @@ -39,7 +39,8 @@ def _ctx(check_level: str): - return sc.Context(sc.NumpyOps(), dtype=np.float64, check_level=check_level) + sc.set_check_level(check_level) + return sc.Context(sc.NumpyOps(), dtype=np.float64) # =========================================================================== diff --git a/tests/linalg/test_solver_contracts.py b/tests/linalg/test_solver_contracts.py index ba2277d..d2c1c15 100644 --- a/tests/linalg/test_solver_contracts.py +++ b/tests/linalg/test_solver_contracts.py @@ -19,7 +19,8 @@ def _ctx(dtype=np.float64, check_level="standard"): - return sc.Context(sc.NumpyOps(), dtype=dtype, check_level=check_level) + sc.set_check_level(check_level) + return sc.Context(sc.NumpyOps(), dtype=dtype) # =========================================================================== diff --git a/tests/linalg/test_utils.py b/tests/linalg/test_utils.py index 359a3d6..3fcd06b 100644 --- a/tests/linalg/test_utils.py +++ b/tests/linalg/test_utils.py @@ -163,24 +163,30 @@ def _dense(self, ctx, matrix): return sc.DenseLinOp(ctx.asarray(matrix), space, space, ctx) def test_noop_when_not_strict(self): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="standard") - op = self._dense(ctx, np.asarray([[1.0, 2.0], [0.0, 3.0]])) # non-Hermitian + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + with sc.use_check_level("standard"): + op = self._dense(ctx, np.asarray([[1.0, 2.0], [0.0, 3.0]])) # non-Hermitian _utils.require_strict_cg_preconditions(op) # no raise: not strict def test_strict_accepts_spd(self): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="strict") - op = self._dense(ctx, np.asarray([[4.0, 1.0], [1.0, 3.0]])) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + with sc.use_check_level("strict"): + op = self._dense(ctx, np.asarray([[4.0, 1.0], [1.0, 3.0]])) _utils.require_strict_cg_preconditions(op) # no raise def test_strict_rejects_non_hermitian(self): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="strict") - op = self._dense(ctx, np.asarray([[1.0, 2.0], [0.0, 3.0]])) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + with sc.use_check_level("strict"): + op = self._dense(ctx, np.asarray([[1.0, 2.0], [0.0, 3.0]])) with pytest.raises(ValueError, match="Hermitian"): _utils.require_strict_cg_preconditions(op) def test_strict_rejects_nonpositive_curvature(self): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="strict") + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) # Hermitian but indefinite: probe vector lands on zero curvature. - op = sc.DiagonalLinOp(ctx.asarray([1.0, -1.0]), sc.DenseCoordinateSpace((2,), ctx), ctx) + with sc.use_check_level("strict"): + op = sc.DiagonalLinOp( + ctx.asarray([1.0, -1.0]), sc.DenseCoordinateSpace((2,), ctx), ctx + ) with pytest.raises(ValueError, match="positive curvature"): _utils.require_strict_cg_preconditions(op) diff --git a/tests/linops/test_algebra_factories.py b/tests/linops/test_algebra_factories.py index c9583eb..fbc384d 100644 --- a/tests/linops/test_algebra_factories.py +++ b/tests/linops/test_algebra_factories.py @@ -2,8 +2,8 @@ Checklist item 3: -* Private scalar helpers — ``is_scalar_like``, ``_conjugate_scalar``, - ``_scalar_equal``, ``_is_zero_scalar``, ``_is_one_scalar`` truth tables. +* Scalar helpers — ``is_scalar_like``, ``_conjugate_scalar``, and the shared + NaN-reflexive ``scalar_eq`` (in :mod:`spacecore._lazy_algebra`) truth tables. * ``make_sum`` — non-empty requirement, flattens nested ``SumLinOp``, drops ``ZeroLinOp`` terms, returns ``ZeroLinOp`` for all-zero, returns the single survivor when one term remains, raises on mismatched @@ -19,11 +19,10 @@ import pytest import spacecore as sc +from spacecore._lazy_algebra import is_recognizably_nonreal, scalar_eq +from tests._helpers import has_jax from spacecore.linop._algebra import ( _conjugate_scalar, - _is_one_scalar, - _is_zero_scalar, - _scalar_equal, is_scalar_like, make_composed, make_scaled, @@ -60,47 +59,95 @@ def test_numpy_complex(self): assert _conjugate_scalar(np.complex128(1 + 1j)) == np.complex128(1 - 1j) -class TestScalarEqual: - @pytest.mark.parametrize("value, target, expected", [ +class TestScalarEq: + """The shared NaN-reflexive ``scalar_eq`` that folds the old zero/one helpers. + + The zero/one canonicalization in ``make_scaled`` now routes through + ``scalar_eq(scalar, 0)`` / ``scalar_eq(scalar, 1)``; ``scalar_eq`` is + NaN-reflexive so a NaN-scaled node equals itself. + """ + + @pytest.mark.parametrize("a, b, expected", [ (0, 0, True), (0.0, 0, True), (1, 1, True), (2, 1, False), (np.float64(0.0), 0, True), + (np.float64(1.0), 1, True), + (float("nan"), float("nan"), True), # NaN-reflexive + (float("nan"), 0.0, False), ]) - def test_truth_table(self, value, target, expected): - assert _scalar_equal(value, target) is expected + def test_truth_table(self, a, b, expected): + assert scalar_eq(a, b) is expected + + def test_returns_false_when_undecidable(self): + """A ``TypeError`` means "no concrete verdict", reported as False. + + This is the abstract/traced-scalar case: nothing about the value is + knowable, so ``scalar_eq`` answers "not recognizably equal" and the + caller skips the simplification. + """ + class _Undecidable: + def __eq__(self, other): + raise TypeError("no concrete value") + + assert scalar_eq(_Undecidable(), 0) is False + + def test_propagates_a_broken_eq(self): + """A non-``TypeError`` means the operand's ``__eq__`` is broken, and propagates. - def test_returns_false_on_exception(self): - """``_scalar_equal`` swallows exceptions and returns False.""" + Every exception used to be swallowed into ``False``, which hid real + defects behind a silently disabled canonicalization. + """ class _Bad: def __eq__(self, other): raise RuntimeError("boom") - assert _scalar_equal(_Bad(), 0) is False + with pytest.raises(RuntimeError, match="boom"): + scalar_eq(_Bad(), 0) + +class TestIsRecognizablyNonreal: + """Only a *provably* non-real scalar is reported; undecidable answers False.""" -class TestIsZeroScalar: @pytest.mark.parametrize("value, expected", [ - (0, True), - (0.0, True), (1, False), - (np.float64(0.0), True), + (1.5, False), + (1 + 2j, True), + (1j, True), + (np.float64(2.0), False), + (np.complex128(2 + 3j), True), + (np.complex128(2 + 0j), False), # complex dtype, real value + (float("nan"), False), # NaN makes `!=` uninformative ]) def test_truth_table(self, value, expected): - assert _is_zero_scalar(value) is expected + assert is_recognizably_nonreal(value) is expected -class TestIsOneScalar: - @pytest.mark.parametrize("value, expected", [ - (1, True), - (1.0, True), - (0, False), - (2, False), - (np.float64(1.0), True), - ]) - def test_truth_table(self, value, expected): - assert _is_one_scalar(value) is expected +class TestScalarPredicatesUnderTracing: + """Under ``jax.jit`` a traced coefficient is undecidable, never a false verdict. + + Folding is skipped rather than misapplied: the expression tree stays larger + than a concrete one, but every node is correct — and a real scaling of a + real-scalar-field space is never spuriously rejected. + """ + + def test_traced_scalar_is_undecidable(self): + if not has_jax(): + pytest.skip("jax is not installed") + import jax + import jax.numpy as jnp + + verdicts = [] + + def probe(t): + verdicts.append( + (scalar_eq(t, 0), scalar_eq(t, 1), is_recognizably_nonreal(t)) + ) + return t + + jax.jit(probe)(jnp.asarray(1.0)) + assert verdicts == [(False, False, False)] # =========================================================================== @@ -292,16 +339,18 @@ def test_make_sum_ignores_check_level_when_dtype_matches(self): (Folded from test_algebra.py::test_factories_ignore_enable_checks_when_context_dtype_matches.) """ - checked = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="standard") - unchecked = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - X_checked = sc.DenseCoordinateSpace((2,), checked) - X_unchecked = sc.DenseCoordinateSpace((2,), unchecked) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + X_checked = sc.DenseCoordinateSpace((2,), ctx, check_level="standard") + X_unchecked = sc.DenseCoordinateSpace((2,), ctx, check_level="none") A = sc.DenseLinOp( - checked.asarray([[1.0, 0.0], [0.0, 1.0]]), X_checked, X_checked, checked, + ctx.asarray([[1.0, 0.0], [0.0, 1.0]]), + X_checked, X_checked, ctx, check_level="standard", ) B = sc.DenseLinOp( - unchecked.asarray([[2.0, 0.0], [0.0, 3.0]]), X_unchecked, X_unchecked, unchecked, + ctx.asarray([[2.0, 0.0], [0.0, 3.0]]), + X_unchecked, X_unchecked, ctx, check_level="none", ) + assert (A.check_level, B.check_level) == ("standard", "none") assert isinstance(make_sum((A, B)), sc.SumLinOp) assert isinstance(make_composed(A, B), sc.ComposedLinOp) diff --git a/tests/linops/test_algebra_linops.py b/tests/linops/test_algebra_linops.py index 2676b40..d29a3d5 100644 --- a/tests/linops/test_algebra_linops.py +++ b/tests/linops/test_algebra_linops.py @@ -44,10 +44,10 @@ def test_apply_returns_input(self, numpy_ctx): def test_apply_unchecked_returns_literal_input(self, numpy_ctx): """With checks disabled, ``apply(x)`` and ``rapply(x)`` return the literal ``x``.""" - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - X = sc.DenseCoordinateSpace((3,), ctx) - op = sc.IdentityLinOp(X, ctx) - x = ctx.asarray([1.0, 2.0, 3.0]) + X = sc.DenseCoordinateSpace((3,), numpy_ctx, check_level="none") + op = sc.IdentityLinOp(X, numpy_ctx, check_level="none") + x = numpy_ctx.asarray([1.0, 2.0, 3.0]) + assert op.check_level == "none" assert op.apply(x) is x # rapply is a distinct code branch; confirm it is also a literal pass-through. assert op.rapply(x) is x @@ -323,7 +323,11 @@ def rev(y): called.append(y) return y * 5.0 # user's chosen adjoint, untouched by Riesz - op = sc.MatrixFreeLinOp(lambda x: x, rev, X, X, numpy_ctx) + # Deliberately verbatim: this test pins that the supplied reverse is + # stored as-is. Declaring the flag states the intent and skips the advisory. + op = sc.MatrixFreeLinOp( + lambda x: x, rev, X, X, numpy_ctx, euclidean_adjoint=False + ) out = op.rapply(numpy_ctx.asarray([1.0, 2.0])) assert len(called) == 1 np.testing.assert_allclose(out, [5.0, 10.0]) @@ -435,8 +439,10 @@ def _non_euclidean_matrix_free_fixture(ctx): matrix = ctx.asarray(matrix_np) metric_adjoint = ctx.asarray(metric_adjoint_np) + # ``metric_adjoint`` is already the metric adjoint, hence the explicit flag. op = sc.MatrixFreeLinOp( lambda z: matrix @ z, lambda w: metric_adjoint @ w, domain, codomain, ctx, + euclidean_adjoint=False, ) return { "op": op, @@ -520,7 +526,10 @@ def __init__(self, shape, ctx): self.geometry = BrokenInnerProduct() space = BrokenSpace((2,), numpy_ctx) - with pytest.raises(ValueError, match="MatrixFreeLinOp.from_coordinate_adjoint"): + # The Riesz-map requirement moved into ``__init__`` behind + # ``euclidean_adjoint=True``, which this classmethod delegates to, so the + # message now names both entry points rather than only the classmethod. + with pytest.raises(ValueError, match="from_coordinate_adjoint"): sc.MatrixFreeLinOp.from_coordinate_adjoint( lambda x: x, lambda y: y, space, space, numpy_ctx ) @@ -567,7 +576,9 @@ def metric_rapply(w): def coordinate_rapply(w): return matrix.T @ w - direct = sc.MatrixFreeLinOp(apply, metric_rapply, domain, codomain, numpy_ctx) + direct = sc.MatrixFreeLinOp( + apply, metric_rapply, domain, codomain, numpy_ctx, euclidean_adjoint=False + ) wrapped = sc.MatrixFreeLinOp.from_coordinate_adjoint( apply, coordinate_rapply, domain, codomain, numpy_ctx, ) @@ -626,7 +637,7 @@ def riesz_inverse(self, x): def test_convert_preserves_direct_reverse_without_riesz(self, numpy_ctx): """Converting a direct matrix-free op keeps the user's reverse callable as-is.""" - new_ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") + new_ctx = sc.Context(sc.NumpyOps(), dtype=np.complex128) WeightedVectorSpace = _weighted_space_class() domain = WeightedVectorSpace([2.0, 5.0], numpy_ctx) codomain = WeightedVectorSpace([3.0, 7.0, 11.0], numpy_ctx) @@ -643,8 +654,11 @@ def apply(z): def rapply(w): return metric_adjoint @ w - op = sc.MatrixFreeLinOp(apply, rapply, domain, codomain, numpy_ctx) + op = sc.MatrixFreeLinOp( + apply, rapply, domain, codomain, numpy_ctx, euclidean_adjoint=False + ) converted = op.convert(new_ctx) + assert converted is not op x = new_ctx.asarray([0.25, -1.5]) y = new_ctx.asarray([2.0, -0.5, 1.25]) @@ -654,7 +668,7 @@ def rapply(w): def test_convert_preserves_batched_reverse(self, numpy_ctx): """``from_coordinate_adjoint`` with a batched coordinate adjoint survives ``convert``.""" - new_ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") + new_ctx = sc.Context(sc.NumpyOps(), dtype=np.complex128) WeightedVectorSpace = _weighted_space_class() domain = WeightedVectorSpace([2.0, 5.0], numpy_ctx) codomain = WeightedVectorSpace([3.0, 7.0, 11.0], numpy_ctx) @@ -674,6 +688,7 @@ def rvapply(ws): coordinate_rvapply=rvapply, ) converted = op.convert(new_ctx) + assert converted is not op ys = new_ctx.asarray([[2.0, -0.5, 1.25], [-1.0, 3.0, 0.75]]) expected = np.stack([to_numpy(converted.rapply(y)) for y in ys], axis=0) @@ -682,7 +697,7 @@ def rvapply(ws): def test_convert_without_rvapply_uses_fallback(self, numpy_ctx): """When no batched coordinate adjoint is given, ``rvapply_fn`` stays ``None``.""" - new_ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") + new_ctx = sc.Context(sc.NumpyOps(), dtype=np.complex128) WeightedVectorSpace = _weighted_space_class() domain = WeightedVectorSpace([2.0, 5.0], numpy_ctx) codomain = WeightedVectorSpace([3.0, 7.0, 11.0], numpy_ctx) @@ -691,6 +706,7 @@ def test_convert_without_rvapply_uses_fallback(self, numpy_ctx): lambda z: matrix @ z, lambda w: matrix.T @ w, domain, codomain, numpy_ctx, ) converted = op.convert(new_ctx) + assert converted is not op ys = new_ctx.asarray([[2.0, -0.5, 1.25], [-1.0, 3.0, 0.75]]) expected = np.stack([to_numpy(converted.rapply(y)) for y in ys], axis=0) @@ -704,9 +720,11 @@ def test_convert_without_rvapply_uses_fallback(self, numpy_ctx): # (folded from test_algebra_linop.py) # =========================================================================== class TestAlgebraClassHierarchy: - def _op(self, ctx): - X = sc.DenseCoordinateSpace((2,), ctx) - return sc.DenseLinOp(ctx.asarray([[1.0, 2.0], [3.0, 4.0]]), X, X, ctx) + def _op(self, ctx, check_level=None): + X = sc.DenseCoordinateSpace((2,), ctx, check_level=check_level) + return sc.DenseLinOp( + ctx.asarray([[1.0, 2.0], [3.0, 4.0]]), X, X, ctx, check_level=check_level, + ) def test_no_adjoint_linop_symbol_exported(self): """There is no public ``AdjointLinOp``; ``A.H`` returns the private view.""" @@ -732,12 +750,11 @@ def test_classes_subclass_linop(self): assert issubclass(sc.IdentityLinOp, sc.LinOp) assert issubclass(sc.MatrixFreeLinOp, sc.LinOp) - def test_check_policy_mismatch_does_not_block_algebra(self): + def test_check_policy_mismatch_does_not_block_algebra(self, numpy_ctx): """Operands with differing ``check_level`` still combine (dtype matches).""" - checked = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="standard") - unchecked = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - A = self._op(checked) - B = self._op(unchecked) + A = self._op(numpy_ctx, check_level="standard") + B = self._op(numpy_ctx, check_level="none") + assert (A.check_level, B.check_level) == ("standard", "none") assert isinstance(A + B, sc.SumLinOp) assert isinstance(A @ B, sc.ComposedLinOp) @@ -967,17 +984,19 @@ def test_jit_algebra_expression_matches_eager(self): from tests._helpers import jax_real_dtype - ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") - X = sc.DenseCoordinateSpace((2,), ctx) - Y = sc.DenseCoordinateSpace((3,), ctx) - A = sc.DenseLinOp( - ctx.asarray([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]), X, Y, ctx, - ) - B = sc.DenseLinOp( - ctx.asarray([[0.5, -1.0], [2.0, 1.0], [-0.5, 3.0]]), X, Y, ctx, - ) - C = sc.DenseLinOp(ctx.asarray([[2.0, -1.0], [0.25, 1.5]]), X, X, ctx) + ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) + with sc.use_check_level("none"): + X = sc.DenseCoordinateSpace((2,), ctx) + Y = sc.DenseCoordinateSpace((3,), ctx) + A = sc.DenseLinOp( + ctx.asarray([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]), X, Y, ctx, + ) + B = sc.DenseLinOp( + ctx.asarray([[0.5, -1.0], [2.0, 1.0], [-0.5, 3.0]]), X, Y, ctx, + ) + C = sc.DenseLinOp(ctx.asarray([[2.0, -1.0], [0.25, 1.5]]), X, X, ctx) expr = (2 * A + B) @ C + assert expr.check_level == "none" x = ctx.asarray([1.0, -2.0]) apply_jit = jax.jit(lambda op, z: op.apply(z)) diff --git a/tests/linops/test_block_diagonal_linop.py b/tests/linops/test_block_diagonal_linop.py index 08342ac..741e1c3 100644 --- a/tests/linops/test_block_diagonal_linop.py +++ b/tests/linops/test_block_diagonal_linop.py @@ -350,8 +350,7 @@ def test_algebra_and_block_validation(self, numpy_ctx): with pytest.raises(TypeError, match="every block"): sc.BlockDiagonalLinOp((block, object())) - other_ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="cheap") - other_x = sc.DenseCoordinateSpace((1,), other_ctx) + other_x = sc.DenseCoordinateSpace((1,), numpy_ctx, check_level="cheap") other = sc.IdentityLinOp(other_x) with pytest.raises(ValueError, match="check policy"): sc.BlockDiagonalLinOp((block, other)) @@ -363,12 +362,14 @@ def test_from_empty_operators_raises(self): def test_batch_checks_reject_wrong_tuple_layout(self): # Folded from tests/linops/test_tree_linop_batching.py. - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="standard") - x1, x2 = sc.DenseCoordinateSpace((2,), ctx), sc.DenseCoordinateSpace((1,), ctx) - y1, y2 = sc.DenseCoordinateSpace((1,), ctx), sc.DenseCoordinateSpace((2,), ctx) - A1 = sc.DenseLinOp(ctx.asarray([[1.0, 2.0]]), x1, y1, ctx) - A2 = sc.DenseLinOp(ctx.asarray([[3.0], [-1.0]]), x2, y2, ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + with sc.use_check_level("standard"): + x1, x2 = sc.DenseCoordinateSpace((2,), ctx), sc.DenseCoordinateSpace((1,), ctx) + y1, y2 = sc.DenseCoordinateSpace((1,), ctx), sc.DenseCoordinateSpace((2,), ctx) + A1 = sc.DenseLinOp(ctx.asarray([[1.0, 2.0]]), x1, y1, ctx) + A2 = sc.DenseLinOp(ctx.asarray([[3.0], [-1.0]]), x2, y2, ctx) op = sc.BlockDiagonalLinOp.from_operators((A1, A2)) + assert op.check_level == "standard" with pytest.raises(ValueError, match="structure"): op.vapply((ctx.asarray([[1.0, 2.0], [-1.0, 0.5]]),)) diff --git a/tests/linops/test_block_matrix_linop.py b/tests/linops/test_block_matrix_linop.py index 46049d7..2cbe025 100644 --- a/tests/linops/test_block_matrix_linop.py +++ b/tests/linops/test_block_matrix_linop.py @@ -150,7 +150,6 @@ def test_rejects_invalid_block_layouts_and_contexts(self, numpy_ctx): with pytest.raises(ValueError, match="column 1"): sc.BlockMatrixLinOp((rows[0], (rows[1][0], incompatible_column))) - other_ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="cheap") - other_space = sc.DenseCoordinateSpace((1,), other_ctx) + other_space = sc.DenseCoordinateSpace((1,), numpy_ctx, check_level="cheap") with pytest.raises(ValueError, match="check policy"): sc.BlockMatrixLinOp(((rows[0][0], sc.IdentityLinOp(other_space)),)) diff --git a/tests/linops/test_dense_linop.py b/tests/linops/test_dense_linop.py index 8ac1b0c..bb053c5 100644 --- a/tests/linops/test_dense_linop.py +++ b/tests/linops/test_dense_linop.py @@ -268,7 +268,7 @@ def test_fused_mode_selected_and_matches_generic_metric(self, numpy_ctx): def test_fused_mode_recomputed_after_convert(self, numpy_ctx): # Source: legacy test_adjoint_identity.py - new_ctx = sc.Context(sc.NumpyOps(), dtype=np.float32, check_level="none") + new_ctx = sc.Context(sc.NumpyOps(), dtype=np.float32) domain = sc.DenseCoordinateSpace( (2,), numpy_ctx, geometry=sc.WeightedInnerProduct(numpy_ctx.asarray([2.0, 5.0])) ) @@ -487,11 +487,12 @@ def test_tree_flatten_unflatten_round_trip(self, numpy_ctx): class TestBatched: def test_fast_path_vapply_rvapply_without_checks(self): # Source: legacy test_batched_lifting.py - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - dom = sc.DenseCoordinateSpace((2,), ctx) - cod = sc.DenseCoordinateSpace((3,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + dom = sc.DenseCoordinateSpace((2,), ctx, check_level="none") + cod = sc.DenseCoordinateSpace((3,), ctx, check_level="none") matrix = ctx.asarray(_MATRIX) - op = sc.DenseLinOp(matrix, dom, cod, ctx) + op = sc.DenseLinOp(matrix, dom, cod, ctx, check_level="none") + assert op.check_level == "none" xs = ctx.asarray([[7.0, 8.0], [1.0, -1.0], [0.5, 2.0]]) ys = ctx.asarray([[1.0, -1.0, 2.0], [0.0, 3.0, -2.0]]) diff --git a/tests/linops/test_fused_algebra_overhead.py b/tests/linops/test_fused_algebra_overhead.py index 5b569a3..b137413 100644 --- a/tests/linops/test_fused_algebra_overhead.py +++ b/tests/linops/test_fused_algebra_overhead.py @@ -46,8 +46,9 @@ def counting(x): @pytest.fixture def ctx(): # check_level="standard" is the default; make it explicit so the per-call - # validation under test is actually active. - return sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="standard") + # validation under test is actually active on every object built here. + with sc.use_check_level("standard"): + yield sc.Context(sc.NumpyOps(), dtype=np.float64) def _dense(ctx, X, Y, seed): @@ -214,9 +215,11 @@ def test_composed_chain_is_jit_safe(): pytest.importorskip("jax") import jax - jctx = sc.Context(sc.JaxOps(), dtype=np.float32, check_level="none") - X = sc.DenseCoordinateSpace((3,), jctx) - chain, ops = _endo_chain(jctx, X, range(4)) + jctx = sc.Context(sc.JaxOps(), dtype=np.float32) + with sc.use_check_level("none"): + X = sc.DenseCoordinateSpace((3,), jctx) + chain, ops = _endo_chain(jctx, X, range(4)) + assert chain.check_level == "none" x = jctx.asarray(np.asarray([1.0, 2.0, 3.0], dtype=np.float32)) jitted = jax.jit(lambda op, v: op.apply(v)) diff --git a/tests/linops/test_linop_jit.py b/tests/linops/test_linop_jit.py index 2d56f73..e833105 100644 --- a/tests/linops/test_linop_jit.py +++ b/tests/linops/test_linop_jit.py @@ -27,8 +27,15 @@ pytestmark = pytest.mark.skipif(not has_jax(), reason="jax is not installed") +@pytest.fixture(autouse=True) +def _unchecked_objects(): + """Every object built in this module carries ``check_level="none"``.""" + with sc.use_check_level("none"): + yield + + def _jax_ctx(): - return sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") + return sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) # =========================================================================== diff --git a/tests/linops/test_metric_helpers.py b/tests/linops/test_metric_helpers.py index 1fac5f7..123f798 100644 --- a/tests/linops/test_metric_helpers.py +++ b/tests/linops/test_metric_helpers.py @@ -92,11 +92,11 @@ def inner(self, ops, x, y): class _CustomInnerSpace(sc.DenseCoordinateSpace): pass - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - X = sc.DenseCoordinateSpace((2,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + X = sc.DenseCoordinateSpace((2,), ctx, check_level="none") # Build a space whose geometry advertises non-Euclidean but lacks # the inherited Riesz machinery. - bad = _CustomInnerSpace((2,), ctx, geometry=_CustomGeometry()) + bad = _CustomInnerSpace((2,), ctx, geometry=_CustomGeometry(), check_level="none") with pytest.raises(TypeError, match="(?i)non-euclidean.*requires Riesz"): _requires_euclidean_or_riesz(bad, X, "my_op") diff --git a/tests/linops/test_sparse_linop.py b/tests/linops/test_sparse_linop.py index b2209d6..59830d0 100644 --- a/tests/linops/test_sparse_linop.py +++ b/tests/linops/test_sparse_linop.py @@ -319,11 +319,12 @@ def test_convert_preserves_action_and_converts_sparse_storage( # =========================================================================== class TestBatchedLifting: def test_fast_paths_without_checks(self): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - dom = sc.DenseCoordinateSpace((2,), ctx) - cod = sc.DenseCoordinateSpace((3,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + dom = sc.DenseCoordinateSpace((2,), ctx, check_level="none") + cod = sc.DenseCoordinateSpace((3,), ctx, check_level="none") sparse = ctx.assparse(sps.csr_matrix([[1.0, 0.0], [0.0, 4.0], [5.0, 6.0]])) - op = sc.SparseLinOp(sparse, dom, cod, ctx) + op = sc.SparseLinOp(sparse, dom, cod, ctx, check_level="none") + assert op.check_level == "none" xs = ctx.asarray([[7.0, 8.0], [1.0, -1.0], [0.5, 2.0]]) ys = ctx.asarray([[1.0, -1.0, 2.0], [0.0, 3.0, -2.0]]) @@ -339,14 +340,15 @@ class TestJit: def test_jit_apply_and_rapply(self): import jax - ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") + ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) op = sc.SparseLinOp( ctx.assparse( np.asarray([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]], dtype=jax_real_dtype()) ), - sc.DenseCoordinateSpace((2,), ctx), - sc.DenseCoordinateSpace((3,), ctx), + sc.DenseCoordinateSpace((2,), ctx, check_level="none"), + sc.DenseCoordinateSpace((3,), ctx, check_level="none"), ctx, + check_level="none", ) x = ctx.asarray([7.0, 8.0]) diff --git a/tests/linops/test_stacked_linop.py b/tests/linops/test_stacked_linop.py index f05a442..88f3c56 100644 --- a/tests/linops/test_stacked_linop.py +++ b/tests/linops/test_stacked_linop.py @@ -270,13 +270,15 @@ def test_jit_apply_and_rapply(self): # Folded from tests/linops/test_linop_jit.py (test_product_linops_jit_compile). import jax - ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") - X = sc.DenseCoordinateSpace((2,), ctx) - Y1 = sc.DenseCoordinateSpace((2,), ctx) - Y2 = sc.DenseCoordinateSpace((1,), ctx) - A1 = _dense(ctx, [[1.0, 2.0], [3.0, 4.0]], X, Y1) - A2 = _dense(ctx, [[5.0, 6.0]], X, Y2) + ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) + with sc.use_check_level("none"): + X = sc.DenseCoordinateSpace((2,), ctx) + Y1 = sc.DenseCoordinateSpace((2,), ctx) + Y2 = sc.DenseCoordinateSpace((1,), ctx) + A1 = _dense(ctx, [[1.0, 2.0], [3.0, 4.0]], X, Y1) + A2 = _dense(ctx, [[5.0, 6.0]], X, Y2) op = sc.StackedLinOp.from_operators((A1, A2)) + assert op.check_level == "none" x = ctx.asarray([7.0, 8.0]) apply_jit = jax.jit(lambda Aop, z: Aop.apply(z)) diff --git a/tests/linops/test_sum_to_single_linop.py b/tests/linops/test_sum_to_single_linop.py index 3e6f429..61a3f1c 100644 --- a/tests/linops/test_sum_to_single_linop.py +++ b/tests/linops/test_sum_to_single_linop.py @@ -274,11 +274,13 @@ def test_jit_apply_and_rapply(self): # Folded from tests/linops/test_linop_jit.py (test_product_linops_jit_compile). import jax - ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") - X = sc.DenseCoordinateSpace((2,), ctx) - Y = sc.DenseCoordinateSpace((2,), ctx) - A1 = _dense(ctx, [[1.0, 2.0], [3.0, 4.0]], X, Y) + ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) + with sc.use_check_level("none"): + X = sc.DenseCoordinateSpace((2,), ctx) + Y = sc.DenseCoordinateSpace((2,), ctx) + A1 = _dense(ctx, [[1.0, 2.0], [3.0, 4.0]], X, Y) op = sc.SumToSingleLinOp.from_operators((A1, A1)) + assert op.check_level == "none" x = ctx.asarray([7.0, 8.0]) sum_apply = jax.jit(lambda Aop, a, b: Aop.apply((a, b))) diff --git a/tests/linops/test_tree_helpers.py b/tests/linops/test_tree_helpers.py index 132651a..861ee45 100644 --- a/tests/linops/test_tree_helpers.py +++ b/tests/linops/test_tree_helpers.py @@ -55,12 +55,14 @@ def test_rejects_blocks_with_different_dtype(self, numpy_ctx, numpy_f32_ctx): ) def test_rejects_blocks_with_different_check_level(self, numpy_ctx): - X = sc.DenseCoordinateSpace((2,), numpy_ctx) - cheap_ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="cheap") - Xc = sc.DenseCoordinateSpace((2,), cheap_ctx) + X = sc.DenseCoordinateSpace((2,), numpy_ctx, check_level="standard") + Xc = sc.DenseCoordinateSpace((2,), numpy_ctx, check_level="cheap") with pytest.raises(ValueError, match="same check policy"): _validate_blocks( - (sc.IdentityLinOp(X, numpy_ctx), sc.IdentityLinOp(Xc, cheap_ctx)), + ( + sc.IdentityLinOp(X, numpy_ctx, check_level="standard"), + sc.IdentityLinOp(Xc, numpy_ctx, check_level="cheap"), + ), "TestOwner", ) diff --git a/tests/optim/test_cached_member_checks.py b/tests/optim/test_cached_member_checks.py index 710bf6d..4c85e50 100644 --- a/tests/optim/test_cached_member_checks.py +++ b/tests/optim/test_cached_member_checks.py @@ -101,8 +101,9 @@ def test_validation_decisions_unchanged_by_cache(): def test_check_level_none_bypasses_validation(): """``checked_method`` fast path: ``check_level="none"`` calls method directly.""" - ctx_none = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - space = sc.DenseCoordinateSpace((4,), ctx_none) + ctx_none = sc.Context(sc.NumpyOps(), dtype=np.float64) + space = sc.DenseCoordinateSpace((4,), ctx_none, check_level="none") + assert space.check_level == "none" # A shape-wrong element would normally raise — but the fast path skips # validation entirely. The actual arithmetic still runs. bad = ctx_none.asarray(np.zeros(5)) @@ -114,8 +115,9 @@ def test_check_level_none_bypasses_validation(): def test_check_level_cheap_validates_shape_dtype_backend(): """At ``cheap`` level, the cached checks still catch shape mismatches.""" - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="cheap") - space = sc.DenseCoordinateSpace((4,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + space = sc.DenseCoordinateSpace((4,), ctx, check_level="cheap") + assert space.check_level == "cheap" bad = ctx.asarray(np.zeros(5)) with pytest.raises(Exception): space.check_member(bad) @@ -142,16 +144,18 @@ def test_inner_product_space_caches_consistent_under_repeated_ops(): assert float(space.inner(x, y)) == pytest.approx(expected_inner) -def test_check_level_change_via_new_context_uses_fresh_space(): - """``space.convert(new_ctx)`` produces a fresh-cache space. +def test_check_level_change_via_new_space_uses_fresh_cache(): + """A space built at a different ``check_level`` has a fresh cache. - Spaces are immutable; switching ``check_level`` happens by creating a - new context and a new space. The new space has its own cache. + Spaces are immutable; ``check_level`` is a property of the space, so + switching it happens by constructing a new space. The new space has its + own cache. """ - ctx_std = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="standard") - ctx_cheap = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="cheap") - space_std = sc.DenseCoordinateSpace((4,), ctx_std) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + space_std = sc.DenseCoordinateSpace((4,), ctx, check_level="standard") space_std.member_checks() - space_cheap = space_std.convert(ctx_cheap) + space_cheap = sc.DenseCoordinateSpace((4,), ctx, check_level="cheap") assert space_cheap is not space_std + assert space_std.check_level == "standard" + assert space_cheap.check_level == "cheap" assert space_cheap._cached_member_checks is None diff --git a/tests/spaces/test_check_batched.py b/tests/spaces/test_check_batched.py index 7040470..f4610b6 100644 --- a/tests/spaces/test_check_batched.py +++ b/tests/spaces/test_check_batched.py @@ -157,17 +157,13 @@ class TestJaxVectorized: def test_jax_vectorized_checks_work_eager_and_jit_with_checks_disabled(self): import jax - eager_ctx = sc.Context( - sc.JaxOps(), dtype=jax_complex_dtype(), check_level="standard" - ) - eager_space = sc.HermitianSpace(2, ctx=eager_ctx) + eager_ctx = sc.Context(sc.JaxOps(), dtype=jax_complex_dtype()) + eager_space = sc.HermitianSpace(2, ctx=eager_ctx, check_level="standard") x = eager_ctx.asarray(np.broadcast_to(np.eye(2), (4, 2, 2)).copy()) _check_batched(eager_space, x) - jit_ctx = sc.Context( - sc.JaxOps(), dtype=jax_real_dtype(), check_level="none" - ) - jit_space = sc.DenseCoordinateSpace((2,), ctx=jit_ctx) + jit_ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) + jit_space = sc.DenseCoordinateSpace((2,), ctx=jit_ctx, check_level="none") @jax.jit def add_batch(xs): diff --git a/tests/spaces/test_dense_coordinate_space.py b/tests/spaces/test_dense_coordinate_space.py index 20cf0d6..aa9e277 100644 --- a/tests/spaces/test_dense_coordinate_space.py +++ b/tests/spaces/test_dense_coordinate_space.py @@ -182,7 +182,7 @@ def test_convert_to_jax_backend(self): pytest.skip("jax not installed") dt = jax_real_dtype() src = sc.Context(sc.NumpyOps(), dtype=dt) - dst = sc.Context(sc.JaxOps(), dtype=dt, check_level="none") + dst = sc.Context(sc.JaxOps(), dtype=dt) space = sc.DenseCoordinateSpace((2, 3), src) out = space.convert(dst) assert out.ctx == dst diff --git a/tests/spaces/test_elementwise_jordan_spaces.py b/tests/spaces/test_elementwise_jordan_spaces.py index 6c5c6fb..290b764 100644 --- a/tests/spaces/test_elementwise_jordan_spaces.py +++ b/tests/spaces/test_elementwise_jordan_spaces.py @@ -122,9 +122,9 @@ def test_round_trip_restores_euclidean(self, numpy_ctx, numpy_complex_ctx): @pytest.mark.skipif(not has_jax(), reason="jax is not installed") def test_jax_pytree_revalidates_invariant(self): """JAX unflatten with a complex ctx must refuse the Euclidean class.""" - real_ctx = sc.Context(sc.JaxOps(), dtype=np.float32, check_level="none") - complex_ctx = sc.Context(sc.JaxOps(), dtype=np.complex64, check_level="none") - space = sc.EuclideanElementwiseJordanSpace((2,), real_ctx) + real_ctx = sc.Context(sc.JaxOps(), dtype=np.float32) + complex_ctx = sc.Context(sc.JaxOps(), dtype=np.complex64) + space = sc.EuclideanElementwiseJordanSpace((2,), real_ctx, check_level="none") import jax leaves, treedef = jax.tree_util.tree_flatten(space) diff --git a/tests/spaces/test_generated_dense_coordinate.py b/tests/spaces/test_generated_dense_coordinate.py index 12860c6..ba427b0 100644 --- a/tests/spaces/test_generated_dense_coordinate.py +++ b/tests/spaces/test_generated_dense_coordinate.py @@ -117,8 +117,8 @@ def test_generated_dense_coordinate_conversion_preserves_structure(case): [(np.float32, "real"), (np.float64, "real"), (np.complex64, "complex"), (np.complex128, "complex")], ) def test_generated_dense_coordinate_field_and_exact_dtype_are_distinct(dtype, field): - ctx = sc.Context(sc.NumpyOps(), dtype=dtype, check_level="cheap") - space = sc.DenseCoordinateSpace((2,), ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=dtype) + space = sc.DenseCoordinateSpace((2,), ctx, check_level="cheap") assert space.field == field if field == "real": diff --git a/tests/spaces/test_hermitian_space.py b/tests/spaces/test_hermitian_space.py index 02df198..d71d816 100644 --- a/tests/spaces/test_hermitian_space.py +++ b/tests/spaces/test_hermitian_space.py @@ -41,6 +41,65 @@ def test_is_jordan_star_inner_product_space(self, numpy_ctx): assert isinstance(H, sc.EuclideanJordanAlgebraSpace) +# =========================================================================== +# Scalar field — closed under real scaling only +# =========================================================================== +class TestScalarField: + """``Herm(n)`` has complex *entries* but is a **real** vector space. + + ``(aH)* = conj(a) H``, so Hermitian structure survives real scaling only; + ``i H`` is anti-Hermitian. ``field`` reports the entry dtype (used by + equality and repr), ``scalar_field`` reports what the space is closed under. + """ + + def test_complex_dtype_still_has_real_scalar_field(self, numpy_complex_ctx): + H = sc.HermitianSpace(2, ctx=numpy_complex_ctx) + assert H.field == "complex" + assert H.scalar_field == "real" + + def test_real_dtype_agrees_on_both(self, numpy_ctx): + H = sc.HermitianSpace(2, ctx=numpy_ctx) + assert H.field == "real" + assert H.scalar_field == "real" + + def test_coordinate_space_scalar_field_defaults_to_field(self, numpy_complex_ctx): + X = sc.DenseVectorSpace((3,), ctx=numpy_complex_ctx) + assert X.scalar_field == X.field == "complex" + + def test_real_scaling_stays_in_the_space(self, numpy_complex_ctx): + H = sc.HermitianSpace(2, ctx=numpy_complex_ctx) + x = numpy_complex_ctx.asarray([[1 + 0j, 2 - 1j], [2 + 1j, 3 + 0j]]) + out = H.scale(2.0, x) + H.check_member(out) + assert np.allclose(to_numpy(out), 2.0 * to_numpy(x)) + + @pytest.mark.parametrize("bad", [1j, 2 + 3j, np.complex128(1j)]) + def test_non_real_scaling_is_rejected(self, numpy_complex_ctx, bad): + """The failure is caught at the source, not at a later membership check.""" + H = sc.HermitianSpace(2, ctx=numpy_complex_ctx) + x = numpy_complex_ctx.asarray([[1 + 0j, 2 - 1j], [2 + 1j, 3 + 0j]]) + with pytest.raises(sc.SpaceValidationError, match="over ℝ"): + H.scale(bad, x) + + def test_non_real_batch_scaling_is_rejected(self, numpy_complex_ctx): + H = sc.HermitianSpace(2, ctx=numpy_complex_ctx) + x = numpy_complex_ctx.asarray([[[1 + 0j, 2 - 1j], [2 + 1j, 3 + 0j]]]) + with pytest.raises(sc.SpaceValidationError, match="over ℝ"): + H.scale_batch(1j, x) + + def test_complex_scalar_really_would_have_left_the_space(self, numpy_complex_ctx): + """Pins the reason for the guard: the unchecked product is not Hermitian.""" + H = sc.HermitianSpace(2, ctx=numpy_complex_ctx) + x = numpy_complex_ctx.asarray([[1 + 0j, 2 - 1j], [2 + 1j, 3 + 0j]]) + assert not H.is_hermitian(1j * to_numpy(x)) + + def test_check_level_none_opts_out(self, numpy_complex_ctx): + """Consistent with every other membership check: ``"none"`` disables it.""" + H = sc.HermitianSpace(2, ctx=numpy_complex_ctx, check_level="none") + x = numpy_complex_ctx.asarray([[1 + 0j, 2 - 1j], [2 + 1j, 3 + 0j]]) + assert np.allclose(to_numpy(H.scale(1j, x)), 1j * to_numpy(x)) + + # =========================================================================== # Membership + symmetrize # =========================================================================== @@ -104,8 +163,8 @@ def test_spectral_decompose_reconstruction(self): """A·v_i = λ_i·v_i identity. Run at check_level=none to avoid the strict Hermitian membership gate on the reconstructed matrix. """ - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - H = sc.HermitianSpace(3, ctx=ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + H = sc.HermitianSpace(3, ctx=ctx, check_level="none") rng = np.random.default_rng(0) A = H.symmetrize(ctx.asarray(rng.standard_normal((3, 3)))) evals, evecs = H.spectral_decompose(A) @@ -120,8 +179,8 @@ def test_eig_to_dense_reconstructs_via_independent_formula(self): tautology) — this genuinely verifies the reconstruction einsum and eigenvector handling. """ - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - H = sc.HermitianSpace(3, ctx=ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + H = sc.HermitianSpace(3, ctx=ctx, check_level="none") rng = np.random.default_rng(1) A = H.symmetrize(ctx.asarray(rng.standard_normal((3, 3)))) evals, evecs = H.spectral_decompose(A) @@ -146,8 +205,8 @@ def test_psd_proj_yields_nonneg_spectrum(self, numpy_ctx): def test_psd_proj_is_idempotent_on_psd_input(self): """psd_proj is a projector — applying it twice equals once.""" - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - H = sc.HermitianSpace(2, ctx=ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + H = sc.HermitianSpace(2, ctx=ctx, check_level="none") # SPD input M = ctx.asarray([[2.0, 0.5], [0.5, 3.0]]) A = H.symmetrize(M) diff --git a/tests/spaces/test_jordan_algebra_spaces.py b/tests/spaces/test_jordan_algebra_spaces.py index 5a21f4c..83af221 100644 --- a/tests/spaces/test_jordan_algebra_spaces.py +++ b/tests/spaces/test_jordan_algebra_spaces.py @@ -58,8 +58,8 @@ def test_elementwise_spectral_decompose_round_trip(self, numpy_ctx): def test_hermitian_spectral_decompose_round_trip(self): """Reconstruction from spectrum is bit-inexact; run at check_level=none.""" - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - space = sc.HermitianSpace(3, ctx=ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + space = sc.HermitianSpace(3, ctx=ctx, check_level="none") rng = np.random.default_rng(0) M = ctx.asarray(rng.standard_normal((3, 3))) H = space.symmetrize(M) @@ -90,8 +90,8 @@ def test_spectral_apply_exp_on_hermitian(self): symmetric, so the strict Hermitian membership gate would refuse it; run with check_level=none. """ - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - space = sc.HermitianSpace(2, ctx=ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + space = sc.HermitianSpace(2, ctx=ctx, check_level="none") H = space.symmetrize(ctx.asarray([[2.0, 0.5], [0.5, 1.0]])) applied = space.spectral_apply(H, lambda t: ctx.ops.exp(t)) expected = scipy.linalg.expm(to_numpy(H)) diff --git a/tests/spaces/test_jordan_invariants.py b/tests/spaces/test_jordan_invariants.py index 749c723..725e3c1 100644 --- a/tests/spaces/test_jordan_invariants.py +++ b/tests/spaces/test_jordan_invariants.py @@ -19,7 +19,7 @@ def _numpy_ctx(dtype=np.float64): - return sc.Context(sc.NumpyOps(), dtype=dtype, check_level="none") + return sc.Context(sc.NumpyOps(), dtype=dtype) # =========================================================================== @@ -94,7 +94,7 @@ def test_trace_inner_oracle(case): # =========================================================================== def test_hermitian_trace_determinant_preserve_batch_axis(): ctx = _numpy_ctx() - X = sc.HermitianSpace(3, ctx=ctx) + X = sc.HermitianSpace(3, ctx=ctx, check_level="none") rng = np.random.default_rng(0) A = rng.standard_normal((4, 3, 3)) A = 0.5 * (A + np.swapaxes(A, -1, -2)) # symmetric batch @@ -111,7 +111,7 @@ def test_hermitian_trace_determinant_preserve_batch_axis(): def test_elementwise_multidim_reduces_all_element_axes(): ctx = _numpy_ctx() - X = sc.ElementwiseJordanSpace((2, 3), ctx) # matrix-shaped elementwise algebra + X = sc.ElementwiseJordanSpace((2, 3), ctx, check_level="none") # matrix-shaped elementwise algebra rng = np.random.default_rng(1) A = rng.standard_normal((2, 3)) arr = ctx.asarray(A) @@ -134,8 +134,8 @@ def test_elementwise_multidim_reduces_all_element_axes(): # =========================================================================== def test_stacked_is_direct_sum_of_copies(): ctx = _numpy_ctx() - base = sc.HermitianSpace(2, ctx=ctx) - X = sc.StackedSpace(base, 3, ctx) + base = sc.HermitianSpace(2, ctx=ctx, check_level="none") + X = sc.StackedSpace(base, 3, ctx, check_level="none") rng = np.random.default_rng(2) copies = [] for _ in range(3): @@ -163,7 +163,7 @@ def test_stacked_trace_determinant_preserve_batch_axis(): # A batch of stacked elements (B, count, n, n) must reduce only the copy axis, # yielding (B,) — the stacked reduction must not collapse the leading batch axis. ctx = _numpy_ctx() - X = sc.StackedSpace(sc.HermitianSpace(2, ctx=ctx), 3, ctx) + X = sc.StackedSpace(sc.HermitianSpace(2, ctx=ctx, check_level="none"), 3, ctx, check_level="none") rng = np.random.default_rng(4) A = rng.standard_normal((4, 3, 2, 2)) A = 0.5 * (A + np.swapaxes(A, -1, -2)) diff --git a/tests/spaces/test_space_base.py b/tests/spaces/test_space_base.py index e4f69c6..aac09e9 100644 --- a/tests/spaces/test_space_base.py +++ b/tests/spaces/test_space_base.py @@ -22,8 +22,8 @@ class _FiniteSetSpace(sc.Space): """Minimal concrete ``Space`` for the base-class contract tests.""" - def __init__(self, values: set[Any], ctx=None) -> None: - super().__init__(ctx) + def __init__(self, values: set[Any], ctx=None, check_level=None) -> None: + super().__init__(ctx, check_level=check_level) self.values = values def _check_member(self, x: Any) -> None: @@ -61,21 +61,21 @@ def test_field_updates_on_convert(self, numpy_ctx, numpy_complex_ctx): # =========================================================================== class TestCheckMember: def test_none_skips_membership(self): - ctx = sc.Context(sc.NumpyOps(), check_level="none") - space = _FiniteSetSpace({"a", "b"}, ctx) + ctx = sc.Context(sc.NumpyOps()) + space = _FiniteSetSpace({"a", "b"}, ctx, check_level="none") # "c" is not a member but ``none`` skips ``_check_member``. space.check_member("c") def test_membership_runs_at_standard(self): - ctx = sc.Context(sc.NumpyOps(), check_level="standard") - space = _FiniteSetSpace({"a", "b"}, ctx) + ctx = sc.Context(sc.NumpyOps()) + space = _FiniteSetSpace({"a", "b"}, ctx, check_level="standard") space.check_member("a") with pytest.raises(ValueError, match="not a member"): space.check_member("c") def test_membership_runs_at_strict(self): - ctx = sc.Context(sc.NumpyOps(), check_level="strict") - space = _FiniteSetSpace({"a", "b"}, ctx) + ctx = sc.Context(sc.NumpyOps()) + space = _FiniteSetSpace({"a", "b"}, ctx, check_level="strict") with pytest.raises(ValueError, match="not a member"): space.check_member("c") @@ -104,7 +104,7 @@ def test_convert_round_trip_preserves_state(self, numpy_ctx, numpy_f32_ctx): assert roundtrip.values == space.values def test_convert_accepts_family_string(self): - space = _FiniteSetSpace({"a"}, sc.Context(sc.NumpyOps(), check_level="strict")) + space = _FiniteSetSpace({"a"}, sc.Context(sc.NumpyOps()), check_level="strict") out = space.convert("numpy") # 'numpy' resolves through normalize_context; default check_level differs. assert isinstance(out, _FiniteSetSpace) diff --git a/tests/spaces/test_space_checks.py b/tests/spaces/test_space_checks.py index ad177da..c78b94b 100644 --- a/tests/spaces/test_space_checks.py +++ b/tests/spaces/test_space_checks.py @@ -268,7 +268,7 @@ def test_passing_checks_run_in_full(self, numpy_ctx): def test_below_minimum_level_check_is_skipped(self): """A check whose ``minimum_level`` exceeds the active level is skipped.""" - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="cheap") + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) @dataclass(frozen=True) class _StrictOnlyAlwaysFail(sc.SpaceCheck): @@ -286,7 +286,8 @@ class _SpaceWithStrictCheck(sc.DenseCoordinateSpace): def _local_checks(self): return (_StrictOnlyAlwaysFail(),) - space = _SpaceWithStrictCheck((2,), ctx) + space = _SpaceWithStrictCheck((2,), ctx, check_level="cheap") + assert space.check_level == "cheap" # check_level=cheap < strict ⇒ the strict-only check is skipped. _run_checks(space, ctx.asarray([1.0, 2.0]), allow_leading=False) @@ -363,12 +364,13 @@ def _local_checks(self): space.check_member(numpy_ctx.asarray([3.0])) def test_disabled_context_skips_local_checks(self): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) class _AlwaysRejecting(sc.DenseCoordinateSpace): checks = (_RejectFirstEntryCheck("child_reject", 0.0),) - space = _AlwaysRejecting((1,), ctx) + space = _AlwaysRejecting((1,), ctx, check_level="none") + assert space.check_level == "none" # check_level=none silences class-level checks too. space.check_member(ctx.asarray([0.0])) diff --git a/tests/spaces/test_stacked_space.py b/tests/spaces/test_stacked_space.py index 37ce85d..d7e2f19 100644 --- a/tests/spaces/test_stacked_space.py +++ b/tests/spaces/test_stacked_space.py @@ -337,8 +337,9 @@ def test_negative_count_raises(self, numpy_ctx): class TestJaxPytree: def test_pytree_round_trip(self): import jax - ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") - space = sc.DenseCoordinateSpace((2,), ctx).stacked(3) + ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) + with sc.use_check_level("none"): + space = sc.DenseCoordinateSpace((2,), ctx).stacked(3) leaves, treedef = jax.tree_util.tree_flatten(space) rebuilt = jax.tree_util.tree_unflatten(treedef, leaves) assert rebuilt == space diff --git a/tests/spaces/test_tree_space.py b/tests/spaces/test_tree_space.py index b4f4765..7e1f348 100644 --- a/tests/spaces/test_tree_space.py +++ b/tests/spaces/test_tree_space.py @@ -269,8 +269,10 @@ def test_check_reports_leaf_path(self, numpy_ctx): tree.check(invalid) def test_check_skipped_when_check_level_none(self, numpy_ctx): - none_ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="none") - tree = sc.TreeSpace(_nested_template(), _three_spaces(none_ctx), ctx=none_ctx) + none_ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + with sc.use_check_level("none"): + tree = sc.TreeSpace(_nested_template(), _three_spaces(none_ctx), ctx=none_ctx) + assert tree.check_level == "none" leaves = list(tree.element(_nested_value(none_ctx)).leaves) leaves[2] = none_ctx.asarray([1.0, 2.0]) # wrong shape # No raise — checks are disabled. @@ -457,8 +459,9 @@ def test_riesz_componentwise_on_mixed_geometry(self, numpy_ctx): class TestJaxPytree: def test_tree_space_round_trip(self): import jax - ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") - tree = sc.TreeSpace(_nested_template(), _three_spaces(ctx), ctx=ctx) + ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) + with sc.use_check_level("none"): + tree = sc.TreeSpace(_nested_template(), _three_spaces(ctx), ctx=ctx) leaves, treedef = jax.tree_util.tree_flatten(tree) rebuilt = jax.tree_util.tree_unflatten(treedef, leaves) assert leaves == [] @@ -467,8 +470,9 @@ def test_tree_space_round_trip(self): def test_tree_element_round_trip(self): import jax - ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") - tree = sc.TreeSpace(_nested_template(), _three_spaces(ctx), ctx=ctx) + ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) + with sc.use_check_level("none"): + tree = sc.TreeSpace(_nested_template(), _three_spaces(ctx), ctx=ctx) element = tree.element(_nested_value(ctx)) leaves, treedef = jax.tree_util.tree_flatten(element) rebuilt = jax.tree_util.tree_unflatten(treedef, leaves) @@ -478,11 +482,11 @@ def test_tree_element_round_trip(self): def test_named_tuple_structure_preserved_through_jax(self): import jax - ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") + ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) template = _State(ctx.asarray([1.0, 2.0]), ctx.asarray([3.0, 4.0, 5.0])) - leaves = (sc.DenseCoordinateSpace((2,), ctx), - sc.DenseCoordinateSpace((3,), ctx)) - tree = sc.TreeSpace.from_template(template, leaves, ctx=ctx) + leaves = (sc.DenseCoordinateSpace((2,), ctx, check_level="none"), + sc.DenseCoordinateSpace((3,), ctx, check_level="none")) + tree = sc.TreeSpace.from_template(template, leaves, ctx=ctx, check_level="none") flat, treedef = jax.tree_util.tree_flatten(tree) rebuilt = jax.tree_util.tree_unflatten(treedef, flat) assert rebuilt == tree diff --git a/tests/spaces/test_tree_spectral_decomposition.py b/tests/spaces/test_tree_spectral_decomposition.py index e02943e..52d3333 100644 --- a/tests/spaces/test_tree_spectral_decomposition.py +++ b/tests/spaces/test_tree_spectral_decomposition.py @@ -68,7 +68,7 @@ def test_spectral_decompose_returns_decomposition(self, numpy_ctx): class TestJaxPytree: def test_round_trip(self): import jax - ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") + ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) decomp = sc.TreeSpectralDecomposition( eigvals=(ctx.asarray([1.0, 2.0]), ctx.asarray([3.0])), frames=(None, None), @@ -82,9 +82,10 @@ def test_round_trip(self): def test_round_trip_preserves_treedef(self): import jax - ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype(), check_level="none") - leaves = (sc.ElementwiseJordanSpace((2,), ctx), sc.ElementwiseJordanSpace((1,), ctx)) - space = sc.TreeSpace.from_leaf_spaces(leaves, ctx) + ctx = sc.Context(sc.JaxOps(), dtype=jax_real_dtype()) + leaves = (sc.ElementwiseJordanSpace((2,), ctx, check_level="none"), + sc.ElementwiseJordanSpace((1,), ctx, check_level="none")) + space = sc.TreeSpace.from_leaf_spaces(leaves, ctx, check_level="none") x = space.element((ctx.asarray([1.0, 2.0]), ctx.asarray([3.0]))) decomp = space.spectral_decompose(x) flat, treedef = jax.tree_util.tree_flatten(decomp) @@ -181,18 +182,18 @@ def test_nested_hermitian_leaf_round_trip(self, numpy_ctx): class TestFlatSpectralRegression: - """W3 must not disturb SpectralLpNormFunctional on flat (non-tree) Jordan domains.""" + """W3 must not disturb the spectral lift on flat (non-tree) Jordan domains.""" def test_hermitian_nuclear_norm_unchanged(self, numpy_ctx): X = sc.HermitianSpace(2, ctx=numpy_ctx) A = numpy_ctx.asarray([[2.0, 0.0], [0.0, -3.0]]) - f = sc.SpectralLpNormFunctional(X, 1) # nuclear norm |2| + |-3| + f = sc.spectralize(X, lambda s: sc.LpNormFunctional(s, 1)) # nuclear norm |2| + |-3| assert float(f.value(A)) == 5.0 def test_elementwise_frobenius_unchanged(self, numpy_ctx): X = sc.ElementwiseJordanSpace((3,), numpy_ctx) x = numpy_ctx.asarray([3.0, -4.0, 0.0]) - f = sc.SpectralLpNormFunctional(X, 2) # sqrt(9 + 16) = 5 + f = sc.spectralize(X, lambda s: sc.LpNormFunctional(s, 2)) # sqrt(9 + 16) = 5 assert float(f.value(x)) == pytest.approx(5.0) @@ -214,9 +215,9 @@ def test_single_leaf_tree_flat_vs_structured(self, numpy_ctx): def test_structured_spectrum_is_not_a_space_member(self): # A Hermitian leaf's spectrum (n,) differs from its element shape (n, n), so # the structured spectrum matches the treedef but is NOT a valid member. - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="standard") + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) leaves = (sc.HermitianSpace(2, ctx=ctx), sc.ElementwiseJordanSpace((2,), ctx)) - space = sc.TreeSpace.from_template((0, (0,)), leaves, ctx=ctx) + space = sc.TreeSpace.from_template((0, (0,)), leaves, ctx=ctx, check_level="standard") x = space.unflatten_tree( (ctx.asarray([[2.0, 0.0], [0.0, 3.0]]), ctx.asarray([1.0, -1.0])) ) @@ -227,9 +228,9 @@ def test_structured_spectrum_is_not_a_space_member(self): def test_complex_hermitian_nested_round_trip(self): # Complex eigenvalues/frames survive the structure-preserving round-trip. - ctx = sc.Context(sc.NumpyOps(), dtype=np.complex128, check_level="standard") + ctx = sc.Context(sc.NumpyOps(), dtype=np.complex128) leaves = (sc.HermitianSpace(2, ctx=ctx), sc.ElementwiseJordanSpace((2,), ctx)) - space = sc.TreeSpace.from_template((0, (0,)), leaves, ctx=ctx) + space = sc.TreeSpace.from_template((0, (0,)), leaves, ctx=ctx, check_level="standard") herm = ctx.asarray([[2.0, 0.5 + 0.5j], [0.5 - 0.5j, 1.0]]) x = space.unflatten_tree((herm, ctx.asarray([1.0 + 0j, -2.0 + 0j]))) decomp = space.spectral_decompose(x) diff --git a/tests/test_equality.py b/tests/test_equality.py new file mode 100644 index 0000000..7b3a41f --- /dev/null +++ b/tests/test_equality.py @@ -0,0 +1,310 @@ +"""Tests for the tiered ``__eq__`` rule: backend compatibility first, then algebra. + +Contract under test (see docs/dev/0.4.0-eq-implementation-spec.md): + +* Tier 1 — backend compatibility: same concrete type AND same backend + (ops family + dtype), ignoring ``check_level``. Failure -> ``NotImplemented`` + (so foreign-type comparisons stay symmetric and fall back to identity). +* Tier 2 — cheap structural checks (shapes, spaces, counts, treedef) before any + numerical comparison. +* Tier 3 — ``allclose`` on values, with ``equal_nan=True`` for reflexivity. +""" +from __future__ import annotations + +import numpy as np +import pytest + +import spacecore as sc + + +@pytest.fixture +def ctx(): + return sc.Context(sc.NumpyOps(), dtype=np.float64) + + +@pytest.fixture +def ctx32(): + return sc.Context(sc.NumpyOps(), dtype=np.float32) + + +@pytest.fixture +def cctx(): + return sc.Context(sc.NumpyOps(), dtype=np.complex128) + + +def _maybe_ctx(family, dtype): + from spacecore.contextual._state import normalize_context + + try: + return normalize_context(family, dtype=dtype) + except Exception: + return None + + +# =========================================================================== +# Tier 1 — backend compatibility (shared gate) +# =========================================================================== +class TestBackendGate: + def test_check_level_ignored_for_spaces(self, ctx): + # check_level is policy, not backend: must not affect identity. + assert sc.DenseVectorSpace((3,), ctx) == sc.DenseVectorSpace( + (3,), ctx, check_level="strict" + ) + + def test_check_level_ignored_for_linops(self, ctx): + v = sc.DenseVectorSpace((3,), ctx) + vs = sc.DenseVectorSpace((3,), ctx, check_level="strict") + A = sc.DenseLinOp(ctx.asarray(np.eye(3)), v, v, ctx) + As = sc.DenseLinOp(ctx.asarray(np.eye(3)), vs, vs, ctx, check_level="strict") + assert A == As + + def test_dtype_is_part_of_identity(self, ctx, ctx32): + # float32 vs float64: different representation -> not equal. + assert sc.DenseVectorSpace((3,), ctx) != sc.DenseVectorSpace((3,), ctx32) + + def test_cross_backend_not_equal(self, ctx): + other = _maybe_ctx("jax", np.float64) or _maybe_ctx("torch", np.float64) + if other is None: + pytest.skip("no second backend available") + assert sc.DenseVectorSpace((3,), ctx) != sc.DenseVectorSpace((3,), other) + + def test_foreign_type_symmetric(self, ctx): + v = sc.DenseVectorSpace((3,), ctx) + A = sc.DenseLinOp(ctx.asarray(np.eye(3)), v, v, ctx) + assert (v == 5) is False and (5 == v) is False + assert (A == 5) is False and (5 == A) is False + + def test_gate_returns_notimplemented_on_foreign_type(self, ctx): + v = sc.DenseVectorSpace((3,), ctx) + assert v.__eq__(5) is NotImplemented + A = sc.DenseLinOp(ctx.asarray(np.eye(3)), v, v, ctx) + assert A.__eq__(object()) is NotImplemented + + +# =========================================================================== +# Spaces — field, shape, geometry, decisions +# =========================================================================== +class TestSpaceEquality: + def test_field_distinguishes(self, ctx, cctx): + assert sc.DenseVectorSpace((3,), ctx) != sc.DenseVectorSpace((3,), cctx) + + def test_shape(self, ctx): + assert sc.DenseVectorSpace((3,), ctx) == sc.DenseVectorSpace((3,), ctx) + assert sc.DenseVectorSpace((3,), ctx) != sc.DenseVectorSpace((4,), ctx) + + def test_geometry_kind_and_weights(self, ctx): + w = sc.WeightedInnerProduct(ctx.asarray(np.arange(1, 4, dtype=float))) + w2 = sc.WeightedInnerProduct(ctx.asarray(np.arange(1, 4, dtype=float))) + wd = sc.WeightedInnerProduct(ctx.asarray(np.arange(2, 5, dtype=float))) + assert sc.DenseVectorSpace((3,), ctx, geometry=w) == sc.DenseVectorSpace((3,), ctx, geometry=w2) + assert sc.DenseVectorSpace((3,), ctx, geometry=w) != sc.DenseVectorSpace((3,), ctx, geometry=wd) + assert sc.DenseVectorSpace((3,), ctx, geometry=w) != sc.DenseVectorSpace((3,), ctx) # euclidean + + def test_weighted_geometry_shape_guard(self, ctx): + # Regression: np.allclose used to broadcast [2.] vs [2.,2.,2.] to True. + w1 = sc.WeightedInnerProduct(ctx.asarray(np.array([2.0]))) + w3 = sc.WeightedInnerProduct(ctx.asarray(np.array([2.0, 2.0, 2.0]))) + assert w1 != w3 + + def test_hermitian_tolerances_excluded(self, ctx): + # Decision: atol/rtol/enforce_herm are membership policy, not identity. + assert sc.HermitianSpace(3, atol=0.0, ctx=ctx) == sc.HermitianSpace(3, atol=1e-6, ctx=ctx) + assert sc.HermitianSpace(3, enforce_herm=True, ctx=ctx) == sc.HermitianSpace( + 3, enforce_herm=False, ctx=ctx + ) + assert sc.HermitianSpace(3, ctx=ctx) != sc.HermitianSpace(4, ctx=ctx) + + def test_stacked_base_is_load_bearing(self, ctx): + a = sc.DenseVectorSpace((3,), ctx).stacked(4) + b = sc.DenseVectorSpace((3,), ctx).stacked(4) + c = sc.DenseVectorSpace((5,), ctx).stacked(4) + assert a == b and a != c + + def test_tree_treedef_and_leaves(self, ctx): + leaves = (sc.DenseVectorSpace((3,), ctx), sc.DenseVectorSpace((2,), ctx)) + t1 = sc.TreeSpace((0, 0), leaves, ctx=ctx) + t2 = sc.TreeSpace((0, 0), leaves, ctx=ctx) + t3 = sc.TreeSpace((0, 0), (sc.DenseVectorSpace((3,), ctx), sc.DenseVectorSpace((9,), ctx)), ctx=ctx) + assert t1 == t2 and t1 != t3 + + def test_spaces_unhashable(self, ctx): + with pytest.raises(TypeError): + hash(sc.DenseVectorSpace((3,), ctx)) + + +# =========================================================================== +# Linear operators +# =========================================================================== +class TestLinOpEquality: + def _v(self, ctx, n=3): + return sc.DenseVectorSpace((n,), ctx) + + def test_dense_values(self, ctx): + v = self._v(ctx) + A = sc.DenseLinOp(ctx.asarray(np.eye(3)), v, v, ctx) + assert A == sc.DenseLinOp(ctx.asarray(np.eye(3)), v, v, ctx) + assert A != sc.DenseLinOp(ctx.asarray(2 * np.eye(3)), v, v, ctx) + + def test_dense_nan_reflexive(self, ctx): + v = self._v(ctx) + m = np.eye(3) + m[0, 0] = np.nan + A = sc.DenseLinOp(ctx.asarray(m), v, v, ctx) + assert A == A # equal_nan=True + + def test_dense_geometry_in_domain_matters(self, ctx): + # Same matrix, domains differ only in geometry -> different operator. + v = self._v(ctx) + w = sc.WeightedInnerProduct(ctx.asarray(np.arange(1, 4, dtype=float))) + vw = sc.DenseVectorSpace((3,), ctx, geometry=w) + assert sc.DenseLinOp(ctx.asarray(np.eye(3)), v, v, ctx) != sc.DenseLinOp( + ctx.asarray(np.eye(3)), vw, vw, ctx + ) + + def test_sparse_nan_not_equal_to_finite(self, ctx): + # Regression: allclose_sparse was NaN-blind, so a NaN entry compared + # "close" to a finite one and two different operators returned True. + import scipy.sparse as sps + + v = self._v(ctx) + + def mk(corner): + m = sps.eye(3, format="csr").tolil() + m[0, 0] = corner + return sc.SparseLinOp(ctx.assparse(m.tocsr()), v, v, ctx) + + assert (mk(np.nan) == mk(5.0)) is False + assert (mk(5.0) == mk(np.nan)) is False + assert (mk(np.nan) == mk(np.nan)) is False # sparse NaN is non-reflexive (documented) + assert (mk(2.0) == mk(2.0)) is True + + def test_cross_type_not_equal(self, ctx): + # Mathematically identity, structurally different types -> not equal. + v = self._v(ctx) + ident = sc.IdentityLinOp(v, ctx) + diag = sc.DiagonalLinOp(ctx.asarray(np.ones(3)), v, ctx) + dense = sc.DenseLinOp(ctx.asarray(np.eye(3)), v, v, ctx) + assert ident != diag and ident != dense and diag != dense + + def test_scaled_returns_python_bool(self, ctx): + # Regression: 0-d array scalar comparison used to leak np.bool_. + v = self._v(ctx) + A = sc.DenseLinOp(ctx.asarray(np.eye(3)), v, v, ctx) + for scal in (2.0, np.float64(2.0)): + res = (scal * A) == (2.0 * A) + assert type(res) is bool and res is True + assert ((2.0 * A) == (3.0 * A)) is False + + def test_scaled_leak_through_composition(self, ctx): + v = self._v(ctx) + A = sc.DenseLinOp(ctx.asarray(np.eye(3)), v, v, ctx) + res = ((2.0 * A) @ A) == ((3.0 * A) @ A) + assert type(res) is bool and res is False + + def test_scaled_nan_scalar_reflexive(self, ctx): + # Regression: bool(nan == nan) is False, so a NaN-scaled op must use the + # NaN-reflexive scalar comparison to stay equal to itself. + v = self._v(ctx) + A = sc.DenseLinOp(ctx.asarray(np.eye(3)), v, v, ctx) + for scal in (float("nan"), np.float64("nan")): + s = scal * A + assert (s == s) is True + assert ((s @ A) == (s @ A)) is True # through composition + assert (s.H == s.H) is True # through adjoint + # distinct scalars (one NaN) still unequal + assert ((float("nan") * A) == (2.0 * A)) is False + + def test_sum_is_ordered(self, ctx): + v = self._v(ctx) + A = sc.DenseLinOp(ctx.asarray(np.eye(3)), v, v, ctx) + ident = sc.IdentityLinOp(v, ctx) + # Distinct operands in different order must not be equal. + assert (ident + A) != (A + ident) + + def test_matrixfree_callable_identity(self, ctx): + v = self._v(ctx) + def f(x): + return x + + a = sc.MatrixFreeLinOp(f, f, v, v, ctx) + assert a == sc.MatrixFreeLinOp(f, f, v, v, ctx) + assert a != sc.MatrixFreeLinOp(lambda x: x, lambda x: x, v, v, ctx) + + def test_adjoint(self, ctx): + v = self._v(ctx) + A = sc.DenseLinOp(ctx.asarray(np.arange(9.0).reshape(3, 3)), v, v, ctx) + B = sc.DenseLinOp(ctx.asarray(np.arange(9.0).reshape(3, 3)), v, v, ctx) + assert A.H == B.H + + def test_tree_linop_parts_and_structure(self, ctx): + v2 = sc.DenseVectorSpace((2,), ctx) + v3 = sc.DenseVectorSpace((3,), ctx) + def blk(): + return sc.BlockDiagonalLinOp( + [ + sc.DenseLinOp(ctx.asarray(np.eye(3)), v3, v3, ctx), + sc.DenseLinOp(ctx.asarray(np.eye(2)), v2, v2, ctx), + ] + ) + assert blk() == blk() + other = sc.BlockDiagonalLinOp( + [ + sc.DenseLinOp(ctx.asarray(2 * np.eye(3)), v3, v3, ctx), + sc.DenseLinOp(ctx.asarray(np.eye(2)), v2, v2, ctx), + ] + ) + assert blk() != other + + +# =========================================================================== +# Functionals +# =========================================================================== +class TestFunctionalEquality: + def _v(self, ctx, n=3): + return sc.DenseVectorSpace((n,), ctx) + + def test_base_returns_notimplemented(self, ctx): + f = sc.InnerProductFunctional(ctx.asarray(np.ones(3)), self._v(ctx), ctx) + assert sc.Functional.__eq__(f, object()) is NotImplemented + + def test_inner_product(self, ctx): + v = self._v(ctx) + a = sc.InnerProductFunctional(ctx.asarray(np.ones(3)), v, ctx) + assert a == sc.InnerProductFunctional(ctx.asarray(np.ones(3)), v, ctx) + assert a != sc.InnerProductFunctional(ctx.asarray(np.arange(3.0)), v, ctx) + + def test_quadratic_linear_none_safe(self, ctx): + v = self._v(ctx) + Q = sc.DenseLinOp(ctx.asarray(np.eye(3)), v, v, ctx) + lin = sc.InnerProductFunctional(ctx.asarray(np.ones(3)), v, ctx) + q_none = sc.LinOpQuadraticForm(Q, ctx=ctx) + q_lin = sc.LinOpQuadraticForm(Q, linear=lin, ctx=ctx) + assert q_none == sc.LinOpQuadraticForm(Q, ctx=ctx) + assert q_none != q_lin # None vs functional, no crash + assert q_lin == sc.LinOpQuadraticForm(Q, linear=lin, ctx=ctx) + + def test_composed(self, ctx): + v, w = self._v(ctx), sc.DenseVectorSpace((2,), ctx) + A = sc.DenseLinOp(ctx.asarray(np.ones((2, 3))), v, w, ctx) + def fn(x): + return ctx.ops.vdot(ctx.asarray(np.ones(2)), x) + + F = sc.MatrixFreeLinearFunctional(fn, w, ctx) + comp = F.compose(A) + assert type(comp).__name__ == "ComposedFunctional" + assert comp == F.compose(A) + + +# =========================================================================== +# Cross-backend gate (parametrized) +# =========================================================================== +@pytest.mark.parametrize("family", ["jax", "torch"]) +def test_same_backend_equal_cross_backend_not(family): + np_ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + other = _maybe_ctx(family, np.float64) + if other is None: + pytest.skip(f"{family} unavailable") + a = sc.DenseVectorSpace((3,), other) + b = sc.DenseVectorSpace((3,), other) + assert a == b # same backend + assert a != sc.DenseVectorSpace((3,), np_ctx) # cross-backend diff --git a/tests/test_opfamily.py b/tests/test_opfamily.py new file mode 100644 index 0000000..e0cae1d --- /dev/null +++ b/tests/test_opfamily.py @@ -0,0 +1,289 @@ +"""Point-indexed operator families and the functional-weighted map ``F · A``. + +``m(x) = F(x) A x`` is not linear, so it is *not* a ``LinOp``. These tests pin +the three things that follow from that: + +* the map is what it claims and is genuinely non-linear; +* freezing the point (``at``) recovers an ordinary ``LinOp`` that participates + in the operator algebra; +* the derivative (``linearize_at``) is a *different* operator from the frozen + member, and its adjoint is the **metric** adjoint. + +The metric and complex cases are the load-bearing ones: a Euclidean-only suite +would pass with a coordinate adjoint and wrong conjugation. +""" +from __future__ import annotations + +import numpy as np +import pytest + +import spacecore as sc + +from tests._helpers import to_numpy + + +# ``spacecore.opfamily`` is a top-level module spanning linop and functional, so +# this suite lives at the tests root and declares the fixture the per-object +# suites (tests/functional, tests/linops) provide via their own conftest. +@pytest.fixture +def numpy_ctx(): + """Float64 NumPy context (default ``standard`` check level).""" + return sc.Context(sc.NumpyOps(), dtype=np.float64) + + +def _matrix(): + return np.array([[1.0, 2.0, 0.0], [0.0, 1.0, 3.0], [2.0, 0.0, 1.0]]) + + +def _euclidean(ctx): + X = sc.DenseCoordinateSpace((3,), ctx) + A = sc.DenseLinOp(ctx.asarray(_matrix()), X, X, ctx) + return X, A, sc.SquaredL2NormFunctional(X) + + +def _weighted(ctx): + weights = ctx.asarray(np.array([2.0, 5.0, 11.0])) + X = sc.DenseCoordinateSpace((3,), ctx, geometry=sc.WeightedInnerProduct(weights)) + A = sc.DenseLinOp(ctx.asarray(_matrix()), X, X, ctx) + return X, A, sc.SquaredL2NormFunctional(X) + + +def _complex_weighted(): + ctx = sc.Context(sc.NumpyOps(), dtype=np.complex128) + weights = ctx.asarray(np.array([2.0, 5.0, 11.0])) + X = sc.DenseCoordinateSpace((3,), ctx, geometry=sc.WeightedInnerProduct(weights)) + A = sc.DenseLinOp( + ctx.asarray(np.array([[1 + 1j, 2, 0], [0, 1, 3 - 2j], [2, 0, 1j]])), X, X, ctx + ) + # A complex-VALUED weight, so the conjugations in the adjoint are exercised. + F = sc.InnerProductFunctional(ctx.asarray([1 + 1j, 2 - 0.5j, -1 + 0.3j]), X, ctx) + return ctx, X, A, F + + +# =========================================================================== +# The map +# =========================================================================== +class TestFunctionalWeightedMap: + def test_value_is_the_weighted_image(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + np.testing.assert_allclose( + to_numpy((F * A).apply(x)), float(F.value(x)) * to_numpy(A.apply(x)) + ) + + def test_is_not_a_linop(self, numpy_ctx): + """The whole point: it must not enter code that assumes linearity.""" + X, A, F = _euclidean(numpy_ctx) + m = F * A + assert isinstance(m, sc.OperatorFamily) + assert not isinstance(m, sc.LinOp) + + def test_is_genuinely_nonlinear(self, numpy_ctx): + """``m(2x) != 2 m(x)`` — this is why it cannot be a ``LinOp``.""" + X, A, F = _euclidean(numpy_ctx) + m = F * A + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + assert not np.allclose(to_numpy(m.apply(X.scale(2.0, x))), 2.0 * to_numpy(m.apply(x))) + + def test_both_operand_orders_work(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + np.testing.assert_allclose(to_numpy((F * A).apply(x)), to_numpy((A * F).apply(x))) + + def test_call_is_apply(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + np.testing.assert_allclose(to_numpy((F * A)(x)), to_numpy((F * A).apply(x))) + + +# =========================================================================== +# Freezing the point recovers a LinOp +# =========================================================================== +class TestFrozenMember: + def test_at_returns_a_scaled_linop(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + member = (F * A).at(x) + assert isinstance(member, sc.LinOp) + h = numpy_ctx.asarray([0.5, -1.0, 2.0]) + np.testing.assert_allclose( + to_numpy(member.apply(h)), float(F.value(x)) * to_numpy(A.apply(h)) + ) + + def test_at_x_applied_to_x_is_the_map(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + m = F * A + np.testing.assert_allclose(to_numpy(m.at(x).apply(x)), to_numpy(m.apply(x))) + + def test_member_participates_in_the_operator_algebra(self, numpy_ctx): + """Freezing buys the whole algebra back: ``@``, ``+``, ``.H``.""" + X, A, F = _euclidean(numpy_ctx) + member = (F * A).at(numpy_ctx.asarray([1.0, 2.0, -1.0])) + assert isinstance(member @ A, sc.LinOp) + assert isinstance(member + A, sc.LinOp) + h = numpy_ctx.asarray([0.5, -1.0, 2.0]) + np.testing.assert_allclose( + to_numpy(member.H.apply(h)), + float(F.value(numpy_ctx.asarray([1.0, 2.0, -1.0]))) * to_numpy(A.H.apply(h)), + ) + + +# =========================================================================== +# The derivative is a different operator from the frozen member +# =========================================================================== +class TestLinearization: + @pytest.mark.parametrize("build", [_euclidean, _weighted], ids=["euclidean", "weighted"]) + def test_matches_finite_differences(self, numpy_ctx, build): + X, A, F = build(numpy_ctx) + m = F * A + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + h = numpy_ctx.asarray([0.5, -1.0, 2.0]) + eps = 1e-6 + finite_difference = ( + to_numpy(m.apply(X.axpy(eps, h, x))) - to_numpy(m.apply(X.axpy(-eps, h, x))) + ) / (2.0 * eps) + np.testing.assert_allclose( + to_numpy(m.linearize_at(x).apply(h)), finite_difference, rtol=1e-6, atol=1e-6 + ) + + @pytest.mark.parametrize("build", [_euclidean, _weighted], ids=["euclidean", "weighted"]) + def test_adjoint_identity_holds_in_the_space_geometry(self, numpy_ctx, build): + """``_Y == _X`` — the *metric* adjoint (ADR-009).""" + X, A, F = build(numpy_ctx) + D = (F * A).linearize_at(numpy_ctx.asarray([1.0, 2.0, -1.0])) + h = numpy_ctx.asarray([0.5, -1.0, 2.0]) + w = numpy_ctx.asarray([1.0, 1.0, 1.0]) + np.testing.assert_allclose( + to_numpy(X.inner(D.apply(h), w)), to_numpy(X.inner(h, D.rapply(w))) + ) + + def test_complex_weight_conjugations_are_right(self): + """A complex-valued weight on a weighted metric: both conjugations bite.""" + ctx, X, A, F = _complex_weighted() + m = F * A + x = ctx.asarray([1 + 0j, 2 - 1j, -1 + 2j]) + h = ctx.asarray([0.5 + 1j, -1 + 0j, 2 + 0j]) + w = ctx.asarray([1 + 1j, 1 + 0j, 1 - 1j]) + D = m.linearize_at(x) + + eps = 1e-6 + finite_difference = ( + to_numpy(m.apply(X.axpy(eps, h, x))) - to_numpy(m.apply(X.axpy(-eps, h, x))) + ) / (2.0 * eps) + np.testing.assert_allclose( + to_numpy(D.apply(h)), finite_difference, rtol=1e-6, atol=1e-6 + ) + np.testing.assert_allclose( + to_numpy(X.inner(D.apply(h), w)), to_numpy(X.inner(h, D.rapply(w))) + ) + + def test_derivative_differs_from_the_frozen_member(self, numpy_ctx): + """They are two different operators; conflating them is the easy mistake.""" + X, A, F = _euclidean(numpy_ctx) + m = F * A + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + h = numpy_ctx.asarray([0.5, -1.0, 2.0]) + assert not np.allclose( + to_numpy(m.linearize_at(x).apply(h)), to_numpy(m.at(x).apply(h)) + ) + + def test_derivative_equals_frozen_member_for_a_constant_weight(self, numpy_ctx): + """A constant family has zero rank-one correction, so they coincide.""" + X, A, _ = _euclidean(numpy_ctx) + C = sc.ConstantFunctional(X, 3.0, numpy_ctx) + m = sc.FunctionalScaledOperator(C, A) # built directly: the factory would collapse it + x = numpy_ctx.asarray([1.0, 2.0, -1.0]) + h = numpy_ctx.asarray([0.5, -1.0, 2.0]) + np.testing.assert_allclose( + to_numpy(m.linearize_at(x).apply(h)), to_numpy(m.at(x).apply(h)) + ) + + def test_base_class_has_no_linearization(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + + class Bare(sc.OperatorFamily): + """A family that only implements ``at`` — the minimal subclass.""" + + def at(self, x): + return A + + def tree_flatten(self): + return (), (self.domain, self.codomain, self.ctx) + + @classmethod + def tree_unflatten(cls, aux, children): + return cls(*aux) + + with pytest.raises(NotImplementedError, match="linearize_at"): + Bare(X, X, numpy_ctx).linearize_at(numpy_ctx.asarray([1.0, 2.0, -1.0])) + + +# =========================================================================== +# Factory: the linear case is not forced through the non-linear type +# =========================================================================== +class TestFactory: + def test_constant_weight_collapses_to_a_scaled_linop(self, numpy_ctx): + X, A, _ = _euclidean(numpy_ctx) + result = sc.ConstantFunctional(X, 3.0, numpy_ctx) * A + assert isinstance(result, sc.LinOp) + h = numpy_ctx.asarray([0.5, -1.0, 2.0]) + np.testing.assert_allclose(to_numpy(result.apply(h)), 3.0 * to_numpy(A.apply(h))) + + def test_zero_weight_collapses_to_zero(self, numpy_ctx): + X, A, _ = _euclidean(numpy_ctx) + assert isinstance(sc.ZeroFunctional(X, numpy_ctx) * A, sc.ZeroLinOp) + + def test_collapse_is_structural_not_value_based(self, numpy_ctx): + """A functional that merely *happens* to be constant is not recognized.""" + X, A, _ = _euclidean(numpy_ctx) + constant_valued = sc.InnerProductFunctional( + numpy_ctx.asarray([0.0, 0.0, 0.0]), X, numpy_ctx + ) + assert isinstance(constant_valued * A, sc.FunctionalScaledOperator) + + def test_domain_mismatch_raises(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + other = sc.SquaredL2NormFunctional(sc.DenseCoordinateSpace((2,), numpy_ctx)) + with pytest.raises(ValueError, match="same domain|domain =="): + other * A + + def test_type_guards(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + with pytest.raises(TypeError, match="Functional"): + sc.FunctionalScaledOperator("nope", A) + with pytest.raises(TypeError, match="LinOp"): + sc.FunctionalScaledOperator(F, "nope") + + +# =========================================================================== +# Container protocol +# =========================================================================== +class TestContainerProtocol: + def test_equality(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + assert sc.FunctionalScaledOperator(F, A) == sc.FunctionalScaledOperator(F, A) + other = sc.DenseLinOp(numpy_ctx.asarray(np.eye(3)), X, X, numpy_ctx) + assert sc.FunctionalScaledOperator(F, A) != sc.FunctionalScaledOperator(F, other) + + def test_pytree_round_trip(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + m = sc.FunctionalScaledOperator(F, A) + children, aux = m.tree_flatten() + assert sc.FunctionalScaledOperator.tree_unflatten(aux, children) == m + + def test_convert_moves_both_operands(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + target = sc.Context(sc.NumpyOps(), dtype=np.float32) + moved = sc.FunctionalScaledOperator(F, A).convert(target) + assert moved.ctx == target + assert moved.domain.ctx == target and moved.op.ctx == target + + def test_domain_and_codomain(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + m = F * A + assert m.domain == A.domain and m.codomain == A.codomain + + def test_repr_marks_it_as_nonlinear(self, numpy_ctx): + X, A, F = _euclidean(numpy_ctx) + assert "⇝" in repr(F * A) # not "→", which denotes a linear arrow diff --git a/tests/test_repr.py b/tests/test_repr.py index b0564f4..6df8e25 100644 --- a/tests/test_repr.py +++ b/tests/test_repr.py @@ -404,7 +404,7 @@ def test_no_check_level_leak(self, ctx): # Cross-backend consistency (Goal #5) # =========================================================================== def _maybe_ctx(family, dtype): - from spacecore._contextual._state import normalize_context + from spacecore.contextual._state import normalize_context try: return normalize_context(family, dtype=dtype) diff --git a/tests/test_weighted_tikhonov.py b/tests/test_weighted_tikhonov.py index 3c8822b..ac758b5 100644 --- a/tests/test_weighted_tikhonov.py +++ b/tests/test_weighted_tikhonov.py @@ -23,14 +23,15 @@ def _problem(n=16, m=24, lam=1e-2, seed=3): def _spaces(M, x_weights, y_weights): - ctx = sc.Context(sc.NumpyOps(), dtype=np.float64, check_level="standard") - X = sc.DenseVectorSpace( - (M.shape[1],), ctx, geometry=sc.WeightedInnerProduct(ctx.asarray(x_weights)) - ) - Y = sc.DenseVectorSpace( - (M.shape[0],), ctx, geometry=sc.WeightedInnerProduct(ctx.asarray(y_weights)) - ) - A = sc.DenseLinOp(ctx.asarray(M), X, Y, ctx) + ctx = sc.Context(sc.NumpyOps(), dtype=np.float64) + with sc.use_check_level("standard"): + X = sc.DenseVectorSpace( + (M.shape[1],), ctx, geometry=sc.WeightedInnerProduct(ctx.asarray(x_weights)) + ) + Y = sc.DenseVectorSpace( + (M.shape[0],), ctx, geometry=sc.WeightedInnerProduct(ctx.asarray(y_weights)) + ) + A = sc.DenseLinOp(ctx.asarray(M), X, Y, ctx) return ctx, X, Y, A