From dc1debdf58000a6b6c487e32268b08278fb6a813 Mon Sep 17 00:00:00 2001 From: Pavlo Pelikh Date: Sat, 25 Jul 2026 00:57:11 -0300 Subject: [PATCH 1/6] Refactor contextual submodule --- bench/_dashboard.py | 9 +- bench/_operations.py | 61 +-- bench/dashboard.py | 5 +- bench/macro/cg_poisson.py | 56 +-- bench/macro/density_pipeline.py | 55 +-- bench/macro/jax_full_loop.py | 71 ++- bench/macro/operator_stress.py | 39 +- bench/macro/pdhg.py | 76 ++- bench/macro/power_lanczos.py | 22 +- bench/macro/qot_barycenter.py | 30 +- spacecore/__init__.py | 43 +- spacecore/_contextual/_policies.py | 17 - spacecore/_contextual/_state.py | 446 ------------------ spacecore/_errors.py | 29 ++ spacecore/_lazy_algebra.py | 128 +++++ spacecore/_version.py | 2 +- spacecore/backend/__init__.py | 53 ++- spacecore/backend/_container.py | 215 +++++++++ spacecore/backend/_ops.py | 23 + spacecore/backend/_optional.py | 215 +++++++++ spacecore/backend/_registry.py | 193 ++++++++ spacecore/backend/jax/__init__.py | 11 +- spacecore/backend/jax/_ops.py | 15 + spacecore/backend/jax/_pytree.py | 32 -- spacecore/backend/torch/_ops.py | 37 ++ .../{_contextual => contextual}/__init__.py | 12 +- .../{_contextual => contextual}/_bound.py | 109 +++-- spacecore/{backend => contextual}/_context.py | 138 ++++-- spacecore/contextual/_contextual.py | 258 ++++++++++ spacecore/contextual/_state.py | 220 +++++++++ spacecore/functional/_algebra.py | 155 +++--- spacecore/functional/_base.py | 42 +- spacecore/functional/_composed.py | 21 +- spacecore/functional/_linear.py | 15 +- spacecore/functional/_quadratic.py | 11 +- spacecore/functional/tools/_entropy.py | 17 +- spacecore/functional/tools/_huber.py | 14 +- spacecore/functional/tools/_norms.py | 30 +- spacecore/functional/tools/_spectral.py | 20 +- spacecore/kernels/core/algebra.py | 4 +- spacecore/kernels/specs/_dispatch.py | 15 +- spacecore/linalg/_power.py | 2 +- spacecore/linop/_algebra.py | 241 +++++----- spacecore/linop/_base.py | 40 +- spacecore/linop/_dense.py | 16 +- spacecore/linop/_diagonal.py | 16 +- spacecore/linop/_sparse.py | 16 +- spacecore/linop/tree/_base.py | 19 +- spacecore/linop/tree/_block.py | 43 +- spacecore/linop/tree/_from_single.py | 15 +- spacecore/linop/tree/_to_single.py | 15 +- spacecore/optimize/_optax.py | 4 +- spacecore/space/base/_coordinate.py | 17 +- spacecore/space/base/_space.py | 19 +- spacecore/space/concrete/_dense_coordinate.py | 11 +- spacecore/space/concrete/_dense_vector.py | 38 +- spacecore/space/concrete/_hermitian.py | 11 +- spacecore/space/concrete/_stacked.py | 60 ++- spacecore/space/concrete/_tree_space.py | 39 +- tests/backend/test_context.py | 81 ++-- tests/backend/test_jax_pytree_class.py | 165 ------- tests/backend/test_optional_guard.py | 188 ++++++++ tests/backend/test_pytree_registry.py | 429 +++++++++++++++++ tests/bench/test_bench_smoke.py | 16 +- tests/conftest.py | 18 + tests/context/_contracts.py | 101 ++++ tests/context/conftest.py | 5 + tests/context/test_ambient_scoping.py | 216 +++++++++ tests/context/test_check_policy.py | 151 +++--- tests/context/test_check_policy_helpers.py | 8 + tests/context/test_checked_method.py | 14 +- tests/context/test_compatibility.py | 115 ++--- tests/context/test_context_bound.py | 76 ++- tests/context/test_context_contracts.py | 133 ++++++ .../context/test_context_resolution_policy.py | 102 ++++ tests/context/test_enable_checks.py | 108 +++-- tests/context/test_policies_errors.py | 36 +- tests/context/test_state_free_functions.py | 72 +-- tests/functional/test_algebra.py | 6 +- .../functional/test_generated_functionals.py | 53 ++- tests/functional/test_linop_quadratic_form.py | 11 +- .../test_matrix_free_linear_functional.py | 14 +- tests/functional/test_metric_gradient.py | 142 +++--- tests/generators/_arrays.py | 2 +- tests/generators/_contexts.py | 17 +- tests/generators/_hermitian.py | 2 +- tests/generators/_metrics.py | 2 +- tests/generators/_trees.py | 2 +- tests/generators/functionals.py | 39 +- tests/generators/linops.py | 36 +- tests/generators/spaces.py | 12 +- tests/integration/test_imports.py | 2 +- tests/integration/test_public_api.py | 12 +- tests/integration/test_smoke_torch.py | 45 +- tests/kernels/test_kernel_dispatch.py | 50 +- tests/linalg/_helpers.py | 12 +- tests/linalg/test_core_resolution.py | 3 +- tests/linalg/test_generated_solver_matrix.py | 3 +- tests/linalg/test_solver_contracts.py | 3 +- tests/linalg/test_utils.py | 22 +- tests/linops/test_algebra_factories.py | 67 ++- tests/linops/test_algebra_linops.py | 56 ++- tests/linops/test_block_diagonal_linop.py | 15 +- tests/linops/test_block_matrix_linop.py | 3 +- tests/linops/test_dense_linop.py | 11 +- tests/linops/test_fused_algebra_overhead.py | 13 +- tests/linops/test_linop_jit.py | 9 +- tests/linops/test_metric_helpers.py | 6 +- tests/linops/test_sparse_linop.py | 16 +- tests/linops/test_stacked_linop.py | 14 +- tests/linops/test_sum_to_single_linop.py | 10 +- tests/linops/test_tree_helpers.py | 10 +- tests/optim/test_cached_member_checks.py | 28 +- tests/spaces/test_check_batched.py | 12 +- tests/spaces/test_dense_coordinate_space.py | 2 +- .../spaces/test_elementwise_jordan_spaces.py | 6 +- .../spaces/test_generated_dense_coordinate.py | 4 +- tests/spaces/test_hermitian_space.py | 12 +- tests/spaces/test_jordan_algebra_spaces.py | 8 +- tests/spaces/test_jordan_invariants.py | 12 +- tests/spaces/test_space_base.py | 18 +- tests/spaces/test_space_checks.py | 10 +- tests/spaces/test_stacked_space.py | 5 +- tests/spaces/test_tree_space.py | 24 +- .../test_tree_spectral_decomposition.py | 17 +- tests/test_equality.py | 310 ++++++++++++ tests/test_repr.py | 2 +- tests/test_weighted_tikhonov.py | 17 +- 128 files changed, 4754 insertions(+), 2073 deletions(-) delete mode 100644 spacecore/_contextual/_policies.py delete mode 100644 spacecore/_contextual/_state.py create mode 100644 spacecore/_errors.py create mode 100644 spacecore/_lazy_algebra.py create mode 100644 spacecore/backend/_container.py create mode 100644 spacecore/backend/_optional.py create mode 100644 spacecore/backend/_registry.py delete mode 100644 spacecore/backend/jax/_pytree.py rename spacecore/{_contextual => contextual}/__init__.py (73%) rename spacecore/{_contextual => contextual}/_bound.py (63%) rename spacecore/{backend => contextual}/_context.py (58%) create mode 100644 spacecore/contextual/_contextual.py create mode 100644 spacecore/contextual/_state.py delete mode 100644 tests/backend/test_jax_pytree_class.py create mode 100644 tests/backend/test_optional_guard.py create mode 100644 tests/backend/test_pytree_registry.py create mode 100644 tests/context/_contracts.py create mode 100644 tests/context/test_ambient_scoping.py create mode 100644 tests/context/test_context_contracts.py create mode 100644 tests/context/test_context_resolution_policy.py create mode 100644 tests/test_equality.py 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/spacecore/__init__.py b/spacecore/__init__.py index eb5efee..4716ffc 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, @@ -116,23 +105,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", @@ -228,15 +228,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/_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..30628c4 --- /dev/null +++ b/spacecore/_lazy_algebra.py @@ -0,0 +1,128 @@ +"""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 scalar_eq(a: Any, b: Any) -> bool: + """Return whether two scalar-likes are equal, NaN-reflexive, as a real ``bool``. + + 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 Exception: + 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..ea2433c 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. diff --git a/spacecore/backend/_optional.py b/spacecore/backend/_optional.py new file mode 100644 index 0000000..97209e5 --- /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: + """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..b9b705c 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} diff --git a/spacecore/_contextual/__init__.py b/spacecore/contextual/__init__.py similarity index 73% rename from spacecore/_contextual/__init__.py rename to spacecore/contextual/__init__.py index 63b824f..fcf15f0 100644 --- a/spacecore/_contextual/__init__.py +++ b/spacecore/contextual/__init__.py @@ -1,14 +1,19 @@ +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 +21,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..eb4fc7d --- /dev/null +++ b/spacecore/contextual/_contextual.py @@ -0,0 +1,258 @@ +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: + """Active default context: 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..1e3a777 --- /dev/null +++ b/spacecore/contextual/_state.py @@ -0,0 +1,220 @@ +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.""" + 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. + """ + _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. + """ + 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. + + 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/_algebra.py b/spacecore/functional/_algebra.py index 02b425e..94381bc 100644 --- a/spacecore/functional/_algebra.py +++ b/spacecore/functional/_algebra.py @@ -12,33 +12,19 @@ """ 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 +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) -> None: @@ -56,7 +42,6 @@ def _require_same_domain(terms: Any) -> None: ) -@jax_pytree_class class ScaledFunctional(Functional): """ Lazy scalar multiple ``scalar * functional``. @@ -67,14 +52,27 @@ 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) @@ -106,9 +104,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,18 +147,19 @@ 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. @@ -169,9 +168,13 @@ class SumFunctional(Functional): ---------- 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 +184,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 @@ -224,7 +229,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 +249,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 +272,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. @@ -308,8 +306,13 @@ class ZeroFunctional(Functional): Backend context specification. """ - 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") def value(self, x: Any, *args: Any, **kwargs: Any) -> Any: @@ -330,7 +333,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 +352,6 @@ def _convert(self, new_ctx: Context) -> "ZeroFunctional": return ZeroFunctional(self.domain.convert(new_ctx), new_ctx) -@jax_pytree_class class ShiftedFunctional(Functional): """ Affine shift ``functional + offset``: value shifted, gradient unchanged. @@ -360,14 +362,25 @@ 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 @@ -391,9 +404,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,7 +446,7 @@ 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) diff --git a/spacecore/functional/_base.py b/spacecore/functional/_base.py index 420226c..0293fb0 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: @@ -234,13 +245,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..dfa756b 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,11 +65,22 @@ 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 @@ -92,7 +103,7 @@ def value(self, x: Any) -> Any: 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..699fdba 100644 --- a/spacecore/functional/_linear.py +++ b/spacecore/functional/_linear.py @@ -6,7 +6,8 @@ from ._base import Domain, Functional from .._batching import _check_scalar_shape, _leading_batch_size 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 @@ -72,7 +73,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. @@ -100,8 +100,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) @@ -126,7 +127,7 @@ def vvalue(self, xs: Any) -> Any: 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 +156,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. @@ -190,6 +190,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 +219,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 @@ -286,7 +287,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..17d7c2d 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 @@ -44,7 +45,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. @@ -89,6 +89,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 +108,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) @@ -159,7 +160,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/tools/_entropy.py b/spacecore/functional/tools/_entropy.py index 558af38..1ab5b85 100644 --- a/spacecore/functional/tools/_entropy.py +++ b/spacecore/functional/tools/_entropy.py @@ -5,12 +5,12 @@ 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``. @@ -39,8 +39,13 @@ 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") def value(self, x: Any) -> Any: @@ -69,7 +74,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``. @@ -110,8 +114,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) diff --git a/spacecore/functional/tools/_huber.py b/spacecore/functional/tools/_huber.py index 2025d0a..e12bb49 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)``. @@ -40,8 +40,14 @@ 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}.") diff --git a/spacecore/functional/tools/_norms.py b/spacecore/functional/tools/_norms.py index 57288bf..473ff36 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``. @@ -41,8 +41,13 @@ 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") def value(self, x: Any) -> Any: @@ -73,7 +78,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``. @@ -105,8 +109,14 @@ 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}.") @@ -137,7 +147,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)``. @@ -154,4 +166,4 @@ def L1NormFunctional( 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/_spectral.py b/spacecore/functional/tools/_spectral.py index 8be3fba..0a4120c 100644 --- a/spacecore/functional/tools/_spectral.py +++ b/spacecore/functional/tools/_spectral.py @@ -21,13 +21,13 @@ 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 ...space import JordanAlgebraSpace from ..._checks import checked_method from ._coordinate import _CoordinateFunctional, lp_coordinate_grad, lp_value -@jax_pytree_class class SpectralLpNormFunctional(_CoordinateFunctional[Domain]): r""" Schatten ``p``-norm ``F(X) = (sum_i |lambda_i(X)|^p)^{1/p}`` for ``p >= 1``. @@ -62,8 +62,14 @@ class SpectralLpNormFunctional(_CoordinateFunctional[Domain]): 5.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) if not isinstance(self.domain, JordanAlgebraSpace): raise TypeError( "SpectralLpNormFunctional requires a Jordan-algebra domain with a " @@ -104,7 +110,9 @@ def _convert(self, new_ctx: Context) -> "SpectralLpNormFunctional": def NuclearNormFunctional( - dom: Domain, ctx: Context | str | None = None + dom: Domain, + ctx: Context | str | None = None, + check_level: CheckLevel | bool | None = None, ) -> "SpectralLpNormFunctional[Domain]": r""" Nuclear (trace) norm, a thin wrapper for ``SpectralLpNormFunctional(X, 1)``. @@ -123,4 +131,4 @@ def NuclearNormFunctional( SpectralLpNormFunctional The ``p = 1`` instance of :class:`SpectralLpNormFunctional`. """ - return SpectralLpNormFunctional(dom, 1.0, ctx) + return SpectralLpNormFunctional(dom, 1.0, ctx, check_level=check_level) diff --git a/spacecore/kernels/core/algebra.py b/spacecore/kernels/core/algebra.py index 25c1434..b07f23a 100644 --- a/spacecore/kernels/core/algebra.py +++ b/spacecore/kernels/core/algebra.py @@ -88,9 +88,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/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..1ab6c20 100644 --- a/spacecore/linop/_algebra.py +++ b/spacecore/linop/_algebra.py @@ -1,16 +1,15 @@ from __future__ import annotations 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 ..contextual import resolve_context_priority 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,25 @@ ) -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. +def _require_same_context(ops: Sequence[LinOp]) -> Context: + """Return the common mathematical context for algebra operands or raise. - 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). + 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. """ - 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 - - -def _require_same_context(ops: Sequence[LinOp]) -> Context: - """Return the common context for algebra operands or raise.""" 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 +53,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 +68,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 +91,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 +109,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 +143,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 +197,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 +217,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 +231,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 +311,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 +339,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 +357,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 +369,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 +389,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 +470,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 +500,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 +519,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 +533,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 +547,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 +624,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 +649,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 +669,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 +681,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 +728,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 +750,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 +767,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 +827,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 +854,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 +900,11 @@ 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`. Returns ------- @@ -918,6 +928,7 @@ 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, ) -> None: """ Initialize a matrix-free linear operator. @@ -943,6 +954,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 +972,7 @@ 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) self.apply_fn = apply self.rapply_fn = rapply self.vapply_fn = vapply @@ -1183,7 +1197,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. @@ -1249,7 +1263,6 @@ def _convert(self, new_ctx: Context) -> MatrixFreeLinOp: @core_kernels("adjoint") -@jax_pytree_class class _AdjointViewLinOp(LinOp[Codomain, Domain]): """ Hermitian-adjoint view of a linear operator. @@ -1260,11 +1273,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 +1331,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..92201e0 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. @@ -39,6 +41,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 ---------- @@ -62,8 +69,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: @@ -358,13 +371,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..7f0002e 100644 --- a/spacecore/linop/_dense.py +++ b/spacecore/linop/_dense.py @@ -5,6 +5,7 @@ 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 +15,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. @@ -48,6 +48,11 @@ class DenseLinOp(LinOp[Domain, Codomain]): ``A``. 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 +76,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 @@ -81,7 +87,7 @@ def __init__( _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 +223,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/_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/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..08a71e8 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: diff --git a/spacecore/space/base/_space.py b/spacecore/space/base/_space.py index 417cfc2..85e1b11 100644 --- a/spacecore/space/base/_space.py +++ b/spacecore/space/base/_space.py @@ -3,9 +3,9 @@ from typing import Any, ClassVar, Literal from ..._check_policy import CheckLevel, check_level_at_least, normalize_check_level -from ..._contextual import ContextBound +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 +17,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` @@ -43,7 +52,7 @@ def field(self) -> Literal["real", "complex"]: 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/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..c1d8dd8 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. @@ -113,6 +119,11 @@ class ElementwiseJordanSpace(JordanAlgebraSpace, DenseCoordinateSpace, StarSpace used. inner_product : InnerProduct or None, optional Alias for ``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`. """ def __new__( @@ -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. @@ -227,6 +239,11 @@ class EuclideanElementwiseJordanSpace(ElementwiseJordanSpace, EuclideanJordanAlg with Euclidean coordinate geometry. inner_product : InnerProduct or None, optional Alias for ``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`. """ def __init__( @@ -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..bcddae7 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 ---------- @@ -52,12 +58,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 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/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..45000e8 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])) + 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..6f8e3f7 100644 --- a/tests/functional/test_algebra.py +++ b/tests/functional/test_algebra.py @@ -85,10 +85,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]) diff --git a/tests/functional/test_generated_functionals.py b/tests/functional/test_generated_functionals.py index a0dec16..b6b07a4 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,9 +247,11 @@ 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"): functional.value(ctx.asarray([1.0, 2.0])) @@ -251,16 +259,17 @@ def test_standard_functional_checks_reject_nonscalar_output(check_level): @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_linop_quadratic_form.py b/tests/functional/test_linop_quadratic_form.py index 58e4927..efea9fb 100644 --- a/tests/functional/test_linop_quadratic_form.py +++ b/tests/functional/test_linop_quadratic_form.py @@ -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 batch"): + 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..856d1a1 100644 --- a/tests/functional/test_matrix_free_linear_functional.py +++ b/tests/functional/test_matrix_free_linear_functional.py @@ -65,9 +65,9 @@ def test_value_enforces_scalar_output_under_standard_checks(self, numpy_ctx): 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/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..4caf0b5 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]: @@ -502,20 +504,21 @@ 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.float64) for d in dtypes): + cases.extend(_battery_cases(np.float64, check_level)) + cases.append(_spectral_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/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..a9b523e 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,9 @@ import pytest import spacecore as sc +from spacecore._lazy_algebra import scalar_eq from spacecore.linop._algebra import ( _conjugate_scalar, - _is_one_scalar, - _is_zero_scalar, - _scalar_equal, is_scalar_like, make_composed, make_scaled, @@ -60,47 +58,34 @@ 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_on_exception(self): - """``_scalar_equal`` swallows exceptions and returns False.""" + """``scalar_eq`` swallows a raising ``__eq__`` and returns False.""" class _Bad: def __eq__(self, other): raise RuntimeError("boom") - assert _scalar_equal(_Bad(), 0) is False - - -class TestIsZeroScalar: - @pytest.mark.parametrize("value, expected", [ - (0, True), - (0.0, True), - (1, False), - (np.float64(0.0), True), - ]) - def test_truth_table(self, value, expected): - assert _is_zero_scalar(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 + assert scalar_eq(_Bad(), 0) is False # =========================================================================== @@ -292,16 +277,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..a892266 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 @@ -626,7 +626,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) @@ -645,6 +645,7 @@ def rapply(w): op = sc.MatrixFreeLinOp(apply, rapply, domain, codomain, numpy_ctx) 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 +655,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 +675,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 +684,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 +693,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 +707,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 +737,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 +971,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..58852e9 100644 --- a/tests/spaces/test_hermitian_space.py +++ b/tests/spaces/test_hermitian_space.py @@ -104,8 +104,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 +120,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 +146,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..7c471b0 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) @@ -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_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 From b4957ff8edd405383c363d0eba727d2c198f450a Mon Sep 17 00:00:00 2001 From: Pavlo Pelikh Date: Sat, 25 Jul 2026 17:50:10 -0300 Subject: [PATCH 2/6] Refactor functional submodule --- CHANGELOG.md | 202 +++++++++++ docs/source/api/functionals.rst | 90 ++++- spacecore/__init__.py | 28 +- spacecore/_batching.py | 10 +- spacecore/_checks.py | 39 +++ spacecore/_lazy_algebra.py | 57 ++- spacecore/backend/_ops.py | 31 ++ spacecore/functional/__init__.py | 19 +- spacecore/functional/_algebra.py | 324 ++++++++++++++++- spacecore/functional/_base.py | 47 ++- spacecore/functional/_composed.py | 54 ++- spacecore/functional/_linear.py | 24 +- spacecore/functional/_quadratic.py | 9 +- spacecore/functional/_realified.py | 238 +++++++++++++ spacecore/functional/tools/__init__.py | 11 +- spacecore/functional/tools/_coordinate.py | 15 + spacecore/functional/tools/_entropy.py | 26 +- spacecore/functional/tools/_huber.py | 22 +- spacecore/functional/tools/_norms.py | 4 +- spacecore/functional/tools/_proximal.py | 20 ++ spacecore/functional/tools/_spectral.py | 283 +++++++++++---- spacecore/kernels/core/algebra.py | 15 +- spacecore/kernels/core/functional.py | 17 +- spacecore/linop/_algebra.py | 129 +++++-- spacecore/linop/_base.py | 19 + spacecore/linop/_dense.py | 32 +- spacecore/linop/_metric.py | 21 +- spacecore/opfamily.py | 329 ++++++++++++++++++ spacecore/space/base/_space.py | 51 ++- spacecore/space/concrete/_hermitian.py | 29 ++ tests/context/test_check_policy.py | 2 +- tests/functional/test_algebra.py | 174 ++++++++- tests/functional/test_composed_functional.py | 166 +++++++++ tests/functional/test_functional_base.py | 103 +++++- .../functional/test_generated_functionals.py | 2 +- tests/functional/test_linop_quadratic_form.py | 2 +- .../test_matrix_free_linear_functional.py | 2 +- tests/functional/tools/test_spectral.py | 38 +- tests/generators/functionals.py | 95 ++++- tests/linops/test_algebra_factories.py | 70 +++- tests/linops/test_algebra_linops.py | 21 +- tests/spaces/test_hermitian_space.py | 59 ++++ .../test_tree_spectral_decomposition.py | 6 +- tests/test_opfamily.py | 289 +++++++++++++++ 44 files changed, 3019 insertions(+), 205 deletions(-) create mode 100644 spacecore/functional/_realified.py create mode 100644 spacecore/opfamily.py create mode 100644 tests/test_opfamily.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 35ded2a..3ab458d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,208 @@ and the project adheres to [Semantic Versioning](https://semver.org/). ## [Unreleased] +### Added + +- **`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 + +- **`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 + +- **`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.3] — 2026-07-25 + +### 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`. + +### 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. + +### 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.** + ## [0.4.2] — 2026-07-01 ### Added 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/spacecore/__init__.py b/spacecore/__init__.py index 4716ffc..e65d33e 100644 --- a/spacecore/__init__.py +++ b/spacecore/__init__.py @@ -26,8 +26,14 @@ make_sum, ) from .functional import ( + RealifiedFunctional, + realify, + SpectralFunctional, + eigenvalue_space, + spectralize, ComposedFunctional, Functional, + ConstantFunctional, HuberFunctional, InnerProductFunctional, KLDivergenceFunctional, @@ -38,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, @@ -55,6 +63,11 @@ prox_l1, prox_l2sq, ) +from .opfamily import ( + FunctionalScaledOperator, + OperatorFamily, + make_functional_scaled_operator, +) from .linalg import ( CGResult, ExpmMultiplyResult, @@ -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", 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/_lazy_algebra.py b/spacecore/_lazy_algebra.py index 30628c4..01b03be 100644 --- a/spacecore/_lazy_algebra.py +++ b/spacecore/_lazy_algebra.py @@ -29,8 +29,38 @@ def is_scalar_like(value: Any) -> bool: 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 equal, NaN-reflexive, as a real ``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 @@ -43,7 +73,30 @@ def scalar_eq(a: Any, b: Any) -> bool: # ``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 Exception: + 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 diff --git a/spacecore/backend/_ops.py b/spacecore/backend/_ops.py index ea2433c..0e94043 100644 --- a/spacecore/backend/_ops.py +++ b/spacecore/backend/_ops.py @@ -415,6 +415,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): 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 94381bc..dfd0686 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 @@ -27,17 +36,24 @@ ) -def _require_same_domain(terms: Any) -> None: +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}." ) @@ -76,7 +92,7 @@ def __init__( 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) @@ -194,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) @@ -314,7 +330,7 @@ def __init__( ) -> 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) @@ -352,6 +368,121 @@ def _convert(self, new_ctx: Context) -> "ZeroFunctional": return ZeroFunctional(self.domain.convert(new_ctx), new_ctx) +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. @@ -384,7 +515,7 @@ def __init__( 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) @@ -451,3 +582,178 @@ def make_shifted_functional(functional: Functional, offset: Any) -> 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 0293fb0..83995d0 100644 --- a/spacecore/functional/_base.py +++ b/spacecore/functional/_base.py @@ -195,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``. + + 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 + + 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, scalar: Any) -> "Functional": - """Return the lazy left scalar multiple ``scalar * self``.""" - from ._algebra import is_scalar_like, make_scaled_functional + 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 not is_scalar_like(scalar): - return NotImplemented - return make_scaled_functional(scalar, self) + if isinstance(other, Functional): + return make_functional_product(other, self) + if isinstance(other, LinOp): + from ..opfamily import make_functional_scaled_operator + + 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) + @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)) diff --git a/spacecore/functional/_composed.py b/spacecore/functional/_composed.py index dfa756b..3fb301c 100644 --- a/spacecore/functional/_composed.py +++ b/spacecore/functional/_composed.py @@ -84,7 +84,7 @@ def __init__( 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)``. @@ -101,6 +101,58 @@ 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.same_math(other): # Tier 1: backend diff --git a/spacecore/functional/_linear.py b/spacecore/functional/_linear.py index 699fdba..e69757b 100644 --- a/spacecore/functional/_linear.py +++ b/spacecore/functional/_linear.py @@ -4,7 +4,7 @@ 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 .._check_policy import CheckLevel from ..contextual import Context @@ -112,18 +112,15 @@ 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.""" @@ -239,7 +236,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. @@ -254,11 +251,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: """ diff --git a/spacecore/functional/_quadratic.py b/spacecore/functional/_quadratic.py index 17d7c2d..838b2a6 100644 --- a/spacecore/functional/_quadratic.py +++ b/spacecore/functional/_quadratic.py @@ -121,7 +121,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) @@ -145,13 +145,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: diff --git a/spacecore/functional/_realified.py b/spacecore/functional/_realified.py new file mode 100644 index 0000000..093e2c2 --- /dev/null +++ b/spacecore/functional/_realified.py @@ -0,0 +1,238 @@ +"""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 1ab5b85..a1cf2d1 100644 --- a/spacecore/functional/tools/_entropy.py +++ b/spacecore/functional/tools/_entropy.py @@ -1,4 +1,24 @@ -"""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 @@ -47,7 +67,7 @@ def __init__( ) -> 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 @@ -126,7 +146,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 e12bb49..55ddebf 100644 --- a/spacecore/functional/tools/_huber.py +++ b/spacecore/functional/tools/_huber.py @@ -18,7 +18,25 @@ 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``. + + 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). Parameters ---------- @@ -53,7 +71,7 @@ def __init__( 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 473ff36..9758d6e 100644 --- a/spacecore/functional/tools/_norms.py +++ b/spacecore/functional/tools/_norms.py @@ -49,7 +49,7 @@ def __init__( ) -> 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)) @@ -122,7 +122,7 @@ def __init__( 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) diff --git a/spacecore/functional/tools/_proximal.py b/spacecore/functional/tools/_proximal.py index 3878ccf..e920514 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 diff --git a/spacecore/functional/tools/_spectral.py b/spacecore/functional/tools/_spectral.py index 0a4120c..a832943 100644 --- a/spacecore/functional/tools/_spectral.py +++ b/spacecore/functional/tools/_spectral.py @@ -1,54 +1,125 @@ -"""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 .._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 -class SpectralLpNormFunctional(_CoordinateFunctional[Domain]): - r""" - Schatten ``p``-norm ``F(X) = (sum_i |lambda_i(X)|^p)^{1/p}`` for ``p >= 1``. +def eigenvalue_space(dom: Any, check_level: CheckLevel | bool | None = None) -> Any: + r"""Return the real coordinate space the spectrum of ``dom`` lives in. - ``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. + 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 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 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) + + +class SpectralFunctional(_CoordinateFunctional[Domain]): + r"""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. + + 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 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,68 +127,138 @@ 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, + base: Functional, ctx: Context | str | None = None, check_level: CheckLevel | bool | None = None, ) -> None: super().__init__(dom, ctx, check_level=check_level) - if not isinstance(self.domain, JordanAlgebraSpace): + 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__}." ) - 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 + 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." + ) + 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 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. - def _convert(self, new_ctx: Context) -> "SpectralLpNormFunctional": - """Convert this functional to ``new_ctx``.""" - return SpectralLpNormFunctional(self.domain.convert(new_ctx), self.p, new_ctx) + 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, check_level: CheckLevel | bool | None = None, -) -> "SpectralLpNormFunctional[Domain]": +) -> "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 ---------- @@ -125,10 +266,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, check_level=check_level) + 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 b07f23a..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: 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/linop/_algebra.py b/spacecore/linop/_algebra.py index 1ab6c20..3c69689 100644 --- a/spacecore/linop/_algebra.py +++ b/spacecore/linop/_algebra.py @@ -1,5 +1,6 @@ from __future__ import annotations +import warnings from math import prod from typing import Any, Callable, Sequence, cast @@ -7,7 +8,6 @@ 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 .._repr import summarize_value from ..contextual import Context from ..kernels import core_kernels @@ -28,6 +28,13 @@ ) +#: 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 mathematical context for algebra operands or raise. @@ -900,6 +907,22 @@ class MatrixFreeLinOp(LinOp[Domain, Codomain]): Optional callable with signature ``rvapply(ys: Any) -> Any`` for batched adjoint application. If omitted, backend ``vmap`` fallback is used. + 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". 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 @@ -929,6 +952,8 @@ def __init__( 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. @@ -973,6 +998,38 @@ def __init__( 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, 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 @@ -980,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)): @@ -1078,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: @@ -1251,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, @@ -1259,6 +1329,7 @@ def _convert(self, new_ctx: Context) -> MatrixFreeLinOp: new_ctx, self.vapply_fn, self.rvapply_fn, + euclidean_adjoint=False, ) diff --git a/spacecore/linop/_base.py b/spacecore/linop/_base.py index 92201e0..ac82914 100644 --- a/spacecore/linop/_base.py +++ b/spacecore/linop/_base.py @@ -32,6 +32,25 @@ class LinOp(PyTreeNode, 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**, 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. + + 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``. + Parameters ---------- dom : Space diff --git a/spacecore/linop/_dense.py b/spacecore/linop/_dense.py index 7f0002e..519e4b2 100644 --- a/spacecore/linop/_dense.py +++ b/spacecore/linop/_dense.py @@ -1,5 +1,7 @@ from __future__ import annotations +import warnings + from functools import cached_property from math import prod from typing import Any, cast @@ -45,7 +47,19 @@ 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 @@ -84,6 +98,22 @@ 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") diff --git a/spacecore/linop/_metric.py b/spacecore/linop/_metric.py index 2cf7de8..4cbfe05 100644 --- a/spacecore/linop/_metric.py +++ b/spacecore/linop/_metric.py @@ -82,7 +82,26 @@ 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. + + References + ---------- + .. [Conway] J. B. Conway, *A Course in Functional Analysis*, 2nd ed., + Springer, 1990, II.2.2/II.2.4 — existence and uniqueness of the adjoint + from the Riesz representation theorem. The formula here is that + existence proof made computational: ``R`` is the Riesz isomorphism. + """ 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/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/space/base/_space.py b/spacecore/space/base/_space.py index 85e1b11..d65e8b0 100644 --- a/spacecore/space/base/_space.py +++ b/spacecore/space/base/_space.py @@ -3,6 +3,7 @@ from typing import Any, ClassVar, Literal from ..._check_policy import CheckLevel, check_level_at_least, normalize_check_level +from ..._lazy_algebra import is_recognizably_nonreal from ...contextual import ContextBound from ..._repr import field_symbol from ...contextual import Context @@ -40,15 +41,61 @@ def __init__( 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. diff --git a/spacecore/space/concrete/_hermitian.py b/spacecore/space/concrete/_hermitian.py index bcddae7..170f558 100644 --- a/spacecore/space/concrete/_hermitian.py +++ b/spacecore/space/concrete/_hermitian.py @@ -51,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, @@ -96,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/tests/context/test_check_policy.py b/tests/context/test_check_policy.py index 45000e8..8b8e6be 100644 --- a/tests/context/test_check_policy.py +++ b/tests/context/test_check_policy.py @@ -128,7 +128,7 @@ def test_functional_scalar_output_shape_is_standard(): standard_functional = sc.MatrixFreeLinearFunctional( lambda _x: ctx.asarray([1.0]), standard_space, ctx, check_level="standard" ) - with pytest.raises(ValueError, match="scalar batch output"): + with pytest.raises(ValueError, match="scalar output"): standard_functional.value(ctx.asarray([1.0, 2.0])) diff --git a/tests/functional/test_algebra.py b/tests/functional/test_algebra.py index 6f8e3f7..1cefb05 100644 --- a/tests/functional/test_algebra.py +++ b/tests/functional/test_algebra.py @@ -73,7 +73,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) @@ -276,6 +279,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..caa18b0 100644 --- a/tests/functional/test_composed_functional.py +++ b/tests/functional/test_composed_functional.py @@ -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..ab0ac4d 100644 --- a/tests/functional/test_functional_base.py +++ b/tests/functional/test_functional_base.py @@ -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 b6b07a4..1a6a458 100644 --- a/tests/functional/test_generated_functionals.py +++ b/tests/functional/test_generated_functionals.py @@ -253,7 +253,7 @@ def test_standard_functional_checks_reject_nonscalar_output(check_level): 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])) diff --git a/tests/functional/test_linop_quadratic_form.py b/tests/functional/test_linop_quadratic_form.py index efea9fb..7f29cd0 100644 --- a/tests/functional/test_linop_quadratic_form.py +++ b/tests/functional/test_linop_quadratic_form.py @@ -75,7 +75,7 @@ def test_rejects_nonscalar_constant(self): with sc.use_check_level("none"): space = sc.DenseCoordinateSpace((2,), ctx) Q = sc.IdentityLinOp(space, ctx) - with pytest.raises(ValueError, match="scalar batch"): + 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): diff --git a/tests/functional/test_matrix_free_linear_functional.py b/tests/functional/test_matrix_free_linear_functional.py index 856d1a1..ffec53f 100644 --- a/tests/functional/test_matrix_free_linear_functional.py +++ b/tests/functional/test_matrix_free_linear_functional.py @@ -61,7 +61,7 @@ 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): 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/functionals.py b/tests/generators/functionals.py index 4caf0b5..8cb60fe 100644 --- a/tests/generators/functionals.py +++ b/tests/generators/functionals.py @@ -246,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 @@ -261,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, @@ -461,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 (W4) 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", @@ -493,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}", ) @@ -518,7 +604,10 @@ def functional_cases( 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/linops/test_algebra_factories.py b/tests/linops/test_algebra_factories.py index a9b523e..fbc384d 100644 --- a/tests/linops/test_algebra_factories.py +++ b/tests/linops/test_algebra_factories.py @@ -19,7 +19,8 @@ import pytest import spacecore as sc -from spacecore._lazy_algebra import scalar_eq +from spacecore._lazy_algebra import is_recognizably_nonreal, scalar_eq +from tests._helpers import has_jax from spacecore.linop._algebra import ( _conjugate_scalar, is_scalar_like, @@ -79,13 +80,74 @@ class TestScalarEq: def test_truth_table(self, a, b, expected): assert scalar_eq(a, b) is expected - def test_returns_false_on_exception(self): - """``scalar_eq`` swallows a raising ``__eq__`` and returns False.""" + 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. + + 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_eq(_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.""" + + @pytest.mark.parametrize("value, expected", [ + (1, False), + (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_recognizably_nonreal(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)] # =========================================================================== diff --git a/tests/linops/test_algebra_linops.py b/tests/linops/test_algebra_linops.py index a892266..d29a3d5 100644 --- a/tests/linops/test_algebra_linops.py +++ b/tests/linops/test_algebra_linops.py @@ -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, ) @@ -643,7 +654,9 @@ 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]) diff --git a/tests/spaces/test_hermitian_space.py b/tests/spaces/test_hermitian_space.py index 58852e9..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 # =========================================================================== diff --git a/tests/spaces/test_tree_spectral_decomposition.py b/tests/spaces/test_tree_spectral_decomposition.py index 7c471b0..52d3333 100644 --- a/tests/spaces/test_tree_spectral_decomposition.py +++ b/tests/spaces/test_tree_spectral_decomposition.py @@ -182,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) 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 From 0dfe7296c87d53ddab53a8963a74ae084561ad3f Mon Sep 17 00:00:00 2001 From: Pavlo Pelikh Date: Fri, 7 Aug 2026 14:00:14 -0300 Subject: [PATCH 3/6] Refactor functional submodule --- spacecore/backend/_ops.py | 236 ++++++++++++++---- spacecore/backend/torch/_ops.py | 16 +- spacecore/functional/tools/_proximal.py | 33 ++- spacecore/linop/_base.py | 17 +- spacecore/linop/_metric.py | 7 - spacecore/space/base/_coordinate.py | 12 +- tests/backend/_references.py | 187 ++++++++++++++ tests/functional/test_algebra.py | 9 +- tests/functional/test_composed_functional.py | 2 +- tests/functional/test_functional_base.py | 2 +- .../test_inner_product_functional.py | 2 +- tests/functional/test_linear_functional.py | 2 +- tests/functional/test_linop_quadratic_form.py | 2 +- .../test_matrix_free_linear_functional.py | 2 +- tests/functional/test_quadratic_form.py | 2 +- tests/functional/tools/test_proximal.py | 139 +++++++++++ tests/generators/functionals.py | 4 +- .../generators/test_registry_completeness.py | 2 +- 18 files changed, 605 insertions(+), 71 deletions(-) diff --git a/spacecore/backend/_ops.py b/spacecore/backend/_ops.py index 0e94043..fefce42 100644 --- a/spacecore/backend/_ops.py +++ b/spacecore/backend/_ops.py @@ -320,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"): @@ -536,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) @@ -823,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( @@ -837,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, @@ -846,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, @@ -855,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, @@ -865,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, @@ -924,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: + """Unit phase of each column's largest-magnitude 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, @@ -941,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, @@ -958,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, @@ -967,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, @@ -1041,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/torch/_ops.py b/spacecore/backend/torch/_ops.py index b9b705c..bf19379 100644 --- a/spacecore/backend/torch/_ops.py +++ b/spacecore/backend/torch/_ops.py @@ -499,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, @@ -520,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, @@ -544,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/functional/tools/_proximal.py b/spacecore/functional/tools/_proximal.py index e920514..e04f1eb 100644 --- a/spacecore/functional/tools/_proximal.py +++ b/spacecore/functional/tools/_proximal.py @@ -164,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 @@ -179,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 ------- @@ -188,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)``). @@ -201,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 @@ -209,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 ------- @@ -218,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/linop/_base.py b/spacecore/linop/_base.py index ac82914..d0b739e 100644 --- a/spacecore/linop/_base.py +++ b/spacecore/linop/_base.py @@ -129,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.""" @@ -152,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: diff --git a/spacecore/linop/_metric.py b/spacecore/linop/_metric.py index 4cbfe05..2fc8450 100644 --- a/spacecore/linop/_metric.py +++ b/spacecore/linop/_metric.py @@ -94,13 +94,6 @@ def metric_rapply(domain, codomain, euclidean_rapply, y): Getting this wrong is silent: a coordinate transpose satisfies the identity on every Euclidean space, so only a non-Euclidean test can detect it. - - References - ---------- - .. [Conway] J. B. Conway, *A Course in Functional Analysis*, 2nd ed., - Springer, 1990, II.2.2/II.2.4 — existence and uniqueness of the adjoint - from the Riesz representation theorem. The formula here is that - existence proof made computational: ``R`` is the Riesz isomorphism. """ if domain.is_euclidean and codomain.is_euclidean: return euclidean_rapply(y) diff --git a/spacecore/space/base/_coordinate.py b/spacecore/space/base/_coordinate.py index 08a71e8..9f13d22 100644 --- a/spacecore/space/base/_coordinate.py +++ b/spacecore/space/base/_coordinate.py @@ -58,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/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/functional/test_algebra.py b/tests/functional/test_algebra.py index 1cefb05..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 diff --git a/tests/functional/test_composed_functional.py b/tests/functional/test_composed_functional.py index caa18b0..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``, diff --git a/tests/functional/test_functional_base.py b/tests/functional/test_functional_base.py index ab0ac4d..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. 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 7f29cd0..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. diff --git a/tests/functional/test_matrix_free_linear_functional.py b/tests/functional/test_matrix_free_linear_functional.py index ffec53f..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). 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/generators/functionals.py b/tests/generators/functionals.py index 8cb60fe..bdd976b 100644 --- a/tests/generators/functionals.py +++ b/tests/generators/functionals.py @@ -430,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 @@ -463,7 +463,7 @@ def _algebra_case( def _algebra_cases(dtype: Any, check_level: sc.CheckLevel | str) -> tuple[FunctionalCase, ...]: """Generated cases for the functional-algebra nodes. - Covers Scaled/Sum/Shifted/Zero (W4) plus the multiplicative nodes + Covers Scaled/Sum/Shifted/Zero plus the multiplicative nodes Constant/Product. """ c2 = np.asarray([0.5, -0.25, 1.0], dtype=dtype) 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). From 76f86fe8081167f296843bd9c93429ae9583e37d Mon Sep 17 00:00:00 2001 From: Pavlo Pelikh Date: Sat, 12 Sep 2026 03:09:03 -0400 Subject: [PATCH 4/6] Fix CI: stale Context API in jit_audit, pin the ruff rule set The jit audit still built its context with `Context(..., enable_checks=False)`. Validation policy moved off the context onto the bound object in this release, so that keyword now raises TypeError and the audit step failed. Drop it and lower the ambient default instead -- argument validation would otherwise add non-traced Python work that muddies the retrace counts the script measures. Separately, the lint configuration named no `select`, so `ruff check .` enforced whatever ruff's implicit default happened to be. 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 turned a green tree red without a single source edit (569 findings on master, which last ran CI against an older ruff). Name the rules explicitly; the result is now identical under 0.15.8 and 0.16.7, and widening the set becomes a deliberate commit rather than a release surprise. Co-Authored-By: Claude Opus 5 --- pyproject.toml | 7 +++++++ scripts/jit_audit.py | 8 +++++++- 2 files changed, 14 insertions(+), 1 deletion(-) 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) From 396a4a65eddbcc09f8f1a03af349932e3854eda6 Mon Sep 17 00:00:00 2001 From: Pavlo Pelikh Date: Sat, 12 Sep 2026 03:09:14 -0400 Subject: [PATCH 5/6] Clear the pydocstyle and numpydoc audits Two CI gates were failing on documentation rather than behaviour. `ruff check --select D` found five regressions from the submodule refactors: the newly *public* `spacecore.contextual` package needed a module docstring (the old `_contextual` was private, so D104 never applied), two summaries were not in the imperative mood, one ran its summary into its description, and `functional/_realified.py` carried LaTeX in a non-raw docstring. `scripts/docstring_audit.py` reported 36 numpydoc issues. These predate the branch -- master reports the identical set -- but they only became a CI failure now that numpydoc 1.10 is what gets installed, and they block this PR either way. Most are the `check_level` constructor argument going undocumented across the spaces, operators and functionals that gained it; the rest are summary placement, documented parameter order that no longer matched the signature, section ordering, and missing Returns/Yields on the ambient-state helpers. Also fix two Sphinx warnings introduced by this branch: the `[Conway]` citation on `LinOp` was defined but never cited, and `SumFunctional`'s summary put `F_1 + ... + F_n` inside an inline literal, where autosummary cuts the sentence at the ellipsis and leaves the literal unterminated. Documented behaviour was verified against the runtime rather than the type hints -- notably, a boolean `check_level` is annotated but rejected, so the prose describes the four literals only. Co-Authored-By: Claude Opus 5 --- spacecore/backend/_ops.py | 2 +- spacecore/backend/_optional.py | 2 +- spacecore/contextual/__init__.py | 10 +++++ spacecore/contextual/_contextual.py | 6 ++- spacecore/contextual/_state.py | 51 +++++++++++++++++++++-- spacecore/functional/_algebra.py | 5 ++- spacecore/functional/_linear.py | 9 ++++ spacecore/functional/_quadratic.py | 6 +++ spacecore/functional/_realified.py | 12 +++--- spacecore/functional/tools/_entropy.py | 6 +++ spacecore/functional/tools/_huber.py | 21 ++++++---- spacecore/functional/tools/_norms.py | 9 ++++ spacecore/functional/tools/_spectral.py | 9 ++-- spacecore/linop/_algebra.py | 10 ++--- spacecore/linop/_base.py | 32 +++++++------- spacecore/space/base/_inner_product.py | 3 ++ spacecore/space/base/_jordan.py | 6 +++ spacecore/space/base/_star.py | 3 ++ spacecore/space/base/_vector.py | 3 ++ spacecore/space/concrete/_dense_vector.py | 8 ++-- 20 files changed, 162 insertions(+), 51 deletions(-) diff --git a/spacecore/backend/_ops.py b/spacecore/backend/_ops.py index fefce42..7143be8 100644 --- a/spacecore/backend/_ops.py +++ b/spacecore/backend/_ops.py @@ -995,7 +995,7 @@ def einsum(self, subscripts: str, *operands: DenseArray) -> DenseArray: # Cost is O(n^2) on top of an O(n^3) decomposition. def _column_phase(self, vectors: DenseArray) -> DenseArray: - """Unit phase of each column's largest-magnitude entry, shaped ``(..., 1, k)``.""" + """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) diff --git a/spacecore/backend/_optional.py b/spacecore/backend/_optional.py index 97209e5..b5d7adb 100644 --- a/spacecore/backend/_optional.py +++ b/spacecore/backend/_optional.py @@ -68,7 +68,7 @@ def discover_entry_point_ops() -> list[type[BackendOps]]: def _backend_absent(exc: ModuleNotFoundError, dep: str) -> bool: - """True iff the failure is the backend dependency itself being missing.""" + """Return True iff the failure is the backend dependency itself being missing.""" return exc.name == dep diff --git a/spacecore/contextual/__init__.py b/spacecore/contextual/__init__.py index fcf15f0..aa360f9 100644 --- a/spacecore/contextual/__init__.py +++ b/spacecore/contextual/__init__.py @@ -1,3 +1,13 @@ +"""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 ( diff --git a/spacecore/contextual/_contextual.py b/spacecore/contextual/_contextual.py index eb4fc7d..ee4da92 100644 --- a/spacecore/contextual/_contextual.py +++ b/spacecore/contextual/_contextual.py @@ -124,8 +124,10 @@ def ctx_from_ops(self, ops: BackendOps, dtype: DType | None = None) -> Context: @property def default_ctx(self) -> Context: - """Active default context: the scoped override if one is installed here, - otherwise the process-wide baseline. + """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 diff --git a/spacecore/contextual/_state.py b/spacecore/contextual/_state.py index 1e3a777..27b9ac9 100644 --- a/spacecore/contextual/_state.py +++ b/spacecore/contextual/_state.py @@ -43,26 +43,53 @@ def set_context( def get_check_level() -> CheckLevel: - """Return the ambient default validation level applied to new bound objects.""" + """ + 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. + """ + 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. + """ + 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 @@ -73,7 +100,8 @@ def use_context( ctx: Context | BackendFamily | str | None = None, dtype: Any = None, ): - """Temporarily override the default context within a ``with`` block. + """ + 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 @@ -85,6 +113,21 @@ def use_context( 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 diff --git a/spacecore/functional/_algebra.py b/spacecore/functional/_algebra.py index dfd0686..5619601 100644 --- a/spacecore/functional/_algebra.py +++ b/spacecore/functional/_algebra.py @@ -178,7 +178,7 @@ def make_scaled_functional(scalar: Any, functional: Functional) -> Functional: 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 ---------- @@ -320,6 +320,9 @@ 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__( diff --git a/spacecore/functional/_linear.py b/spacecore/functional/_linear.py index e69757b..68ef2fb 100644 --- a/spacecore/functional/_linear.py +++ b/spacecore/functional/_linear.py @@ -43,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 @@ -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 ---------- @@ -173,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 ------- diff --git a/spacecore/functional/_quadratic.py b/spacecore/functional/_quadratic.py index 838b2a6..7696c40 100644 --- a/spacecore/functional/_quadratic.py +++ b/spacecore/functional/_quadratic.py @@ -27,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: @@ -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 ---------- diff --git a/spacecore/functional/_realified.py b/spacecore/functional/_realified.py index 093e2c2..883f86c 100644 --- a/spacecore/functional/_realified.py +++ b/spacecore/functional/_realified.py @@ -1,4 +1,4 @@ -"""Real-coordinate view of a complex-domain functional. +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 @@ -13,8 +13,8 @@ .. math:: - \\partial F/\\partial a = \\operatorname{Re} g_v, \\qquad - \\partial F/\\partial b = \\operatorname{Im} g_v, + \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 @@ -67,7 +67,8 @@ class RealifiedFunctional(Functional): - r"""View of a complex-domain functional over stacked real coordinates. + 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 @@ -213,7 +214,8 @@ def tree_unflatten(cls, aux: Any, children: Any) -> Self: def realify(F: Functional) -> Functional: - """Return ``F`` unchanged on a real domain, realified on a complex one. + """ + 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 diff --git a/spacecore/functional/tools/_entropy.py b/spacecore/functional/tools/_entropy.py index a1cf2d1..04dd6d5 100644 --- a/spacecore/functional/tools/_entropy.py +++ b/spacecore/functional/tools/_entropy.py @@ -45,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 -------- @@ -115,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 -------- diff --git a/spacecore/functional/tools/_huber.py b/spacecore/functional/tools/_huber.py index 55ddebf..211e96b 100644 --- a/spacecore/functional/tools/_huber.py +++ b/spacecore/functional/tools/_huber.py @@ -28,6 +28,18 @@ class HuberFunctional(_CoordinateFunctional[Domain]): when transcribing a prox or a smoothing constant from either source — the two differ by exactly one factor of ``delta``. + Parameters + ---------- + dom : Space + Domain space ``X``. + delta : float + 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", @@ -38,15 +50,6 @@ class HuberFunctional(_CoordinateFunctional[Domain]): Example 6.62 (its smoothness: gradient ``1/mu``-Lipschitz) and Example 6.66 (prox of the Huber function). - Parameters - ---------- - dom : Space - Domain space ``X``. - delta : float - Transition threshold; must be finite and ``> 0``. - ctx : Context, str, or None, optional - Backend context specification. Default is resolved from ``dom``. - Examples -------- >>> import numpy as np diff --git a/spacecore/functional/tools/_norms.py b/spacecore/functional/tools/_norms.py index 9758d6e..9fde0d0 100644 --- a/spacecore/functional/tools/_norms.py +++ b/spacecore/functional/tools/_norms.py @@ -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 -------- @@ -90,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 ----- @@ -160,6 +166,9 @@ 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 ------- diff --git a/spacecore/functional/tools/_spectral.py b/spacecore/functional/tools/_spectral.py index a832943..c148a2d 100644 --- a/spacecore/functional/tools/_spectral.py +++ b/spacecore/functional/tools/_spectral.py @@ -48,7 +48,8 @@ def eigenvalue_space(dom: Any, check_level: CheckLevel | bool | None = None) -> Any: - r"""Return the real coordinate space the spectrum of ``dom`` lives in. + 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 @@ -83,7 +84,8 @@ def eigenvalue_space(dom: Any, check_level: CheckLevel | bool | None = None) -> class SpectralFunctional(_CoordinateFunctional[Domain]): - r"""Lift a **symmetric** coordinate functional onto a Jordan spectrum. + r""" + 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 @@ -209,7 +211,8 @@ def spectralize( ctx: Context | str | None = None, check_level: CheckLevel | bool | None = None, ) -> "SpectralFunctional[Domain]": - r"""Build the spectral counterpart of a coordinate functional. + 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 diff --git a/spacecore/linop/_algebra.py b/spacecore/linop/_algebra.py index 3c69689..ce8f33c 100644 --- a/spacecore/linop/_algebra.py +++ b/spacecore/linop/_algebra.py @@ -907,6 +907,11 @@ 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 @@ -923,11 +928,6 @@ class MatrixFreeLinOp(LinOp[Domain, Codomain]): already carries the geometry and silences the advisory -- the distinction drawn is between "never considered" and "considered and declared". - 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`. Returns ------- diff --git a/spacecore/linop/_base.py b/spacecore/linop/_base.py index d0b739e..5fa30cc 100644 --- a/spacecore/linop/_base.py +++ b/spacecore/linop/_base.py @@ -32,25 +32,12 @@ class LinOp(PyTreeNode, 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**, 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 + 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. - 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``. - Parameters ---------- dom : Space @@ -75,6 +62,19 @@ class LinOp(PyTreeNode, 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`. 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/_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_vector.py b/spacecore/space/concrete/_dense_vector.py index c1d8dd8..4f76207 100644 --- a/spacecore/space/concrete/_dense_vector.py +++ b/spacecore/space/concrete/_dense_vector.py @@ -117,13 +117,13 @@ class ElementwiseJordanSpace(PyTreeNode, JordanAlgebraSpace, DenseCoordinateSpac geometry : InnerProduct or None, optional Inner-product geometry. If omitted, Euclidean coordinate geometry is used. - inner_product : InnerProduct or None, optional - Alias for ``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``. """ def __new__( @@ -237,13 +237,13 @@ class EuclideanElementwiseJordanSpace(ElementwiseJordanSpace, EuclideanJordanAlg geometry : InnerProduct or None, optional Inner-product geometry. This class is selected only for real contexts with Euclidean coordinate geometry. - inner_product : InnerProduct or None, optional - Alias for ``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``. """ def __init__( From 89f91bdc15470470dfbb4b6fafa1b0842307029f Mon Sep 17 00:00:00 2001 From: Pavlo Pelikh Date: Sat, 12 Sep 2026 03:09:25 -0400 Subject: [PATCH 6/6] Prepare the 0.4.3 release The version, the changelog heading and the public-API version assertion were already at 0.4.3, but an `[Unreleased]` section had since accumulated on top of the dated 0.4.3 section. Fold it in, redate to the actual release day, and keep the empty `[Unreleased]` placeholder the file has always carried. The docs still described the API this release removes, which would have shipped examples that raise: * README and the checking-policy design note both built contexts as `Context(ops, dtype=..., check_level=...)`. The level is now a property of the bound object, so both are rewritten around per-object levels and the ambient `set_check_level` / `use_check_level` helpers. The corrected README example was run; it prints the output the page claims. * `api/context.rst` pointed autodoc at `spacecore.backend.Context`, which no longer imports, and still described the context as carrying validation policy. It now also lists the scoping helpers. * README advertised `SpectralLpNormFunctional`, removed here in favour of `spectralize` / `SpectralFunctional`. Release notes had no 0.4.3 entry, and no 0.4.2 entry either -- the page jumped from 0.4.3's predecessor straight to 0.4.1 -- so both are written up. Sphinx now builds with no warnings, matching master. Co-Authored-By: Claude Opus 5 --- CHANGELOG.md | 182 ++++++++++++------------- README.md | 31 +++-- docs/source/api/context.rst | 32 +++-- docs/source/design/checking_policy.rst | 52 ++++--- docs/source/release_notes.rst | 91 +++++++++++++ 5 files changed, 253 insertions(+), 135 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 3ab458d..84ced0a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,8 +7,68 @@ 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 @@ -60,6 +120,26 @@ and the project adheres to [Semantic Versioning](https://semver.org/). ### 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 @@ -90,6 +170,14 @@ and the project adheres to [Semantic Versioning](https://semver.org/). ### 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 @@ -115,100 +203,6 @@ and the project adheres to [Semantic Versioning](https://semver.org/). so canonicalization is skipped and the expression tree stays unfolded but correct) while letting a genuinely broken `__eq__` propagate. -## [0.4.3] — 2026-07-25 - -### 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`. - -### 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. - -### 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.** - ## [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/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/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 -------------