🤖 AI Summary
This study addresses the substantial memory overhead in Swin Transformers caused by window attention, which impedes simultaneously achieving large receptive fields and high token density. To overcome this limitation, we introduce FlashAttention into the window attention architecture for the first time, reducing memory complexity to O(M²). Furthermore, local learnable 2D rotary position embeddings (RoPE) are incorporated to effectively recover positional information. This approach transcends the conventional memory bottleneck of Swin Transformers, enabling efficient processing of dense tokens within large windows. Experimental results demonstrate that the Tiny model achieves 84.1% accuracy on ImageNet and 44.1 AP on COCO, significantly outperforming SwinV2 while maintaining constant training memory consumption.
📝 Abstract
High-resolution vision backbones have long been forced to trade away local token density to afford larger receptive fields. Hierarchical Swin transformers impose this compromise because standard windowed attention materializes an $M^2\times M^2$ score matrix per window, incurring $O(M^4)$ memory as windows or token grids grow. Furthermore, Swin adds a learned relative-position bias elementwise to attention scores, requiring full materialization of the score matrix and its gradient. This keeps Swin and SwinV2 trapped in a small-window($M=8,16$), coarse-token regime with patch size $4\times4$ ($p=4$), limiting performance for fine-grained tasks. We introduce FLASHSWIN, which replaces standard windowed attention with a FlashAttention implementation that computes exact softmax attention without materializing the score matrix, reducing per-window memory from $O(M^4)$ to $O(M^2)$. This enables higher token density and larger receptive fields without inflating memory overhead. Training memory is flat across window sizes: at a $32\times32$ window, FLASHSWIN-T requires only $12.4$\,GB, unchanged from $8\times8$, compared to $70/90$\,GB for SwinV2/V1-T. However, applying FlashAttention directly to Swin creates a trade-off: bypassing the score matrix precludes Swin's additive relative-position bias, forfeiting spatial information in exchange for memory efficiency. FLASHSWIN restores position information as window-local learnable 2D RoPE, making large windows and dense token grids both affordable and accurate. At matched scale, FLASHSWIN-T outperforms Swin variants. With dense tokens and wide windows ($p=2,M=32$), the same Tiny model reaches $84.1\%$ ImageNet-1K, $44.1$ COCO box AP, and $47.28$ ADE20K mIoU---gains of $+1.3$, $+5.1$, and $+1.82$ over SwinV2-T at $M=16$, respectively. At fixed $M=32$, halving the patch size yields roughly $3\times$ larger gains in boundary quality than in mIoU.