Adding a mixer
This is the page that matters most, because the extension point is the product. A mixer is a 1-D sequence model; everything N-dimensional β the lattice, the schedule, the absent cells, the permutations β is the library's job, and you never write any of it.
The worked example lives in
examples/custom_mixer.py and is executed by
tests/test_examples.py, so nothing on this page is a snippet that used to
work.
1. The contract
def __call__(self, x: torch.Tensor) -> torch.Tensor: # (M, A, H) -> (M, A, H)
That is all of it. A is the length of the axis being swept, H is the
feature width, and M is the batch times every other axis folded together.
Three things a mixer is deliberately never told:
- which axis it is sweeping. Rows, columns, time, or "commodity" all
arrive as the same
(M, A, H). This is why one implementation works at every rank. - what rank the lattice is, or which cells are absent. Absent cells are zeroed before you see them and after every layer, so a mixer never masks.
- which direction it is going. A backward sweep arrives already flipped.
Write it causally; the schedule gives you bidirectionality for free, and
setting
bidirectional=Trueon an inner RNN would double the width and leave the schedule nothing to control.
If your model can process a batch of sequences, it is already a mixer.
2. Write it
class EMAMixer(nn.Module):
"""y_t = a * y_{t-1} + (1 - a) * x_t, with a learned per channel."""
def __init__(self, d_model: int, init_halflife: float = 4.0) -> None:
super().__init__()
a0 = 0.5 ** (1.0 / init_halflife)
self.decay_logit = nn.Parameter(torch.full((d_model,), logit(a0)))
self.out = nn.Linear(d_model, d_model)
def forward(self, x): # (M, A, H)
a = torch.sigmoid(self.decay_logit)
state, out = torch.zeros_like(x[:, 0]), []
for t in range(x.shape[1]):
state = a * state + (1 - a) * x[:, t]
out.append(state)
return self.out(torch.stack(out, dim=1))
The decay is parameterized as a logit rather than clamped into (0, 1),
because clamping gives exactly zero gradient at the boundary the model most
wants to move along. That kind of detail is the mixer author's business; the
axis bookkeeping is not.
The one signature requirement: the first constructor argument is
d_model. The composition layer builds one mixer per layer by calling
Mixer(d_model, **mixer_kwargs).
3. Get a model class for free
from torch_dimensions.models.base import LatticeModel
class EMA(LatticeModel):
_mixer = EMAMixer
Two lines, and you have nd_method=, plan=, lattice=, d_input=,
dropout=, chunk=, .config, .save(), .to_spec() and viewer support β
the same surface td.LSTM and td.Mamba have, because they are written the
same way.
4. Run the conformance suite
This is the step to not skip. td.testing is public API precisely so that a
new mixer is held to the standard the built-ins are held to.
def factory(lattice, d_model, plan=None):
return EMA(d_model, len(lattice.axis_names), lattice, plan=plan)
report = td.testing.check_block(factory, reference=reference)
print(report)
[ ok] shape is preserved β ranks (1, 2, 3)
[ ok] gradients flow and gradcheck passes β 10 tensors, gradcheck clean
[ ok] rank-1 equals the bare 1-D module β bitwise identical
[skip] Kronecker identity (kernel family) β no `kernels` adapter given
[ ok] absent cells cannot influence the output β 2 sparse lattices
[ ok] output is covariant with axis storage order β rank 3, storage order rotated
[skip] torch.compile matches eager β check_compile=False
What each check is actually protecting you from:
| check | the bug it catches |
|---|---|
| shape | a mixer that changes width, or a fold that does not invert |
| gradients + gradcheck | a parameter that never learns; a wrong backward |
| rank-1 equivalence | the N-D machinery doing anything on a 1-D lattice |
| absent-cell inertia | values from cells that do not exist reaching an output |
| storage covariance | an output that depends on axis storage order, not sweep order |
Supply reference=. Without it that check is skipped, and the report
says so β a skip is recorded, never silently passed. The reference is what one
pre-norm residual layer around your bare mixer computes:
def reference(block, x):
return x + block.nd.mixers[0](block.nd.norms[0](x))
On a rank-1 lattice the whole apparatus must reduce to exactly that, bitwise. It is the fastest way to discover that a fold or permutation is subtly wrong.
Then check that it learns:
td.testing.check_trainable(factory, d_model=16, steps=120)
# {'initial': 2.64, 'final': 0.37, 'held_out': 0.32, 'ratio': 8.4}
check_trainable fits a task that genuinely requires axial mixing, on fresh
data every step, scored held out, with a negative control that must fail. A
learning test without a negative control measures capacity, not learning.
5. Register it (optional)
td.register_model("ema", EMA)
model = td.build(
{"kind": "ema", "d_model": 32, "n_layers": 8, "lattice": {"shape": [4, 5, 3], "time": True}}
)
Registration is what lets a config file or a checkpoint name your model, so it can rebuild itself. Everything else works without it.
6. Use it at any rank
lattice = td.Lattice(shape=(4, 5, 3), names=("depth", "row", "col"), valid=observed, time=True)
plan = td.ScanPlan.paired(lattice.axis_names, n_layers=8, bidirectional=("depth", "row", "col"))
model = EMA(d_model=32, lattice=lattice, plan=plan, d_input=2)
print(plan.coverage(lattice))
# Coverage(8 layers)
# time 2β 0β forward
# depth 1β 1β both
# row 1β 1β both
# col 1β 1β both
A sparse 3-D lattice with a time axis and a bidirectional schedule, from a mixer that knows about none of those things.
Common mistakes
Masking inside the mixer. Absent cells are already zero when you get them, and re-masking with a mask you derived yourself is how the two disagree.
Reading x.shape[0] as the batch. It is the batch times every unswept
axis. Anything per-example must come through the feature dimension.
Making it bidirectional internally. Set the schedule, not the module.
Depending on A at construction. The swept axis length varies by axis and
by call β the time axis has no static length at all. Learn per-channel
parameters, not per-position ones. (If you genuinely need per-position
parameters, that is the kernel family, not a mixer.)
Skipping check_block because it "obviously works". Every bug in
DEBUG.md obviously worked first.
Next
- Adding an nd_method β changing how the axes are handled, rather than what happens along one.
- DESIGN.md β why the boundary is where it is.