File size: 12,077 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
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
"""The Vision Transformer, and why it is an N-D model at all.

ViT is usually described as "images as sequences of patches", which hides the
lattice: the patches *are* a grid, and every design choice that distinguishes
ViT variants is a choice about how that grid is handled.

    td.ViT(d_model=192, n_layers=6, image=(32, 32), patch=4)     # joint, as published
    td.ViT(..., method=td.axial_scan)                            # axial ViT

The first is the published model: patchify, flatten the grid to one sequence,
attend over all of it. The second is the axial variant — one argument apart,
which is the entire point of the library. On an 8×8 patch grid the joint form
attends over 64 tokens; the axial form does two passes of 8. Same patches,
same parameter count, different method of multidimensionality, and
BENCHMARKS.md says which is cheaper where.

The kernel family (``td.cafa``, ``td.axial_attention``) is deliberately *not*
available here, and the refusal is informative rather than a gap: those
methods own every spatial axis themselves and leave the mixer to run along
time, so on a time-less patch grid the transformer blocks would be dead
weight. A factorized-attention model over a patch grid is
``td.AxialKernel(mixer=None, ...)`` — the kernels are the model. Give the
lattice a time axis (video) and the hybrid form applies again.

**What this ships and what it does not.** Patch embedding, positional
embedding, and the transformer stack over the patch lattice, returning
per-patch features ``(B, *grid, d_model)``. No class token, no pooling, no
classification head — the same boundary every other model in the library
keeps. A head is three lines of caller code and it is the caller's three
lines. Concretely::

    vit = td.ViT(192, 6, image=(32, 32), patch=4, in_channels=3)
    head = nn.Linear(192, 10)
    logits = head(vit(images).mean(dim=(1, 2)))    # mean-pool the grid

**Positional embedding is where the lattice shows.** ViT learns one embedding
per patch position — a table the size of the grid, which cannot transfer to a
different image size. The factorized alternative learns one table per *axis*
and adds them, which is ``r·A`` parameters instead of ``A^r`` and extends to a
new grid size by interpolating one axis at a time. Both are here; factorized
is the default because on a lattice it is the natural one, and the published
choice is one argument away.
"""

from __future__ import annotations

from collections.abc import Sequence

import torch
import torch.nn as nn

from torch_dimensions.compose import flatten
from torch_dimensions.lattice import Lattice
from torch_dimensions.mixers.attention import AttentionMixer
from torch_dimensions.models.base import LatticeModel

__all__ = ["PatchEmbed", "ViT"]


