File size: 11,684 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
"""The kernel family: per-axis attention contracted across the lattice.

Where :class:`~torch_dimensions.AxialScan` sweeps a 1-D mixer along one axis
per layer, this builds an explicit ``(A, A)`` operator per axis and contracts
them all every layer β€” the factorized joint operator is a Kronecker product,
which is what keeps cost quadratic in axial size rather than in cell count.

Two variants, differing in where the kernel comes from:

- **axial attention** (``per_line=True``): scores are computed per line, so
  every row of the fold gets its own attention pattern.
- **CaFA** (``per_line=False``): features are pooled over the *other* spatial
  axes first, giving one kernel per axis per (batch, timestep) β€” the
  factorized-attention construction, cheaper and more structured.

The hybrid form: on a lattice with a time axis, the kernels own the spatial
axes and the model's 1-D mixer runs along time, each layer. This is why
``td.LSTM(nd_method=td.cafa)`` is meaningful β€” CaFA never consumes the LSTM,
it handles the axes the LSTM does not. Pooling (CaFA) deliberately keeps the
time dimension unpooled so the kernel at time ``t`` sees only time ``t``:
causality along time is the mixer's property and must not leak away through a
pooled kernel.

Scores carry a learnable relative-position bias per axis (spatial axes have
static sizes, so the table is well-defined); gating is ``"softmax"`` or
``"leaky_relu"`` (the CaFA paper's default). Sparse lattices are handled by
:func:`~torch_dimensions.axial_contract`'s per-line renormalization β€” for a
softmax kernel that renormalization *is* masked softmax, and for a signed
gate the relative cancellation guard applies.
"""

from __future__ import annotations

import math
from collections.abc import Callable

import torch
import torch.nn as nn
import torch.nn.functional as F

from torch_dimensions.compose.kernel import axial_contract
from torch_dimensions.compose.scan import axial_apply
from torch_dimensions.lattice import Lattice
from torch_dimensions.plan import ScanPlan

__all__ = ["AxialKernel"]

_GATES = ("softmax", "leaky_relu")


