Isometric tensor-network token mixer / isometric_mixer.py
Mechanism confirmed, baseline not beaten
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))