Flash PD-SSM: Memory-Optimized Structured Sparse State-Space Models
Abstract
State-space models (SSMs) face a fundamental trade-off between efficiency and expressivity that is mainly dictated by the structure of the model's transition matrix. Unstructured transition matrices enable maximal expressivity, as measured by their ability to model finite-state automaton (FSA) transitions, but come at a prohibitively high compute and memory cost. In contrast, diagonal and other structured transition matrices suffer from limited expressivity, but are highly efficient both in runtime and memory consumption. Building on recent work on structured sparse SSMs, we propose Flash PD-SSM, a novel SSM that achieves comparable throughput to diagonal models with optimal expressivity guarantees. Flash PD-SSM keeps a collection of structured sparse matrices, from which a single matrix is selected at each time-step, enabling FSA expressiveness at the level of unstructured matrices while maintaining the efficiency required for training models at scale. First, we validate Flash PD-SSM against a suite of alternative models on common mechanistic and synthetic state-tracking tasks, showing that its theoretical expressivity is achieved in practice. Moreover, on multivariate time-series tasks involving sequences of length over 17,000, Flash PD-SSM defines a new state-of-the-art (SoTA) accuracy among competing SSM methods. Second, we demonstrate that Flash PD-SSM is an effective drop-in replacement for hybrid LLMs, yielding improvements both in natural language state-tracking and in common language modeling scenarios. Finally, we show that our highly efficient design results in increased throughput and decreased memory consumption with respect to SoTA SSMs widely used in frontier language models.