File size: 9,173 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
"""Joint composition: no factorization at all.

The third method of multidimensionality, and the honest baseline for the other
two. Fold **every** lattice axis into one sequence and run the 1-D mixer over
the whole thing. A rank-3 lattice of 8×8×8 becomes one sequence of 512 tokens;
the mixer sees every cell and has no idea a lattice was ever involved.

This is what a Vision Transformer does. ViT does not attend along rows and then
along columns — it flattens the patch grid into one sequence and attends over
all of it, which is why ``td.ViT`` needs this method and not
:func:`~torch_dimensions.axial_scan`. It is also what every "flatten it and
use a sequence model" baseline does, which makes it the comparison the axial
methods have to beat rather than a strawman nobody implemented.

**The cost, stated plainly.** Attention over the flattened lattice is
``O(cells²)``; axial attention is ``O(cells · A)``; the factorized kernel
family is ``O(Σ A²)``. At 8×8 that ordering barely matters and joint attention
is the most expressive of the three. At 64×64×64 it is 2.6e5 tokens and the
scores alone do not fit in memory. BENCHMARKS.md has the crossover.

**Sparse lattices become genuinely cheaper here**, and that is not true of the
other methods. Absent cells are dropped from the sequence rather than masked
into it: a lattice at 40% occupancy is a 40%-length sequence, and quadratic
attention over it costs 16% as much. The scan and kernel families must keep
absent cells in place — a recurrence has to step over them — so they can only
mask. This is the one composition where sparsity is a saving rather than a
bookkeeping obligation.
"""

from __future__ import annotations

from collections.abc import Callable

import torch
import torch.nn as nn

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

__all__ = ["Flatten"]


class Flatten(nn.Module):
    """Stack of pre-norm residual layers over the fully flattened lattice.

    Args:
        mixer: a zero-argument factory (one per layer) or a built module
            (shared), as in :class:`~torch_dimensions.AxialScan`.
        plan: contributes its **depth only**. There is no axis to choose —
            every layer mixes every axis — so the schedule's axis assignments
            and directions are not used. Kept in the signature because depth
            is a property of the plan everywhere else in the library, and a
            second way to say "how many layers" would be a second way to
            disagree.
        join_time: fold the time axis into the same sequence as the spatial
            cells (joint space-time attention, ViViT's first variant). When
            ``False``, time folds into the batch instead and each timestep is
            mixed independently — the model is then not a sequence model along
            time at all, which is right for per-frame encoders and wrong for
            forecasting.

    **This composition supplies no positional information, and with a
    permutation-invariant mixer that makes the model permutation-invariant
    too.** Attention over a set is a set function: flattened and unlabelled,
    the cells are indistinguishable, so a task whose answer depends on *where*
    a cell is — a cumulative sum along an axis, say — is not merely hard here
    but unlearnable. Measured, not argued: on a dense lattice, permuting an
    axis of the input and un-permuting the output changes nothing beyond float
    noise (~1e-6).

    That is a property of joint attention rather than a defect of this class,
    and it is why :class:`~torch_dimensions.ViT` — the flatten-family model
    meant for real use — adds a positional embedding of its own. Reach for a
    mixer that carries position (or add an embedding to the features) before
    using ``flatten`` on anything positional. The scan and kernel families do
    not have this problem: they process an axis in order, so position is
    implicit in the traversal.

    On a sparse lattice the sequence contains only cells that exist. Absent
    cells never reach the mixer, so their values cannot influence anything —
    the same guarantee the other families give by masking, obtained here by
    construction. (Note that masking *does* break the symmetry above, since
    the pattern of present cells is itself information — but incidentally,
    not usefully.)
    """

    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,
        join_time: bool = True,
    ) -> None:
        super().__init__()
        if mixer is None:
            raise ValueError(
                "the flatten method is nothing but a mixer over the flattened lattice; "
                "with mixer=None there would be no operator at all"
            )
        self.lattice = lattice
        # Resolved for consistency with the other families (it validates axis
        # names against the lattice), then used only for its length. Silently:
        # every axis is mixed in every layer here, so an "axis never swept"
        # warning would be true of the schedule and false of the model.
        self.plan = plan.resolve(lattice, warn=False)
        self.d_model = d_model
        self.residual = residual
        self.chunk = chunk
        self.join_time = join_time and lattice.time

        n = len(self.plan)
        if isinstance(mixer, nn.Module):
            self.mixers = nn.ModuleList([mixer] * n)
        else:
            self.mixers = nn.ModuleList([mixer() for _ in range(n)])
        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)

    @property
    def seq_len(self) -> int:
        """Tokens the mixer sees per row, excluding time when it is not joined."""
        return self.lattice.n_valid

    def _to_tokens(self, x: torch.Tensor) -> tuple[torch.Tensor, tuple[int, ...]]:
        """``(B, [T,] *shape, H)`` -> ``(rows, L, H)`` over present cells."""
        lat = self.lattice
        g = lat.gather(x)  # (B, [T,] G, H) — absent cells dropped
        if lat.time:
            b, t, cells, h = g.shape
            # Joined: one sequence of T·G tokens, time-major, so the flattened
            # order agrees with time order and a causal mixer stays causal in
            # time. Not joined: time folds into the batch and each timestep is
            # mixed on its own.
            seq = g.reshape(b, t * cells, h) if self.join_time else g.reshape(b * t, cells, h)
            return seq, (b, t, cells, h)
        b, cells, h = g.shape
        return g, (b, cells, h)

    def _from_tokens(self, seq: torch.Tensor, shape: tuple[int, ...]) -> torch.Tensor:
        return self.lattice.scatter(seq.reshape(*shape))

    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]}")
        # Dropping absent cells from the *sequence* keeps their values out of
        # the mixer, but the residual stream still carries whatever was sitting
        # in them, so the output at an absent cell would echo its input. The
        # other families zero on entry and after every layer; so does this one.
        # Caught by the conformance suite's mask-invariance check, which is
        # exactly the reasoning error that check exists for: "they never reach
        # the mixer" is not the same claim as "they cannot influence output".
        x = self._masked(x)
        for i in range(len(self.plan)):
            h = self.norms[i](x) if self.norms is not None else x
            seq, shape = self._to_tokens(h)
            if self.chunk is None or seq.shape[0] <= self.chunk:
                out = self.mixers[i](seq)
            else:
                rows = range(0, seq.shape[0], self.chunk)
                out = torch.cat([self.mixers[i](seq[j : j + self.chunk]) for j in rows], 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)."
                )
            h = self._from_tokens(out, shape)
            x = x + self.drop(h) if self.residual else self.drop(h)
            x = self._masked(x)
        return x

    def extra_repr(self) -> str:
        span = "space+time" if self.join_time else "space only"
        return (
            f"d_model={self.d_model}, layers={len(self.plan)}, tokens={self.seq_len} "
            f"({span}), lattice={self.lattice}"
        )