File size: 7,752 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
"""Sequential composition: sweep a 1-D mixer along one lattice axis at a time.

This is what makes N-D tractable without an N-D kernel. Permute the swept axis
to the sequence position, fold every other axis into the batch, run an ordinary
1-D operator, and permute back. Alternating the axis across layers recovers N-D
context, so a 1-D CUDA kernel is reused unchanged.

:func:`axial_apply` is the whole mechanism and is deliberately a pure function:
no parameters, no state, nothing to configure. :class:`AxialScan` is the thin
module that wraps it in pre-norm residual layers and walks a
:class:`~torch_dimensions.ScanPlan`. Keeping them separate means the axis
bookkeeping can be tested against an explicit per-line reference without any
norm or residual in the way.
"""

from __future__ import annotations

import inspect
from collections.abc import Callable

import torch
import torch.nn as nn

from torch_dimensions.lattice import AxisSpec, Lattice
from torch_dimensions.plan import ScanPlan

__all__ = ["AxialScan", "axial_apply"]

# What a mixer factory may ask to be told about the layer it is being built
# for. Both are opt-in by signature: a factory that does not name the argument
# is called exactly as before, so this is invisible to every mixer that does
# not want it.
#
# `sweep` — not `layer` — is the one that matters for N-D. A dilated
# convolution's window should grow with how many times *its own axis* has been
# swept, because that is what its receptive field along that axis composes
# over. Growing with the global layer index instead would give the third sweep
# of a rank-3 lattice a dilation of 64 where 4 is meant, and the model would
# quietly be looking past the end of every line.
_LAYER_ARGS = ("layer", "sweep")


def _layer_kwargs(factory: Callable[..., nn.Module], layer: int, sweep: int) -> dict[str, int]:
    try:
        params = inspect.signature(factory).parameters
    except (TypeError, ValueError):  # builtins, C callables
        return {}
    available = {"layer": layer, "sweep": sweep}
    # Only names the signature spells out. A factory with `**kwargs` would
    # swallow these silently, and a mixer that never asked for a `sweep`
    # should not have one appear in its config.
    return {k: v for k, v in available.items() if k in params and k in _LAYER_ARGS}


def axial_apply(
    x: torch.Tensor,
    lattice: Lattice,
    axis: AxisSpec,
    fn: Callable[[torch.Tensor], torch.Tensor],
    *,
    reverse: bool = False,
    chunk: int | None = None,
) -> torch.Tensor:
    """Apply a 1-D operator independently along every line of ``axis``.

    ``fn`` receives ``(M, A, H)`` — a plain batch of 1-D sequences — and must
    return the same shape. It is never told which axis it is sweeping, how many
    axes exist, or how long the other axes are. That ignorance is the extension
    point: any 1-D module is a valid mixer.

    ``reverse`` flips the sequence before the call and unflips after, so the
    operator sees the line back-to-front while the output stays in lattice
    order.

    ``chunk`` caps how many folded rows go through ``fn`` at once. Scanning a
    short axis folds every other axis into ``M``, which can reach tens of
    thousands of rows and overrun a fused kernel's grid limits. Pure-torch
    mixers have no such limit, so this defaults to off; the kernel adapters
    that need it set it themselves rather than exposing a magic constant.
    """
    if chunk is not None and chunk < 1:
        raise ValueError(f"chunk must be >= 1 or None; got {chunk}")
    seq, restore = lattice.to_sequence(x, axis)
    if reverse:
        seq = seq.flip(1)

    if chunk is None or seq.shape[0] <= chunk:
        out = fn(seq)
    else:
        out = torch.cat([fn(seq[i : i + chunk]) for i in range(0, seq.shape[0], chunk)], dim=0)

    if out.shape != seq.shape:
        raise ValueError(
            f"mixer changed shape: got {tuple(out.shape)}, expected {tuple(seq.shape)}. "
            "A mixer must map (M, A, H) -> (M, A, H)."
        )
    if reverse:
        out = out.flip(1)
    return lattice.from_sequence(out, restore)


class AxialScan(nn.Module):
    """Stack of pre-norm residual layers, each sweeping one axis of a lattice.

    Args:
        mixer: either a zero-argument factory called once per layer (each layer
            gets its own weights, as in Mamba-ND), or an already-built
            ``nn.Module``, in which case every layer *shares* it.
        plan: which axis each layer sweeps and in which direction. Resolved
            against ``lattice`` at construction, so unknown axes fail here
            rather than at the first forward pass.
        lattice: the grid being swept.
        d_model: feature width. A mixer maps ``d_model -> d_model``; this stack
            never changes the feature dimension.
        chunk: see :func:`axial_apply`.

    On a lattice with absent cells, those cells are zeroed on entry and after
    every layer. Zeroing on *entry* is what makes the outputs at present cells
    independent of whatever values were sitting in absent ones. Note that an
    absent cell still occupies a position in its line, so it advances a
    recurrence — the guarantee is invariance to its *value*, not to its
    existence.
    """

    def __init__(
        self,
        mixer: Callable[[], nn.Module] | nn.Module,
        plan: ScanPlan,
        lattice: Lattice,
        d_model: int,
        *,
        dropout: float = 0.0,
        norm: bool = True,
        residual: bool = True,
        chunk: int | None = None,
    ) -> None:
        super().__init__()
        self.lattice = lattice
        self.plan = plan.resolve(lattice)
        self.d_model = d_model
        self.residual = residual
        self.chunk = chunk

        n = len(self.plan)
        if isinstance(mixer, nn.Module):
            # Shared weights, by request. A shared mixer is one object, so it
            # cannot vary per layer — no dilation schedule, no per-sweep
            # anything. That is the trade the caller made by passing an
            # instance instead of a factory.
            self.mixers = nn.ModuleList([mixer] * n)
        else:
            swept: dict[int, int] = {}
            built = []
            for i, step in enumerate(self.plan):
                axis = int(step.axis)
                built.append(mixer(**_layer_kwargs(mixer, i, swept.get(axis, 0))))
                swept[axis] = swept.get(axis, 0) + 1
            self.mixers = nn.ModuleList(built)
        self.norms = nn.ModuleList([nn.LayerNorm(d_model) for _ in range(n)]) if norm else None
        self.drop = nn.Dropout(dropout)

        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 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)
        for i, step in enumerate(self.plan):
            h = self.norms[i](x) if self.norms is not None else x
            h = axial_apply(
                h,
                self.lattice,
                step.axis,
                self.mixers[i],
                reverse=step.reverse,
                chunk=self.chunk,
            )
            x = x + self.drop(h) if self.residual else self.drop(h)
            x = self._masked(x)
        return x

    def extra_repr(self) -> str:
        return f"d_model={self.d_model}, plan={self.plan}, lattice={self.lattice}"