"""LinearNO attention block (Stage 2 / Gate B). Reimplemented from the equations of: "Transolver is a Linear Transformer: Revisiting Physics-Attention through the Lens of Linear Attention" — Hu, Liu, Qiao, Sun, Dou (NUDT), AAAI 2026, arXiv:2511.06294. Reconstructed from the equations BEFORE the official code was released (github.com/HiPRL/LinearNO, Jan 2026); we did not consult it. A post-hoc check confirms structural agreement (softmax axes, associativity). Differences from the official Elasticity block: it adds a learnable per-head softmax temperature (init 0.5, clamped [0.01,1]) that we omit, and uses one shared in_project_x lift + small per-head q/k/v maps vs our full C->inner projections. Only the attention block differs from Transolver. Core idea: Physics-Attention is the special case of linear attention ``Attention(Q,K,V) ~ phi(Q) (psi^T(K) V)`` in which (a) phi and psi come from the SAME linear layer (differing only by normalization) and (b) there is an extra slice self-attention step. LinearNO removes BOTH constraints: 1. learn Q and K projections independently (break weight sharing), and 2. drop the slice self-attention (identity). LinearNO(H) = phi(Q) @ ( psi(K)^T @ V ) Q = H Wq ; K = H Wk ; V = Linear_V(H) phi(Q) = softmax_over_M( Linear_Q(Q) ) # (N, M) rows sum to 1 <- softmax along M psi(K) = softmax_over_N( Linear_K(K) ) # (N, M) cols sum to 1 <- softmax along N ### THE make-or-break detail (paper Table 6, Elasticity): phi softmax over M (slices), psi softmax over N (points) -> 0.0050 (correct) swapping these dims -> 0.0081..0.0112 (wrong). If LinearNO lands at 0.008-0.011, the softmax dimensions are almost certainly swapped. Associativity: compute ``psi^T V`` first (M x d), then ``phi @ (...)`` -> O(N*M*d), linear in N. Two variants (see docs/RECONCILIATION.md for the parameter-count tension): - "independent" (paper skeleton, default): separate Wq, Wk, Wv as full dim->inner projections. Most literal to the plan's reference code. ~0.85M params (> baseline). - "shared_qk": share one dim->inner base for the slice projections (phi, psi), keep a separate dim->inner V; breaks slice-weight sharing via two separate small slice layers. ~0.72M params (~= baseline). Use this to satisfy the Gate-B "<= baseline params" constraint. An optional ``project_out=False`` drops the output projection (folds the per-head concat directly), reaching ~0.59M (the plan's stated LinearNO size). """ from __future__ import annotations import torch import torch.nn as nn from einops import rearrange class LinearNO(nn.Module): def __init__( self, dim, heads=8, dim_head=16, slice_num=64, dropout=0.0, variant: str = "independent", project_out: bool = True, temperature: bool = False, ): super().__init__() inner = heads * dim_head self.h, self.m, self.dh = heads, slice_num, dim_head self.variant = variant self.project_out = project_out self.temperature = temperature if temperature: # Matches the official LinearNO `temp` block: learnable per-head temperature on both # softmaxes, init 0.5, clamped to [0.01, 1] (github.com/HiPRL/LinearNO). self.temp_q = nn.Parameter(torch.ones(1, heads, 1, 1) * 0.5) self.temp_k = nn.Parameter(torch.ones(1, heads, 1, 1) * 0.5) if variant == "independent": # Independent Q, K, V projections (paper modification 1, literal). self.to_q = nn.Linear(dim, inner, bias=False) self.to_k = nn.Linear(dim, inner, bias=False) self.to_v = nn.Linear(dim, inner, bias=False) elif variant == "shared_qk": # Share one base projection for the slice space (Q==K base), keep V separate. self.to_qk = nn.Linear(dim, inner, bias=False) self.to_v = nn.Linear(dim, inner, bias=False) else: raise ValueError(f"unknown variant {variant!r} (expected 'independent' or 'shared_qk')") # Asymmetric slice projections: query->slices and key->slices are SEPARATE weights. self.lin_q = nn.Linear(dim_head, slice_num) self.lin_k = nn.Linear(dim_head, slice_num) if project_out: self.to_out = nn.Sequential(nn.Linear(inner, dim), nn.Dropout(dropout)) else: self.to_out = nn.Dropout(dropout) def _heads(self, t, B, N): return t.reshape(B, N, self.h, self.dh).permute(0, 2, 1, 3).contiguous() # (B,H,N,dh) def forward(self, x): # x: (B, N, C) B, N, C = x.shape if self.variant == "independent": q = self._heads(self.to_q(x), B, N) # (B,H,N,dh) k = self._heads(self.to_k(x), B, N) v = self._heads(self.to_v(x), B, N) else: # shared_qk base = self._heads(self.to_qk(x), B, N) q = base k = base v = self._heads(self.to_v(x), B, N) sq = self.lin_q(q) # (B,H,N,M) sk = self.lin_k(k) # (B,H,N,M) if self.temperature: sq = sq / self.temp_q.clamp(0.01, 1.0) sk = sk / self.temp_k.clamp(0.01, 1.0) phi = sq.softmax(dim=-1) # softmax OVER M (slices) <- rows sum to 1 psi = sk.softmax(dim=-2) # softmax OVER N (points) <- cols sum to 1 kv = torch.einsum("bhnm,bhnd->bhmd", psi, v) # (B,H,M,dh) cheap inner product first out = torch.einsum("bhnm,bhmd->bhnd", phi, kv) # (B,H,N,dh) linear in N out = rearrange(out, "b h n d -> b n (h d)") return self.to_out(out)