Vision Transformers devote most of their parameters to MLPs for channel mixing, but still rely on quadratic multi-head self-attention for token interactions. While linear attention fixes the complexity problem, bringing it down to O(N), it is usually just paired with the same fixed-activation MLP as before. Kolmogorov-Arnold Networks take a different approach, placing learnable univariate functions on the edges instead. However, existing vision KANs either retain standard attention or remove attention entirely, so the two ideas have not been effectively combined. We introduce LKAT (Linear Kolmogorov-Arnold Transformer) to close this gap: an isotropic ViT-style encoder that couples chunk-wise Gated Linear Attention with a two-layer KAN feed-forward block, backed by an I/O-aware fused RBF-KAN kernel to make radial-basis grid functions efficient in practice. Under a shared DeiT-style training recipe, LKAT-B outperforms ViT-B/16, ViT-5-B, and Mixer-B/16 on ImageNet-100, while Tiny, Small, and Base variants scale consistently on CIFAR-10/100. ImageNet-100 pretraining also transfers effectively to CIFAR fine-tuning, suggesting that gated linear attention and KAN-based radial basis functions provide complementary inductive biases for mid-scale visual representation learning. Code: https://github.com/mehizelali/linear-kan-transformer
Figures & tables
Input: flattened x∈RN in HBM
centers c∈RG , weights w∈RG in HBM
inv←1/d ; block size BLOCK
Output: y∈RN in HBM
1:
for each program p in parallel do
2:
Load tile x[i] , i∈Block(p) (HBM → on-chip)
3:
acc←0 (registers)
Algorithm 1 Fused RBF–grid linear forward
Figure 1: Peak memory reduction of the fused Triton RBF–grid linear operator versus an unfused torch baseline for a joint forward+backward pass, as a function of channel width D . Bars report Memtorch/MemTriton at grid sizes G∈{4,8,16} (dotted line =1 ). The reduction is essentially independent of D and scales nearly linearly with G ( ≈5.5× , 10.5× , and 20.5× ), consistent with avoiding materialization of the intermediate […,D,G] expansion in HBM.
Figure 2: LKAT architecture. Patch embedding feeds N pre-norm blocks (LayerNorm then Gated Linear Attention with per-head RMSNorm on Q/K; LayerNorm then KAN-based radial-basis functions), each with a residual skip. A linear head reads the final class token.
Figure 3: Learned RBF–KAN edge functions ϕp,q(x) for ImageNet-100 LKAT-B (representative blocks; full set in Figs. 9 – 10 ).
Hyperparameter
Value
Resolution / grid / bases
2242 / G=4 / RBF
Registers / RoPE
4 / 2D RoPE ( 2dv1 )
Optimizer
AdamW
Peak LR / weight decay
3×10−4 / 0
Warmup / min. LR
5 epochs / 10−5
Global batch
256 / 512
Table 1: Shared from-scratch experimental setup. All models (LKAT, rbfKAT, ViT, ViT-5, MLP-Mixer) use this recipe; only the block operators differ.
Table 3: LKAT variants at 2242 . Params (M) and GFLOPs for one forward pass.
Figure 4: ImageNet-100 Top-1 Acc@1 versus GFLOPs ( 2242 ).
Figure 5: Gated residual around the RBF–KAN block (Eqs. ( 9 )–( 12 )).
Variant
Params (M)
GFLOPs
Acc@1
Gated residual
LKAT-B
93.99
37.34
80.14
+ gate
149.96
60.10
82.02
RBF grid size G
G=4 (default)
93.99
37.34
80.14
G=6
93.29
37.34
81.08
Table 4: LKAT-B ImageNet-100 ablations ( 2242 , 300 epochs). Bold denotes the best Acc@1 in each block.
Setting
CIFAR-10
CIFAR-100
RA , 2242
93.01
78.20
RB , 2242
93.62
78.28
RB , 3842
95.45
78.63
Table 5: CIFAR fine-tuning of ImageNet-100 LKAT-B.
Setting
RA
RB
Shared
Init
ImageNet-100 pretrained encoder
Optimizer / peak LR
AdamW / 10−4
Resolution / aug.
2242 / 3842 / DeiT
Mixup / CutMix
0.8 / 1.0
Differing
Table 6: CIFAR fine-tuning recipe ( RA vs. RB ).
Figure 13
Require: x,c,w,d as in Alg. 1 ;
upstream dy∈RN
Ensure: dx∈RN , dw∈RG
1:
inv←1/d ; inv2←1/d2 ; dw←0
2:
for each program p in parallel do
3:
Load tile x[i] , dy[i]
4:
jac←0
Algorithm 2 Fused RBF–grid linear backward
Figure 8: End-to-end wall-clock speedup of the fused Triton RBF–grid linear operator versus an unfused torch baseline for a joint forward+backward pass, as a function of channel width D . Bars report Timetorch/TimeTriton at grid sizes G∈{4,8,16} (dotted line =1 ). Unlike peak-memory reduction, which is nearly constant in D , e2e speedup grows with both D and G : from about 1.4 – 1.7× at D=128 up to about 17× / 29× / 44× at D=4096 for G=4/8/16 . Larger grids amplify the benefit of avoiding the materialised […,D,G] intermediate, so the fused kernel’s latency stays nearly flat while torch scales with G .
Figure 9: Learned RBF–KAN edge functions ϕp,q(x) for ImageNet-100 LKAT-B, first KAN layer (fc1), blocks 0 – 11 (panel labels).
Figure 10: Learned RBF–KAN edge functions ϕp,q(x) for ImageNet-100 LKAT-B, second KAN layer (fc2), blocks 0 – 11 (panel labels).
Institute of Advanced Intelligence and Computing (IAIC), A*STAR, Singapore · Nanyang Technological University · ST Engineering Geo-Insights, Singapore +3