File size: 28,896 Bytes
ecc81b3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 | """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,
)
|