torch-dimensions β Design
Status: Phases 0β11 built. Lattice (with sub-lattices), ScanPlan (with algebra and coverage reporting), both composition families, the conformance suite, the data layer, td.build/save/load (torch or safetensors), and the portable mixers β LSTM/GRU, S4/S4D/Mamba cross-validated against the upstream reference kernels, and Transformer (attention as the swept mixer). Reproductions are in RESULTS.md, measured costs in BENCHMARKS.md, and the viewer ships inside the wheel (td.viz.show). Verified on CPU and Apple Silicon; CUDA has never been executed β see docs/cuda-checklist.md. Remaining: fused-kernel fast paths and autoregressive stepping. See PLAN.md.
Positioning: composition layer. We own the N-D structure, the registry, the config surface, and the nn.Module contract. The 1-D kernels are ours too now, in portable pure torch β the plan was to depend on mamba-ssm / flash-linear-attention / state-spaces/s4, and writing derivative implementations instead is what makes the library work on a laptop at all. Those packages return as optional fast paths, held to agreeing with the portable reference.
1. The claim
Every model in scope is the same object:
a 1-D mixer, plus a plan for sweeping it over an N-D lattice.
| Model | 1-D mixer | Sweep strategy |
|---|---|---|
| Mamba-ND | Mamba-2 / Mamba-3 selective scan | sequential, one axis per layer, alternating direction |
| MDRNN / Grid-LSTM | nn.LSTM / nn.GRU |
sequential, one axis per layer |
| RNN + axial attention | nn.LSTM / nn.GRU |
hybrid β kernel across the lattice, mixer along time |
| S4ND | S4 FFT conv | separable β per-axis kernel, outer product |
| Axial Transformer | attention | td.Transformer (attention sweeps each axis) or per-axis kernels contracted in turn |
| Factorized axial attention | factorized axial cross-attn | per-axis kernel, Kronecker contraction |
Nobody has written this down as one abstraction. Every repo above hardcodes its own axis bookkeeping. That is the entire product: N-D RNNs, N-D transformers, and N-D SSMs fall out of one mechanism, which is why the user-facing API can be as small as S4(dim=2, layers=12).
The mechanism splits into composition strategies, selected per model by nd_method:
AxialScan β sequential. Permute one lattice axis to the sequence position, fold every other axis into batch, run the 1-D mixer, permute back. Residual + pre-norm per layer. This is what makes ND tractable: we never write an N-D kernel. Mamba-ND's real insight is that alternating 1-D scans over permuted axis orderings recover N-D context, so the existing 1-D CUDA kernel is reused unchanged.
AxialKernel β contraction. Build one kernel A_ax β R^{S_ax Γ S_ax} per axis, contract them into the value tensor one axis at a time. On a dense lattice the joint operator is exactly the Kronecker product A_0 β A_1 β β¦ β A_{n-1}, so cost is quadratic in axial size, not in prod(shape).
Hybrid β different operators on different axes. A kernel-family operator mixes across the lattice at each timestep; the 1-D mixer then runs along time. This is the shape of most real forecasting models over a categorical lattice, and it is why LSTM(nd_method=td.cafa) is meaningful: CaFA never consumes the LSTM, it handles the axes the LSTM does not.
Everything else in the library is a mixer, a plan, or a readout.
2. Public API
Three levels, each a strict superset of the one above.
Level 1 β drop-in modules
import torch
import torch_dimensions as td
lattice = td.Lattice(shape=(32, 64, 64), names=("depth", "height", "width"))
model = td.MambaND(d_model=128, n_layers=12, lattice=lattice, time=True)
y = model(x) # (B, T, *shape, d_model)
loss = y.pow(2).mean()
loss.backward() # plain autograd; nothing custom to call
Autograd is free. Composed torch ops give backward automatically; the upstream kernels already ship their own autograd.Function. There is no model.backwards() β you call .backward() on the loss, as with any torch model.
LSTM and Transformer take the same constructor shape. That is the point.
There is no LSTMND. td.LSTM(d_model, n_layers) with no lattice is an ordinary sequence model; adding lattice= makes the same class N-dimensional, because a lattice with no spatial axes has an identity permutation and the 1-D case is the N-D case with nothing to fold. How the extra axes are handled is nd_method's business. Strategies are plain functions exported at top level β td.axial_scan, later td.axial_attention and td.cafa β and a user's own function sits on exactly the same footing. Names are accepted too, but only because YAML cannot hold a callable.
Level 2 β mixer + plan
import torch_dimensions as td
plan = td.ScanPlan.cyclic(axes=("time", "depth", "height"), n_layers=12, bidirectional=True)
block = td.AxialScan(
mixer=partial(td.mixers.Mamba2Mixer, d_model=128, d_state=64), # one per layer
plan=plan,
lattice=lattice,
d_model=128,
)
ScanPlan is data, not control flow β a list of (axis, reverse) steps that is printable, serializable, diffable, and unit-testable independent of any mixer. In every existing ND implementation this schedule is inline list comprehensions welded to the module, which is why none of them can be inspected or swapped without editing the model. Constructors: .cyclic(), .paired(), .from_list(). .paired() is the schedule the official Mamba-ND implementation uses.
Bidirectionality is per-axis, not a global flag. Forward-only along time is correct β that is causality β while forward-only along a spatial or categorical axis is just lost receptive field, so bidirectional= takes a collection of axes. It also costs layers: covering k axes both ways needs roughly 2k layers, and below that budget the request is silently downgraded. ScanPlan warns instead.
Space-filling traversals (Hilbert, Morton) were in an earlier draft of this list and do not belong: a step is (axis, reverse), and a Hilbert curve is not axis-aligned. Supporting it needs a second step kind carrying a full permutation of cell indices, which changes what AxialScan folds and is a different feature wearing this one's clothes. Deferred rather than stubbed.
Any nn.Module mapping (M, A, H) -> (M, A, H) is a valid mixer. That is the extension point β a user drops in a new SSM from a paper published next month without touching the library. Pass a factory for per-layer weights, or a built module to share one set across layers.
Level 3 β config
model:
kind: mamba_nd
d_model: 128
n_layers: 12
d_state: 64
lattice: {shape: [32, 64, 64], names: [depth, height, width], time: true}
plan: {type: cyclic, bidirectional: true}
model = td.build(cfg) # or td.build_from_yaml(path)
One dataclass schema per block, validated at construction, with the error naming the offending key. No stringly-typed variant names.
3. Lattice
Lattice is the object that carries all N-D structure, so that no block ever hardcodes a rank or an axis meaning.
@dataclass
class Lattice:
shape: tuple[int, ...]
names: tuple[str, ...] | None = None
valid: torch.Tensor | None = None # bool, `shape` β sparse support
time: bool = False # prepend a scanned, non-lattice time axis
It owns: axis-name β position resolution, flat indices for scatter/gather, the broadcast validity mask, per-axis valid counts for masked pooling, and permutation/inverse-permutation generation.
Two failure modes it exists to prevent, both endemic to existing ND implementations:
- Conflating an encoding strategy with a rank. Names of the form
<encoder>_4dmean "this encoder" and "four axes" at once, which forces string-rewriting hacks the moment a new encoder appears. Here they are orthogonal arguments:kind=,lattice=,encoder=. - Rank-locked contraction tables. Hand-written einsum strings keyed by axis index fix the rank at three (or two, or four) forever.
Latticegenerates permutations for arbitrary N instead. Preferpermute β reshape β matmul β permute backover generated einsum strings: same FLOPs, no per-call path planning, andtorch.compile-friendly.
Device-dependent limits β chunk sizes for the folded batch dimension, kernel grid bounds β are library-owned and auto-tuned, never user-facing constants.
4. Sparse lattices β the differentiator
S4ND, Mamba-ND, and factorized axial attention all assume a dense grid. Real N-D data frequently is not: not every (sensor, channel, band) triple is instrumented, not every (patient, visit, assay) exists, meshes are irregular, modalities go missing.
Support is mechanical once Lattice.valid exists, and different per family:
- Scan family: scatter to dense, zero invalid cells after each layer, gather back.
- Kernel family: masked-mean pooling over valid cells only, then per-line renormalization by the softmax mass that landed on valid keys, so every output stays a convex combination of valid values.
That per-line rescale is the one departure from a strict Kronecker product, and it costs O(N Β· S_ax) elementwise work rather than O(N Β· S_ax) score memory. This is what keeps the factorized path alive at rank 4, where materializing one attention matrix per lattice line runs out of memory.
Making this a first-class feature β automatic for every mixer, tested once β is the strongest reason for this library to exist rather than for users to keep vendoring kernels and re-deriving the masking themselves.
5. Layout
torch_dimensions/
lattice.py Lattice, permutation + scatter/gather + mask machinery
plan.py ScanPlan and its constructors
compose/
scan.py AxialScan
kernel.py AxialKernel, Kronecker contraction, sparse renormalization
mixers/
ssm.py Mamba2, Mamba3, S4, S5 (thin adapters over upstream kernels)
rnn.py LSTM, GRU (adapters over torch.nn)
attn.py self-attn, cross-attn, factorized axial kernel
models/ MambaND, S4ND, Transformer, LSTM, GRU
data/ LatticeSource protocol, long-format tables, windowing, collate
registry.py register / build / list
config.py dataclass schemas + YAML loader + save/load
testing.py shared conformance suite
mixers/ are adapters, not reimplementations β that is the composition-layer decision made concrete. Optional deps are imported defensively: a missing mamba-ssm unregisters mamba_nd and leaves the rest importable. A missing real implementation must fail loudly rather than silently register a stub, so that a benchmark can never accidentally report an LSTM's numbers under an SSM's name.
6. Conformance suite
One parametrized suite every registered block must pass. This is what keeps an N-model Γ N-dim matrix from rotting.
- Shape β
(B, *shape, H)in, same out, for ranks 1β4. - Gradient β
gradcheckin float64 on a tiny instance; every parameter receives non-Nonegrad. - Equivalence β a single layer on a rank-1 lattice must match the underlying 1-D module bit-for-bit. Catches permutation bugs immediately. Stacks are checked to floating-point tolerance instead: the fold reshapes, which requires contiguity, and torch's RNN kernels are not bit-identical across memory layouts, so from layer two onward a one-ULP drift is expected and is not a bug.
- Kronecker identity β for
AxialKernelon a dense lattice, sequential contraction must equal the explicitβoperator to numerical tolerance. - Mask invariance β invalid cells must not influence valid outputs. Perturb invalid cells, assert valid outputs are bitwise unchanged. This is the sparse-lattice guarantee, and the test most likely to catch a real bug.
- Permutation covariance β permuting lattice axes and the plan together permutes the output correspondingly.
- Compile β
torch.compileproduces matching numerics.
7. v0.1 scope
In: Lattice (dense + sparse), ScanPlan, AxialScan, AxialKernel, mixers for LSTM/GRU/attention (pure torch, no optional deps) and Mamba-2/S4 (adapters), models LSTM/GRU/Transformer/MambaND/S4ND, the data/ construction layer (Β§8), registry, config, save/load, conformance suite. Ranks 1β6 (5 and 6 verified after the fact; nothing counts axes).
Deferred, deliberately:
- Stateful stepping. Autoregressive decode needs per-axis recurrent state caching, and along a non-time axis "state" is not well-defined β you would be caching a cross-section, not a prefix. This is a genuine open research question, not an implementation gap. v0.1 is forward-only; the README must say so.
- Custom kernels. The composition-layer decision. Revisit only when profiling shows permute/contiguous dominating, at which point the fix is a fused permute+scan kernel, not a reimplemented SSM.
- Ranks β₯ 5. Nothing forbids them; nothing tests them.
- Training loops, optimizers, losses, schedules. Out of scope permanently.
data/builds lattices; it never trains them (Β§8).
Known risk. The permutes are not free: AxialScan needs two .contiguous() calls per layer, so a 12-layer ND model does ~24 full-tensor copies. Benchmark this before advertising performance numbers β it may well dominate the mixer at small d_model.
8. Data β construction, not training
One distinction scopes this entire subpackage:
Building a lattice from data is lattice construction, which this library already owns. Running a training loop is not.
The gap is concrete. A user holding long-format rows β (coordβ, coordβ, β¦, t, featuresβ¦) β must otherwise hand-write the coordinate-to-index mapping, infer the shape, build the validity mask, and scatter into (B, T, *shape, H). That is exactly the code that silently produces a mis-shuffled lattice: it trains, it converges, and the numbers are quietly wrong. It is the same bug class Β§3 exists to eliminate, sitting one layer up.
Four pieces, each usable alone:
Lattice.from_coordsβ infer shape, validity mask, and categorical vocabularies from observed coordinates.LatticeWindowβ windowing over time. Pure index arithmetic, no I/O.LatticeSourceβ a protocol, not a base class. This is where customization comes from: memory-mapped arrays, zarr, HDF5, or a database all batch correctly if they satisfy it. Ship the protocol plus two reference implementations.collate_latticeβ stacks windows, keeping theLatticeout of the batch since it is static metadata rather than per-sample data.
Customization comes from composition, never from a god-class with forty constructor arguments.
Deliberately absent: dataset downloads, normalization policy (a hook, shipping nothing), augmentation, splitting strategy, trainers, Lightning integration.
9. Open questions
- Axis order when
time=True. Is time axis 0 of the lattice, or a separate leading dim? Existing implementations special-case it into the scan schedule. Cleaner: time is a normal named axis, and causality is a property of the mixer, not the axis. - Does
AxialKernelneedAxialScan's per-layer alternation, or is one pass over all axes sufficient? Factorized attention does one pass; Mamba-ND alternates across 12 layers. Probably a plan-level choice rather than two mechanisms. Name.Resolved:torch-dimensions, importing astorch_dimensions, aliasedtd. See PLAN.md.