FlashFFN: Multi-Head Decomposition Enables I/O-Aware Feed-Forward Network
Abstract
The Feed-Forward Network (FFN) dominates the activation memory of modern Transformers because the intermediate tensor of width dff is materialized in HBM between the two matmul stages. FlashAttention-style I/O-aware tiling avoids the analogous cost in attention, yet has not been brought to FFN. We identify the reason as a hardware-geometry constraint: under one-pass tiled accumulation a fused FFN kernel must keep an output accumulator of width dmodel in on-chip SRAM, which is infeasible whenever dmodel >= 2048 on currently shipping datacenter GPUs. We formalize this as an SRAM-feasibility proposition and prove a corollary: multi-head decomposition is the simplest architectural transform that brings FFN inside this SRAM-feasible regime, because it makes the per-head accumulator width dh = dmodel/H a free parameter that can be driven below the GPU's SRAM cutoff. Naive multi-head FFN, however, blows up the per-head expansion ratio dff/dh as dmodel grows, drifting away from the well-validated SwiGLU 8/3 optimum and degrading quality at scale; we resolve this with a dense MoE-like sub-network correction that restores the optimal ratio while preserving SRAM-feasibility. The resulting architecture, FlashFFN, is an SRAM-feasible FFN block that replaces SwiGLU under matched parameter and FLOPs budgets and derives directly from the SRAM constraint. Across 128M--1.3B parameter models trained on 60--100B tokens of The Pile, FlashFFN achieves a Pareto improvement over SwiGLU at fixed compute: up to ~4x lower peak activation memory at long sequences, inference latency that matches or improves over SwiGLU (up to ~1.32x faster on H100), and consistent quality gains (up to -0.84 in evaluation perplexity and +1.60 average accuracy across six downstream benchmarks at 1.3B)