import math import torch from torch import nn class IsometricButterflyMixer(nn.Module): """Trainable orthogonal butterfly on (..., tokens, features).""" def __init__(self, n_tokens: int, init_scale: float = 0.02): super().__init__() if n_tokens < 2 or n_tokens & (n_tokens - 1): raise ValueError("n_tokens must be a power of two") self.n = n_tokens self.levels = int(math.log2(n_tokens)) self.angles = nn.Parameter(init_scale * torch.randn(self.levels, n_tokens // 2)) def gates(self): a = self.angles c, s = torch.cos(a), torch.sin(a) return torch.stack((torch.stack((c, -s), -1), torch.stack((s, c), -1)), -2) def apply_transform(self, x, inverse=False): if x.shape[-2] != self.n: raise ValueError(f"expected token axis {-2} to have size {self.n}") z = x gs = self.gates() levels = range(self.levels - 1, -1, -1) if inverse else range(self.levels) for level in levels: half, step = 2 ** level, 2 ** (level + 1) out = z.clone() for base in range(0, self.n, step): for j in range(half): ids = [base + j, base + j + half] g = gs[level, base // step * half + j] if inverse: g = g.transpose(-1, -2) out[..., ids, :] = torch.einsum('ab,...bq->...aq', g, z[..., ids, :]) z = out return z def forward(self, x): return self.apply_transform(x) def adjoint(self, x): return self.apply_transform(x, inverse=True) class IsometricTokenBlock(nn.Module): """Token mixer followed by a pointwise MLP and residual connection.""" def __init__(self, n_tokens, width, isometric=True): super().__init__() self.isometric = isometric self.mixer = IsometricButterflyMixer(n_tokens) if isometric else nn.Linear(n_tokens, n_tokens) self.norm = nn.LayerNorm(width) self.mlp = nn.Sequential(nn.Linear(width, 2 * width), nn.GELU(), nn.Linear(2 * width, width)) self.alpha = nn.Parameter(torch.tensor(0.1)) def forward(self, x): if self.isometric: mixed = self.mixer(x) else: mixed = self.mixer(x.transpose(-1, -2)).transpose(-1, -2) return x + self.alpha * self.mlp(self.norm(mixed))