class AxialKernel(nn.Module):
    """Kernel-family block: per-axis attention over the lattice, optional
    mixer along time. See the module docstring.

    Args:
        mixer: zero-arg factory (or module) for the per-layer *time* mixer.
            Requires the lattice to have a time axis β€” on a purely spatial
            lattice the kernels are the whole model and a mixer would be
            silently dead weight, which is refused rather than allowed.
            Pass ``None`` for a kernel-only block.
        plan: depth and axis coverage. Each layer contracts every *spatial*
            axis the plan mentions, in the plan's first-appearance order;
            an axis the plan never names is never contracted (and
            ``plan.resolve`` warns, same as the scan family).
        per_line: per-line scores (axial attention) vs pooled per-axis
            kernels (CaFA).
        gate: ``"softmax"`` or ``"leaky_relu"``.
        qk_norm: RMS-normalize the query and key before their product. Costs
            nothing and stops the scores' scale drifting with feature norm,
            which is why the CaFA reference implementation offers it.
        kernel_residual: add a learnable ``gamma * I`` to the kernel *before*
            the gate, so a contraction starts near "keep your own value" and
            has to learn to mix. Taken from CaFA's ``LowRankKernel``, where it
            is on by default; here it is off by default so existing models are
            unchanged. ``gamma`` initializes to ``1/sqrt(d_model)``, as theirs
            does.

    Two options CaFA has that this does *not* copy: rotary position embedding
    on the query and key (this uses a learned relative-position bias table
    instead) and their spherical quadrature weights, which are a property of
    the sphere rather than of the method.
    """

    def __init__(
        self,
        mixer: Callable[[], nn.Module] | nn.Module | None,
        plan: ScanPlan,
        lattice: Lattice,
        d_model: int,
        *,
        per_line: bool = True,
        gate: str = "softmax",
        qk_norm: bool = False,
        kernel_residual: bool = False,
        dropout: float = 0.0,
        norm: bool = True,
        residual: bool = True,
        chunk: int | None = None,
    ) -> None:
        super().__init__()
        if gate not in _GATES:
            raise ValueError(f"gate must be one of {_GATES}; got {gate!r}")
        self.lattice = lattice
        self.plan = plan.resolve(lattice)
        self.d_model = d_model
        self.per_line = per_line
        self.gate = gate
        self.qk_norm = qk_norm
        self.kernel_residual = kernel_residual
        self.residual = residual
        self.chunk = chunk

        time_index = 0 if lattice.time else None
        self.spatial_axes = [int(a) for a in self.plan.axes if a != time_index]
        if not self.spatial_axes:
            raise ValueError("the kernel family needs at least one spatial axis in the plan")

        n = len(self.plan)
        h = d_model
        # Flat (layer, axis) indexing β€” layer * n_axes + j β€” keeps every
        # per-axis module addressable without nested ModuleLists.
        n_ax = len(self.spatial_axes)
        self.q = nn.ModuleList(nn.Linear(h, h, bias=False) for _ in range(n * n_ax))
        self.k = nn.ModuleList(nn.Linear(h, h, bias=False) for _ in range(n * n_ax))
        sizes = [lattice.axis_size(a) for a in self.spatial_axes]
        self.bias = nn.ParameterList(
            nn.Parameter(torch.zeros(a, a)) for _ in range(n) for a in sizes
        )
        # One gamma per (layer, axis), like the per-axis kernels it scales.
        # Initialized to 1/sqrt(d_model) β€” CaFA's value.
        self.gamma = (
            nn.ParameterList(nn.Parameter(torch.tensor(h**-0.5)) for _ in range(n * n_ax))
            if kernel_residual
            else None
        )
        self.out = nn.ModuleList(nn.Linear(h, h) for _ in range(n))
        self.norms = nn.ModuleList(nn.LayerNorm(h) for _ in range(n)) if norm else None
        self.drop = nn.Dropout(dropout)

        if lattice.time:
            if mixer is None:
                self.mixers = None
            elif isinstance(mixer, nn.Module):
                self.mixers = nn.ModuleList([mixer] * n)  # shared, as in AxialScan
            else:
                self.mixers = nn.ModuleList([mixer() for _ in range(n)])
            self.time_norms = (
                nn.ModuleList(nn.LayerNorm(h) for _ in range(n))
                if norm and self.mixers is not None
                else None
            )
        else:
            if mixer is not None:
                raise ValueError(
                    "a mixer was given but the lattice has no time axis for it to sweep; "
                    "the kernel family owns every spatial axis, so the mixer would be dead "
                    "weight. Use a lattice with time=True (the hybrid form) or mixer=None."
                )
            self.mixers = None
            self.time_norms = None

        if not lattice.is_dense:
            self.register_buffer("cell_mask", lattice.mask(torch.bool), persistent=False)
        else:
            self.cell_mask = None

    def _masked(self, x: torch.Tensor) -> torch.Tensor:
        if self.cell_mask is None:
            return x
        return x * self.cell_mask.to(x.dtype)

    def _qk_norm(self, x: torch.Tensor) -> torch.Tensor:
        """RMS-normalize the last dimension, or pass through.

        No learnable scale: the score already has one in `scale`, and a second
        would be the same parameter twice.
        """
        if not self.qk_norm:
            return x
        return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + 1e-6)

    def _kernel(self, layer: int, j: int, axis: int, h: torch.Tensor) -> torch.Tensor:
        """Build the ``(M, A, A)`` operator for one axis of one layer."""
        lat = self.lattice
        scale = 1.0 / math.sqrt(self.d_model)
        idx = layer * len(self.spatial_axes) + j
        bias = self.bias[idx]

        if self.per_line:
            seq, _ = lat.to_sequence(h, axis)  # (M, A, H)
            q = self._qk_norm(self.q[idx](seq))
            k = self._qk_norm(self.k[idx](seq))
            scores = q @ k.transpose(1, 2) * scale + bias
        else:
            # Pool over the *other spatial* axes only. Batch stays batch, and
            # time deliberately stays unpooled: a kernel at time t built from
            # future timesteps would leak the future into a "causal" model.
            d = lat.tensor_dim(axis)
            keep = {0, d, h.ndim - 1}
            if lat.time:
                keep.add(1)
            reduce_dims = tuple(i for i in range(h.ndim) if i not in keep)
            counts = lat.valid_counts(axis).to(h.dtype).to(h.device)
            pooled = h.sum(reduce_dims) if reduce_dims else h  # (B, [T,] A, H)
            pooled = pooled / counts.unsqueeze(-1)
            q = self._qk_norm(self.q[idx](pooled))
            k = self._qk_norm(self.k[idx](pooled))
            scores = q @ k.transpose(-1, -2) * scale + bias  # (B, [T,] A, A)
            a = scores.shape[-1]
            # Expand to one kernel per folded line. The fold orders leading
            # dims as (B, [T,] *others), batch-major, so lines sharing a
            # (batch, timestep) are contiguous.
            lines_per = 1
            for i in range(1, h.ndim - 1):
                if i != d:
                    lines_per *= h.shape[i]
            if lat.time:
                lines_per //= h.shape[1]
                scores = scores.reshape(-1, 1, a, a).expand(-1, lines_per, a, a)
            else:
                scores = scores.reshape(-1, 1, a, a).expand(-1, lines_per, a, a)
            scores = scores.reshape(-1, a, a)

        if self.gamma is not None:
            # `gamma * I`, added before the gate exactly as CaFA does β€” after
            # the gate it would be a different model, since softmax is not
            # additive. A contraction therefore starts near the identity and
            # has to learn to mix, rather than starting fully mixed.
            eye = torch.eye(scores.shape[-1], device=scores.device, dtype=scores.dtype)
            scores = scores + self.gamma[idx] * eye

        if self.gate == "softmax":
            return F.softmax(scores, dim=-1)
        return F.leaky_relu(scores)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        if x.shape[-1] != self.d_model:
            raise ValueError(f"expected {self.d_model} features, got {x.shape[-1]}")
        x = self._masked(x)
        valid = None if self.cell_mask is None else self.cell_mask.to(x.dtype)
        for i in range(len(self.plan)):
            h = self.norms[i](x) if self.norms is not None else x
            for j, axis in enumerate(self.spatial_axes):
                kernel = self._kernel(i, j, axis, h)
                h = axial_contract(h, self.lattice, axis, kernel, valid=valid)
            h = self.out[i](h)
            x = x + self.drop(h) if self.residual else self.drop(h)
            x = self._masked(x)
            if self.mixers is not None:
                h = self.time_norms[i](x) if self.time_norms is not None else x
                h = axial_apply(h, self.lattice, 0, self.mixers[i], chunk=self.chunk)
                x = x + self.drop(h)
                x = self._masked(x)
        return x

    def extra_repr(self) -> str:
        kind = "per-line" if self.per_line else "pooled (CaFA)"
        return f"d_model={self.d_model}, {kind}, gate={self.gate}, lattice={self.lattice}"