Download src/torch_dimensions/testing.py from Celsia/torch-dimensions: direct link, hf CLI and curl.
- Browser
- Download file 28.9 kB
-
https://huggingface.co/Celsia/torch-dimensions/resolve/main/src/torch_dimensions/testing.py
- Command line
-
hf download hf://Celsia/torch-dimensions/src/torch_dimensions/testing.py
-
curl -L -o testing.py https://huggingface.co/Celsia/torch-dimensions/resolve/main/src/torch_dimensions/testing.py
28.9 kB
| """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] | |
| 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 ''}" | |
| class Report: | |
| results: list[Result] = field(default_factory=list) | |
| def failed(self) -> list[Result]: | |
| return [r for r in self.results if r.status == "fail"] | |
| 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 | |
| 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 | |
| def linear(self) -> bool: | |
| return max(self.additivity, self.homogeneity) < self.tol | |
| def affine(self) -> bool: | |
| """Linear once its constant response is subtracted — a bias, in short.""" | |
| return self.linear and self.zero_response > self.tol | |
| def time_invariant(self) -> bool: | |
| return self.shift_equivariance < self.tol | |
| 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, | |
| ) | |