M$^2$RNN: Non-Linear RNNs with Matrix-Valued States for Scalable Language Modeling
Mayank Mishra ⋅ Shawn Tan ⋅ Ion Stoica ⋅ Joseph Gonzalez ⋅ Tri Dao
Abstract
Transformers are limited to computations in the TC$^0$ complexity class, excluding tasks such as entity tracking and code execution that provably require greater expressivity. Motivated by this, we revisit non-linear Recurrent Neural Networks (RNNs) and introduce Matrix-to-Matrix RNN (M$^2$RNN): an architecture with matrix-valued hidden states and non-linear state transitions. We show that (i) the language modeling gap of non-linear RNNs is primarily a state-size gap, not a non-linearity penalty, and (ii) outer-product state expansion enables efficient tensor-core utilization. Empirically, M$^2$RNN achieves perfect state-tracking generalization at sequence lengths beyond training. These gains transfer to large-scale language modeling: at 7B MoE, Hybrid M$^2$RNN outperforms Hybrid Gated DeltaNetby $0.5$ perplexity points using $3\times$ smaller recurrent states, and replacing even a single recurrent layer with M$^2$RNN matches Hybrid M$^2$RNN accuracy with minimal throughput cost. Hybrid Gated DeltaNetwith a single M$^2$RNN layer also outperforms state-of-the-art hybrid linear-attention architectures by up to $8$ points on LongBench. Together, these results establish non-linear RNN layers as a compelling building block for efficient and scalable language models.
Chat is not available.
Successful Page Load