🤖 AI Summary
The Kolmogorov-Arnold Transformer (KAT) exhibits only 1/123 the training speed of a standard Transformer despite comparable FLOPs, with its primary bottleneck residing not in computation but in memory-bound stalls during backpropagation—specifically, inefficient gradient accumulation in GR-KAN layers.
Method: We propose a kernel function reconstruction framework: designing lightweight atomic addition operations to eliminate redundant slow-memory accesses and optimizing gradient computation paths to reduce memory pressure.
Contribution/Results: Our approach accelerates KAT training by 86.5×, significantly mitigates coefficient gradient rounding errors, and achieves practical training efficiency approaching that of standard Transformers. This work constitutes the first systematic identification and resolution of KAT’s memory-level performance bottleneck, establishing a new paradigm for efficient implementation of higher-order activation functions in neural architectures.
📝 Abstract
The Kolmogorov-Arnold Network (KAN) has been gaining popularity as an alternative to the multi-layer perceptron (MLP) with its increased expressiveness and interpretability. However, the KAN can be orders of magnitude slower due to its increased computational cost and training instability, limiting its applicability to larger-scale tasks. Recently, the Kolmogorov-Arnold Transformer (KAT) has been proposed, which can achieve FLOPs similar to the traditional Transformer with MLPs by leveraging Group-Rational KAN (GR-KAN). Unfortunately, despite the comparable FLOPs, our characterizations reveal that the KAT is still 123x slower in training speeds, indicating that there are other performance bottlenecks beyond FLOPs. In this paper, we conduct a series of experiments to understand the root cause of the slowdown in KAT. We uncover that the slowdown can be isolated to memory stalls and, more specifically, in the backward pass of GR-KAN caused by inefficient gradient accumulation. To address this memory bottleneck, we propose FlashKAT, which builds on our restructured kernel that minimizes gradient accumulation with atomic adds and accesses to slow memory. Evaluations demonstrate that FlashKAT can achieve a training speedup of 86.5x compared with the state-of-the-art KAT, while reducing rounding errors in the coefficient gradients. Our code is available at https://github.com/OSU-STARLAB/FlashKAT.