class PatchEmbed(nn.Module):
    """Cut an image into patches and embed each one — image to lattice.

    ``(B, *image, C)`` in, ``(B, *grid, d_model)`` out, where
    ``grid[i] = image[i] // patch[i]``. Rank-generic: a 2-D image, a 3-D
    volume, and a 4-D spatio-temporal block all work, because the patching is
    a reshape and a linear map rather than a ``Conv2d`` with a rank baked in.

    Args:
        image: size of each input axis.
        patch: patch size per axis; an int applies to every axis. Must divide
            the image exactly — a partial patch at the edge is a silent crop,
            and cropping the user's data without saying so is not this
            module's decision to make.
        in_channels: channels of the input (3 for RGB, 1 for greyscale).
        d_model: embedding width per patch.
    """

    def __init__(
        self,
        image: Sequence[int],
        patch: Sequence[int] | int,
        in_channels: int,
        d_model: int,
    ) -> None:
        super().__init__()
        self.image = tuple(int(s) for s in image)
        rank = len(self.image)
        self.patch = (
            (int(patch),) * rank if isinstance(patch, int) else tuple(int(p) for p in patch)
        )
        if len(self.patch) != rank:
            raise ValueError(f"patch {self.patch} has {len(self.patch)} axes, image has {rank}")
        bad = [(s, p) for s, p in zip(self.image, self.patch, strict=True) if p < 1 or s % p]
        if bad:
            raise ValueError(
                f"patch size must divide the image exactly; got image={self.image}, "
                f"patch={self.patch}. A partial patch at the edge would silently crop the "
                "input — pad or resize before this point, deliberately."
            )
        self.grid = tuple(s // p for s, p in zip(self.image, self.patch, strict=True))
        self.in_channels = in_channels
        self.d_model = d_model
        self.n_patch_features = in_channels
        for p in self.patch:
            self.n_patch_features *= p
        self.proj = nn.Linear(self.n_patch_features, d_model)

    def lattice(self, *, names: Sequence[str] | None = None, time: bool = False) -> Lattice:
        """The patch grid as a lattice — what the model actually operates on."""
        return Lattice(shape=self.grid, names=tuple(names) if names else None, time=time)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        rank = len(self.image)
        if x.ndim != rank + 2:
            raise ValueError(
                f"expected a {rank + 2}-D tensor (B, *{self.image}, {self.in_channels}); "
                f"got shape {tuple(x.shape)}"
            )
        if tuple(x.shape[1:-1]) != self.image:
            raise ValueError(f"expected image dims {self.image}, got {tuple(x.shape[1:-1])}")
        if x.shape[-1] != self.in_channels:
            raise ValueError(f"expected {self.in_channels} channels, got {x.shape[-1]}")

        b = x.shape[0]
        # Split each image axis into (grid, patch), then move every patch axis
        # next to the channel axis and flatten them together. Written as one
        # rank-generic reshape+permute rather than einops or a per-rank table,
        # for the same reason the fold is: a rank-4 case must not need new code.
        split: list[int] = [b]
        for g, p in zip(self.grid, self.patch, strict=True):
            split += [g, p]
        split.append(self.in_channels)
        h = x.reshape(*split)

        grid_dims = [1 + 2 * i for i in range(rank)]
        patch_dims = [2 + 2 * i for i in range(rank)]
        h = h.permute(0, *grid_dims, *patch_dims, h.ndim - 1).contiguous()
        h = h.reshape(b, *self.grid, self.n_patch_features)
        return self.proj(h)

    def extra_repr(self) -> str:
        return (
            f"image={self.image}, patch={self.patch}, grid={self.grid}, "
            f"in_channels={self.in_channels}, d_model={self.d_model}"
        )


class _PosEmbed(nn.Module):
    """Learned positional embedding over a patch grid.

    ``factorized`` learns one table per axis and adds them (``r·A``
    parameters); ``full`` learns one per cell (``A^r``), which is what ViT
    publishes. Factorized is the default: on a lattice it is the natural
    parameterization, it is what makes a 3-D or 4-D grid affordable, and the
    axial models in this library already assume per-axis structure everywhere
    else.
    """

    def __init__(self, grid: tuple[int, ...], d_model: int, kind: str) -> None:
        super().__init__()
        if kind not in ("factorized", "full", "none"):
            raise ValueError(f"pos_embed must be factorized|full|none; got {kind!r}")
        self.kind = kind
        self.grid = grid
        if kind == "factorized":
            self.tables = nn.ParameterList(nn.Parameter(torch.zeros(n, d_model)) for n in grid)
        elif kind == "full":
            self.tables = nn.ParameterList([nn.Parameter(torch.zeros(*grid, d_model))])
        else:
            self.tables = nn.ParameterList()
        for t in self.tables:
            nn.init.trunc_normal_(t, std=0.02)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        if self.kind == "none":
            return x
        if self.kind == "full":
            return x + self.tables[0]
        rank = len(self.grid)
        for axis, table in enumerate(self.tables):
            # Broadcast the axis's table along every other grid axis.
            shape = [1] * rank
            shape[axis] = self.grid[axis]
            x = x + table.reshape(*shape, -1)
        return x

    def extra_repr(self) -> str:
        n = sum(p.numel() for p in self.tables)
        return f"{self.kind}, grid={self.grid}, {n} parameters"


class ViT(LatticeModel):
    """Vision Transformer over a patch lattice, at any rank.

    Args:
        d_model: embedding width.
        n_layers: transformer blocks.
        image: input size per axis — ``(32, 32)`` for CIFAR, ``(16, 64, 64)``
            for a volume.
        patch: patch size per axis, or one int for all.
        in_channels: input channels.
        pos_embed: ``"factorized"`` (default), ``"full"`` (ViT's own), or
            ``"none"``.
        method: the method of multidimensionality. Defaults to
            :func:`~torch_dimensions.flatten` — attention over all patches at
            once, which is what ViT does. ``td.axial_scan`` gives the axial
            variant. The kernel family needs a time axis; see the module
            docstring.
        mixer_kwargs: forwarded to
            :class:`~torch_dimensions.mixers.attention.AttentionMixer`, e.g.
            ``{"n_heads": 12}``.

    Returns per-patch features ``(B, *grid, d_model)``. Pool and classify
    yourself; see the module docstring.
    """

    _mixer = AttentionMixer

    def __init__(
        self,
        d_model: int,
        n_layers: int = 1,
        *,
        image: Sequence[int],
        patch: Sequence[int] | int = 16,
        in_channels: int = 3,
        pos_embed: str = "factorized",
        names: Sequence[str] | None = None,
        n_heads: int = 4,
        **kw,
    ) -> None:
        # The lattice is derived, so accepting one would be accepting a second
        # answer to a question already answered. A checkpoint records
        # `lattice: None` for exactly this reason, and rebuilding hands it back.
        if kw.pop("lattice", None) is not None:
            raise ValueError(
                "ViT builds its lattice from `image` and `patch`; passing `lattice` too "
                "would let the two disagree. Pass image/patch, or use td.Transformer "
                "directly if you already have a lattice of tokens."
            )
        embed = PatchEmbed(image, patch, in_channels, d_model)
        lat = embed.lattice(names=names)
        # Only default when the caller named neither spelling: setting
        # `nd_method` unconditionally would collide with a caller's `method=`
        # and the base class refuses both, as it should.
        if "method" not in kw and "nd_method" not in kw:
            kw["nd_method"] = flatten
        mixer_kwargs = {"n_heads": n_heads, **kw.pop("mixer_kwargs", {})}
        super().__init__(d_model, n_layers, lat, mixer_kwargs=mixer_kwargs, **kw)
        # Registered after super().__init__ so the base class's parameter
        # accounting and the recorded config are unaffected by their presence.
        self.patch_embed = embed
        self.pos = _PosEmbed(embed.grid, d_model, pos_embed)
        self.config.update(
            {
                "image": list(embed.image),
                "patch": list(embed.patch),
                "in_channels": in_channels,
                "pos_embed": pos_embed,
                "n_heads": n_heads,
                # The lattice is derived from image/patch, so recording it too
                # would let a checkpoint hold two answers to one question.
                "lattice": None,
            }
        )

    @property
    def grid(self) -> tuple[int, ...]:
        """Patch-grid shape — the lattice the transformer runs over."""
        return self.patch_embed.grid

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """``(B, *image, C)`` in, ``(B, *grid, d_model)`` out."""
        return self.nd(self.pos(self.patch_embed(x)))