Isometric tensor-network token mixer / isometric_mixer.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import math
 2import torch
 3from torch import nn
 4
 5
 6class IsometricButterflyMixer(nn.Module):
 7    """Trainable orthogonal butterfly on (..., tokens, features)."""
 8    def __init__(self, n_tokens: int, init_scale: float = 0.02):
 9        super().__init__()
10        if n_tokens < 2 or n_tokens & (n_tokens - 1):
11            raise ValueError("n_tokens must be a power of two")
12        self.n = n_tokens
13        self.levels = int(math.log2(n_tokens))
14        self.angles = nn.Parameter(init_scale * torch.randn(self.levels, n_tokens // 2))
15
16    def gates(self):
17        a = self.angles
18        c, s = torch.cos(a), torch.sin(a)
19        return torch.stack((torch.stack((c, -s), -1),
20                            torch.stack((s, c), -1)), -2)
21
22    def apply_transform(self, x, inverse=False):
23        if x.shape[-2] != self.n:
24            raise ValueError(f"expected token axis {-2} to have size {self.n}")
25        z = x
26        gs = self.gates()
27        levels = range(self.levels - 1, -1, -1) if inverse else range(self.levels)
28        for level in levels:
29            half, step = 2 ** level, 2 ** (level + 1)
30            out = z.clone()
31            for base in range(0, self.n, step):
32                for j in range(half):
33                    ids = [base + j, base + j + half]
34                    g = gs[level, base // step * half + j]
35                    if inverse:
36                        g = g.transpose(-1, -2)
37                    out[..., ids, :] = torch.einsum('ab,...bq->...aq', g, z[..., ids, :])
38            z = out
39        return z
40
41    def forward(self, x):
42        return self.apply_transform(x)
43
44    def adjoint(self, x):
45        return self.apply_transform(x, inverse=True)
46
47
48class IsometricTokenBlock(nn.Module):
49    """Token mixer followed by a pointwise MLP and residual connection."""
50    def __init__(self, n_tokens, width, isometric=True):
51        super().__init__()
52        self.isometric = isometric
53        self.mixer = IsometricButterflyMixer(n_tokens) if isometric else nn.Linear(n_tokens, n_tokens)
54        self.norm = nn.LayerNorm(width)
55        self.mlp = nn.Sequential(nn.Linear(width, 2 * width), nn.GELU(), nn.Linear(2 * width, width))
56        self.alpha = nn.Parameter(torch.tensor(0.1))
57
58    def forward(self, x):
59        if self.isometric:
60            mixed = self.mixer(x)
61        else:
62            mixed = self.mixer(x.transpose(-1, -2)).transpose(-1, -2)
63        return x + self.alpha * self.mlp(self.norm(mixed))