File size: 7,086 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
# 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`](../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

```python
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=True` on 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

```python
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

```python
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.

```python
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:

```python
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*:

```python
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)

```python
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

```python
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](../DEBUG.md) obviously worked first.

## Next

- [Adding an nd_method](adding-a-method.md) β€” changing *how* the axes are
  handled, rather than what happens along one.
- [DESIGN.md](../DESIGN.md) β€” why the boundary is where it is.