"""The shared conformance suite. Public API rather than test-directory scaffolding, because the extension point *is* the product: anyone writing a mixer or an ``nd_method`` should be able to run exactly the checks the library runs on itself. import torch_dimensions as td td.testing.check_block(lambda lat, d: td.LSTM(d, 3, lat)) The checks are ordered so the cheapest and most diagnostic run first. An axis bug in an N-D model presents as "the model trains badly"; these turn it into a specific failing assertion instead. """ from __future__ import annotations import inspect from collections.abc import Callable, Sequence from dataclasses import dataclass, field from functools import reduce from typing import NamedTuple import torch import torch.nn as nn from torch_dimensions.lattice import Lattice from torch_dimensions.plan import ScanPlan __all__ = [ "LTIReport", "Recorder", "Report", "Result", "check_block", "check_data_source", "check_lti", "check_trainable", ] Factory = Callable[..., nn.Module] @dataclass class Result: name: str status: str # "pass" | "fail" | "skip" detail: str = "" def __str__(self) -> str: mark = {"pass": "ok", "fail": "FAIL", "skip": "skip"}[self.status] return f"[{mark:>4}] {self.name}{f' — {self.detail}' if self.detail else ''}" @dataclass class Report: results: list[Result] = field(default_factory=list) @property def failed(self) -> list[Result]: return [r for r in self.results if r.status == "fail"] @property def skipped(self) -> list[Result]: return [r for r in self.results if r.status == "skip"] def __bool__(self) -> bool: return not self.failed def __str__(self) -> str: return "\n".join(str(r) for r in self.results) def _lattice(rank: int, *, sparse: bool = False, time: bool = False, seed: int = 0) -> Lattice: shape = tuple(range(2, 2 + rank)) valid = None if sparse: g = torch.Generator().manual_seed(seed) valid = torch.rand(shape, generator=g) > 0.4 valid.reshape(-1)[0] = True valid.reshape(-1)[-1] = True return Lattice(shape=shape, valid=valid, time=time) def _input(lat: Lattice, d_model: int, batch: int, seq: int, seed: int) -> torch.Tensor: g = torch.Generator().manual_seed(seed) lead = (batch, seq) if lat.time else (batch,) return torch.randn(*lead, *lat.shape, d_model, generator=g, dtype=torch.float64) def _build(factory: Factory, lat: Lattice, d_model: int, seed: int, **kw) -> nn.Module: torch.manual_seed(seed) block = factory(lat, d_model, **kw) return block.double().eval() def _accepts_plan(factory: Factory) -> bool: try: return "plan" in inspect.signature(factory).parameters except (TypeError, ValueError): # builtins, C callables return False def check_block( factory: Factory, *, d_model: int = 4, ranks: Sequence[int] = (1, 2, 3), sparse: bool = True, time: bool = False, batch: int = 2, seq: int = 3, reference: Callable[[nn.Module, torch.Tensor], torch.Tensor] | None = None, kernels: Callable[[nn.Module, torch.Tensor], tuple[Sequence[torch.Tensor], torch.Tensor]] | None = None, check_compile: bool = False, seed: int = 0, raise_on_failure: bool = True, ) -> Report: """Run the conformance checks against a block factory. Args: factory: ``(lattice, d_model) -> nn.Module``. Accepting a keyword ``plan`` additionally enables the permutation-covariance check, which needs to hold the sweep order fixed while the lattice's axis *storage* order changes. ranks: lattice ranks to exercise. Rank 1 is the one that catches permutation bugs fastest. sparse: also build on a lattice with absent cells and verify that their values cannot influence any output. reference: ``(block, x) -> expected`` for the rank-1 equivalence check. Omit and that check is skipped rather than silently passed. kernels: ``(block, x) -> (per_axis_kernels, output)`` for the Kronecker check — run the block's axis-by-axis contraction and hand back the matrices it actually used along with what it produced. The check then builds the joint operator with ``torch.kron`` and compares. A factorized block that cannot produce this is a block whose central claim is untested; it used to be an unconditional skip. check_compile: compare ``torch.compile`` numerics against eager. Off by default because it is slow, not because it is unimportant. raise_on_failure: raise ``AssertionError`` with the full report when any check fails. The report is returned either way. Returns: A :class:`Report`, falsy if anything failed. """ rep = Report() def record(name, fn): try: detail = fn() except _Skip as s: rep.results.append(Result(name, "skip", str(s))) except Exception as e: # noqa: BLE001 — a failing check is data, not a crash rep.results.append(Result(name, "fail", f"{type(e).__name__}: {e}")) else: rep.results.append(Result(name, "pass", detail or "")) # 1. shape --------------------------------------------------------------- def _shapes(): for r in ranks: lat = _lattice(r, time=time) x = _input(lat, d_model, batch, seq, seed) out = _build(factory, lat, d_model, seed)(x) if out.shape != x.shape: raise AssertionError(f"rank {r}: got {tuple(out.shape)}, expected {tuple(x.shape)}") return f"ranks {tuple(ranks)}" record("shape is preserved", _shapes) # 2. gradients ----------------------------------------------------------- def _grads(): # A rank the caller actually asked for. Hardcoding rank 2 "for speed" # gradchecked a lattice the factory was never claimed to support — # ranks=(3, 4) would build and differentiate a rank-2 block behind the # caller's back. lat = _lattice(2 if 2 in ranks else min(ranks), time=time) # The caller's width, not a narrower one chosen here for speed. A # hardcoded `d_model=2` gradchecked a block the factory was never # claimed to support, and factories with a width constraint — an # attention mixer whose head count must divide `d_model` — failed a # check about *gradients* with an error about heads. Same shape as the # rank bug (DEBUG.md #16), one argument over. block = _build(factory, lat, d_model, seed) x = _input(lat, d_model, 1, seq, seed).requires_grad_(True) block(x).pow(2).mean().backward() dead = [n for n, p in block.named_parameters() if p.grad is None] if dead: raise AssertionError(f"parameters received no gradient: {dead}") if not torch.autograd.gradcheck(block, (x.detach().requires_grad_(True),), fast_mode=True): raise AssertionError("gradcheck failed") return f"{sum(1 for _ in block.parameters())} tensors, gradcheck clean" record("gradients flow and gradcheck passes", _grads) # 3. rank-1 equivalence -------------------------------------------------- def _equivalence(): if reference is None: raise _Skip("no `reference` given") if 1 not in ranks: raise _Skip("rank 1 not in `ranks`; the equivalence claim is a rank-1 claim") lat = _lattice(1, time=time) block = _build(factory, lat, d_model, seed) x = _input(lat, d_model, batch, seq, seed) got, want = block(x), reference(block, x) if not torch.equal(got, want): raise AssertionError( f"rank-1 output differs from the reference by " f"{(got - want).abs().max().item():.3e} (must be exact)" ) return "bitwise identical" record("rank-1 equals the bare 1-D module", _equivalence) # 4. Kronecker identity -------------------------------------------------- def _kronecker(): if kernels is None: raise _Skip("no `kernels` adapter given; kernel-family blocks should supply one") r = max(r for r in ranks if r >= 2) if any(r >= 2 for r in ranks) else 0 if not r: raise _Skip("needs rank >= 2; a one-axis Kronecker product is just the kernel") lat = _lattice(r, time=time) block = _build(factory, lat, d_model, seed) # Batch 1: the factorized families build one kernel per (batch, step), # and a single joint operator can only be compared against a single # batch element's kernels. x = _input(lat, d_model, 1, seq, seed) mats, out = kernels(block, x) if len(mats) < 2: raise _Skip(f"adapter returned {len(mats)} kernels; needs >= 2 to form a product") joint = reduce(torch.kron, [m.to(torch.float64) for m in mats]) flat = x.reshape(*x.shape[: -(lat.rank + 1)], -1, x.shape[-1]).to(torch.float64) want = (joint @ flat).reshape(out.shape) diff = (out.to(torch.float64) - want).abs().max().item() if diff > 1e-9: raise AssertionError( f"contracting axis by axis differs from the joint Kronecker operator by " f"{diff:.3e}; the factorization is not the product it claims to be" ) return f"rank {r}, {len(mats)} axes, max diff {diff:.1e}" record("Kronecker identity (kernel family)", _kronecker) # 5. mask invariance ----------------------------------------------------- def _mask(): if not sparse: raise _Skip("sparse=False") checked = 0 for r in ranks: if r < 1: continue lat = _lattice(r, sparse=True, time=time, seed=seed) if lat.n_valid == lat.n_cells: continue block = _build(factory, lat, d_model, seed) x = _input(lat, d_model, batch, seq, seed) noise = torch.randn_like(x) * 1e3 * (~lat.mask()).to(x.dtype) if not torch.equal(block(x), block(x + noise)): raise AssertionError( f"rank {r}: perturbing absent cells changed the output; they must be " "zeroed before the mixer sees them" ) checked += 1 if not checked: raise _Skip("no sparse lattice was generated") return f"{checked} sparse lattices" record("absent cells cannot influence the output", _mask) # 6. permutation covariance ---------------------------------------------- def _covariance(): if not _accepts_plan(factory): raise _Skip("factory does not accept `plan`") r = max(ranks) if r < 2: raise _Skip("needs rank >= 2") names = tuple(f"ax{i}" for i in range(r)) shape = tuple(range(2, 2 + r)) order = tuple(range(1, r)) + (0,) # rotate the storage order plan = ScanPlan.from_list(list(names)) a = Lattice(shape=shape, names=names, time=time) b = Lattice( shape=tuple(shape[i] for i in order), names=tuple(names[i] for i in order), time=time, ) x = _input(a, d_model, batch, seq, seed) # move each lattice dim of x into b's storage order lead = 2 if time else 1 perm = (*range(lead), *(lead + i for i in order), x.ndim - 1) out_a = _build(factory, a, d_model, seed, plan=plan)(x) out_b = _build(factory, b, d_model, seed, plan=plan)(x.permute(*perm)) if not torch.allclose(out_b, out_a.permute(*perm), rtol=0, atol=1e-12): raise AssertionError( "output depends on the order axes happen to be stored in, not just on " "the sweep order" ) return f"rank {r}, storage order rotated" record("output is covariant with axis storage order", _covariance) # 7. compile ------------------------------------------------------------- def _compile(): if not check_compile: raise _Skip("check_compile=False") lat = _lattice(max(ranks), time=time) block = _build(factory, lat, d_model, seed) x = _input(lat, d_model, batch, seq, seed) eager = block(x) got = torch.compile(block)(x) if not torch.allclose(got, eager, rtol=1e-9, atol=1e-9): raise AssertionError(f"max diff {(got - eager).abs().max().item():.3e}") return "matches eager" record("torch.compile matches eager", _compile) if raise_on_failure and rep.failed: raise AssertionError("conformance check failed:\n" + str(rep)) return rep class _Skip(Exception): """Raised inside a check to record it as skipped rather than passed. Deliberately not silent: a skipped check appears in the report, so "we never ran that one" can never read as "that one passed". """ def check_trainable( factory: Factory, *, d_model: int = 16, steps: int = 200, lr: float = 1e-2, batch: int = 8, seq: int = 5, min_ratio: float = 3.0, seed: int = 0, raise_on_failure: bool = True, ) -> dict[str, float]: """Fit a small task that genuinely needs N-D mixing, and check it learns. Separate from :func:`check_block` on purpose. That one asks *is this correct* — deterministic, exact, fast. This one asks *does this learn*, which is stochastic, slower, and catches a different failure entirely: a block can have flawless gradients, pass ``gradcheck``, and still never converge because of initialization, masking that kills the signal, or activations that blow up. "No trainer in the library" must not quietly become "nobody ever checked that it trains". The task is a cumulative sum along the **last lattice axis**, so a model that never sweeps that axis cannot solve it — the check has a meaningful negative, not just a number that goes down. Fresh data is drawn every step and the reported score is on a held-out batch. With a fixed training set this test is worthless: a model with enough capacity memorizes eight examples without doing any axial mixing at all, and every plan passes. Returns a dict of ``initial``, ``final``, ``held_out`` and ``ratio``. """ lat = Lattice(shape=(3, 4), names=("h", "w"), time=True) torch.manual_seed(seed) block = factory(lat, d_model) head = nn.Linear(d_model, 1) opt = torch.optim.Adam([*block.parameters(), *head.parameters()], lr=lr) def draw(g): x = torch.randn(batch, seq, *lat.shape, d_model, generator=g) return x, x[..., :1].cumsum(dim=lat.tensor_dim("w")) g = torch.Generator().manual_seed(seed) initial = final = 0.0 for i in range(steps): x, y = draw(g) loss = (head(block(x)) - y).pow(2).mean() if i == 0: initial = loss.item() final = loss.item() opt.zero_grad() loss.backward() opt.step() block.eval() with torch.no_grad(): x, y = draw(torch.Generator().manual_seed(seed + 9973)) held_out = (head(block(x)) - y).pow(2).mean().item() ratio = initial / max(held_out, 1e-12) result = { "initial": initial, "final": final, "held_out": held_out, "ratio": ratio, } if raise_on_failure and ratio < min_ratio: raise AssertionError( f"block did not learn: held-out loss {held_out:.4f} vs initial {initial:.4f} " f"({ratio:.1f}x, needed {min_ratio}x). Gradients can be correct and the " "block still not converge." ) return result class Recorder(nn.Module): """A mixer that computes nothing and remembers everything. The first question every integration bug asks is *which axis did layer 3 actually sweep, and in which direction* — and until now the only way to answer it was a private helper in this project's own test files. It is a mixer like any other, so it drops into any model in place of the real one:: model = td.LSTM(8, 6, lattice, mixer=td.testing.Recorder) model(x) print(model.nd.mixers[0].calls) # [Call(shape=(24, 5, 8), lines=24, length=5)] It is the identity function, so the model still runs and still has the right output shape; only the mixing is gone. What a call records is what a mixer is actually told: the folded shape. A mixer never learns its axis name — that is the design — so the axis is inferred by the caller from ``length`` against the lattice, which is exactly the reasoning a person does by hand when a sweep goes wrong. """ class Call(NamedTuple): shape: tuple[int, ...] lines: int """The folded batch: batch times every axis except the swept one.""" length: int """The swept axis's length — what identifies the axis on most lattices.""" def __init__(self, d_model: int, **_: object) -> None: super().__init__() self.d_model = d_model # A parameter so that optimizers and the conformance suite's # "everything gets a gradient" check have something to hold; it is # multiplied by one, so the module stays the identity. self.scale = nn.Parameter(torch.ones(())) self.calls: list[Recorder.Call] = [] def forward(self, x: torch.Tensor) -> torch.Tensor: self.calls.append(self.Call(tuple(x.shape), int(x.shape[0]), int(x.shape[1]))) return x * self.scale def reset(self) -> None: self.calls.clear() def extra_repr(self) -> str: return f"d_model={self.d_model}, {len(self.calls)} calls recorded" def check_data_source( source: object, *, n_probe: int = 3, raise_on_failure: bool = True, ) -> Report: """Check that a custom :class:`~torch_dimensions.data.LatticeSource` behaves. The source protocol is an extension point — a memmap, a zarr store, a database cursor — and extension points deserve the same treatment mixers got. A source that satisfies the *types* and gets the semantics wrong produces a model that trains on subtly misaligned data and never says so. td.testing.check_data_source(MyZarrSource(...)) What it checks: the declared lattice matches the shape actually returned; slices are consistent with each other (the concatenation of two adjacent slices is the slice that spans them); reads are repeatable; and the source survives being pickled, because ``DataLoader(num_workers>0)`` pickles it and a source holding an open file handle fails only in a worker process (DEBUG.md #9 — that failure mode *hung* rather than raised). """ rep = Report() def record(name, fn): try: detail = fn() except _Skip as s: rep.results.append(Result(name, "skip", str(s))) except Exception as e: # noqa: BLE001 rep.results.append(Result(name, "fail", f"{type(e).__name__}: {e}")) else: rep.results.append(Result(name, "pass", detail or "")) def _members(): missing = [m for m in ("lattice", "__len__", "__getitem__") if not hasattr(source, m)] if missing: raise AssertionError(f"missing {missing}; see td.data.LatticeSource") if len(source) < 1: # type: ignore[arg-type] raise AssertionError("source is empty; nothing can be checked against it") return f"{len(source)} timesteps" # type: ignore[arg-type] record("has the protocol's members", _members) def _shape(): lat: Lattice = source.lattice # type: ignore[attr-defined] chunk = source[0 : min(n_probe, len(source))] # type: ignore[index] if not isinstance(chunk, torch.Tensor): raise AssertionError(f"__getitem__ returned {type(chunk).__name__}, expected a Tensor") got = tuple(chunk.shape[1:-1]) if got != tuple(lat.shape): raise AssertionError( f"returns lattice dims {got} but declares shape {tuple(lat.shape)}; " "the lattice and the data disagree" ) return f"{tuple(chunk.shape)} for {min(n_probe, len(source))} steps" record("returned shape matches the declared lattice", _shape) def _slices(): n = len(source) # type: ignore[arg-type] if n < 2: raise _Skip("needs at least 2 timesteps") mid = max(1, n // 2) whole = source[0:n] # type: ignore[index] halves = torch.cat([source[0:mid], source[mid:n]]) # type: ignore[index] if not torch.equal(whole, halves): raise AssertionError( "reading in two slices differs from reading in one; windows will " "silently straddle the seam" ) return f"split at {mid} of {n}" record("adjacent slices concatenate to the whole", _slices) def _repeatable(): a = source[0 : min(n_probe, len(source))] # type: ignore[index] b = source[0 : min(n_probe, len(source))] # type: ignore[index] if not torch.equal(a, b): raise AssertionError("two identical reads returned different data") return "two reads agree" record("reads are repeatable", _repeatable) def _picklable(): import pickle try: revived = pickle.loads(pickle.dumps(source)) except Exception as e: # noqa: BLE001 raise AssertionError( f"cannot pickle ({type(e).__name__}: {e}) — DataLoader(num_workers>0) " "pickles the source, and an open file handle fails only in a worker" ) from e k = min(n_probe, len(source)) # type: ignore[arg-type] if not torch.equal(revived[0:k], source[0:k]): # type: ignore[index] raise AssertionError("the unpickled source returns different data") return "survives a worker process" record("pickles for DataLoader workers", _picklable) if raise_on_failure and rep.failed: raise AssertionError("data source check failed:\n" + str(rep)) return rep @dataclass class LTIReport: """What :func:`check_lti` measured. Numbers, not adjectives. Every field is a *relative* error — the deviation divided by the size of the output it deviates from — so the numbers are comparable across mixers with wildly different output scales. Around 1e-16 means the property holds to floating point; anything above ~1e-6 means it does not hold at all. """ name: str additivity: float homogeneity: float zero_response: float shift_equivariance: float tol: float = 1e-9 @property def linear(self) -> bool: return max(self.additivity, self.homogeneity) < self.tol @property def affine(self) -> bool: """Linear once its constant response is subtracted — a bias, in short.""" return self.linear and self.zero_response > self.tol @property def time_invariant(self) -> bool: return self.shift_equivariance < self.tol @property def verdict(self) -> str: if self.linear and self.time_invariant: return "LTI" + (" (affine)" if self.affine else "") if self.time_invariant: return "time-invariant, nonlinear" if self.linear: return "linear, not time-invariant" return "neither" def __str__(self) -> str: return ( f"{self.name}: {self.verdict}\n" f" additivity {self.additivity:.2e}\n" f" homogeneity {self.homogeneity:.2e}\n" f" shift equivariance {self.shift_equivariance:.2e}\n" f" response to zero {self.zero_response:.2e}" ) def _rel(diff: torch.Tensor, scale: torch.Tensor) -> float: """Deviation relative to the magnitude of what it deviates from.""" denom = scale.abs().max().item() return float(diff.abs().max().item() / max(denom, 1e-30)) def check_lti( mixer: nn.Module | Callable[[], nn.Module], *, d_model: int = 4, length: int = 24, batch: int = 2, shift: int = 3, guard: int | None = None, seed: int = 0, tol: float = 1e-9, ) -> LTIReport: """Measure whether a mixer is linear and time-invariant. This is not a pass/fail check and never raises — no mixer is *supposed* to be LTI. It is a classification, and the classification is what decides how a mixer behaves under N-D composition: - **LTI mixers commute across axes.** Sweeping ``h`` then ``w`` equals sweeping ``w`` then ``h``, so the sweep order carries no information and the whole stack collapses to one separable N-D operator. This is why a separable CNN is exactly an N-D convolution, and why S4ND can apply its axes simultaneously instead of in sequence. - **Non-LTI mixers do not.** Order and direction become architectural choices with real consequences, which is the entire reason ``ScanPlan`` exists and why Mamba-ND needed a schedule at all. **How time-invariance is tested, and why that way.** The input is shifted by zero-padding the front rather than by rolling it, because time invariance is a statement about a system *at rest*: feed it nothing, then feed it the signal later, and the same thing should come out later. Rolling would instead wrap a different prefix into place and test memory decay, which is a different question. A consequence worth knowing: a recurrent mixer whose gates have biases does not stay at rest under zero input, so it is not time-invariant even though it is perfectly causal. Args: mixer: a built module, or a zero-argument factory. Run in ``eval`` mode and float64 — dropout would make every measurement noise. shift: how far to delay the signal for the equivariance test. tol: relative error below which a property counts as holding. Returns: An :class:`LTIReport`. See LTI.md for the measured table across every mixer this library ships. """ torch.manual_seed(seed) block = (mixer if isinstance(mixer, nn.Module) else mixer()).double().eval() name = type(block).__name__ g = torch.Generator().manual_seed(seed) shape = (batch, length, d_model) x = torch.randn(*shape, generator=g, dtype=torch.float64) y = torch.randn(*shape, generator=g, dtype=torch.float64) zeros = torch.zeros(*shape, dtype=torch.float64) with torch.no_grad(): # A block may be affine rather than linear (any bias makes it so). # Subtracting its response to zero tests the linear part, and the # response itself is reported separately rather than hidden. f0 = block(zeros) fx, fy, fxy = block(x) - f0, block(y) - f0, block(x + y) - f0 f3x = block(3.0 * x) - f0 additivity = _rel(fxy - (fx + fy), fxy) homogeneity = _rel(f3x - 3.0 * fx, f3x) # Delay by zero-padding the front: the system starts at rest and the # signal arrives `shift` steps later. delayed = torch.cat([torch.zeros(batch, shift, d_model, dtype=torch.float64), x], dim=1)[ :, :length ] fd = block(delayed) - f0 # Measured away from both boundaries, because both ends lie about it. # # At the *start*: a stacked causal convolution with biases is not # actually at rest for its first few positions — its own left-padding # is zero while the interior has settled to the bias, so the response # to a zero input is not constant until the transient passes. Exactly # the same shape as an RNN's state ramp, and measuring inside it # reports a boundary convention as a property of the operator. # # At the *end*: delaying truncates the tail of the signal, which a # causal mixer never notices and a centred one does. band = max(shift, length // 4) if guard is None else guard got = fd[:, shift + band : length - band] want = fx[:, band : length - shift - band] if got.shape[1] < 1: raise ValueError( f"nothing left to compare: length={length}, shift={shift}, guard={band}. " "Lengthen the probe or shrink the guard." ) shift_err = _rel(got - want, want) return LTIReport( name=name, additivity=additivity, homogeneity=homogeneity, zero_response=_rel(f0, block(x)), shift_equivariance=shift_err, tol=tol, )