File size: 5,656 Bytes
3e77c56
 
 
 
 
2c93889
 
 
 
 
3e77c56
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2c93889
3e77c56
 
 
 
 
 
2c93889
 
 
 
 
 
3e77c56
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2c93889
 
 
 
 
 
 
3e77c56
 
 
 
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
"""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)