TensorX
返回文献探索

Paper · arXiv 2506.05229

Diagonal Batching Unlocks Parallelism in Recurrent Memory Transformers for Long Contexts

Danil Sivtsov, Ivan Rodkin, Gleb Kuzmin, Yuri Kuratov, Ivan Oseledets

38 upvotesJune 5, 2025arXiv 预印本
AI 摘要

Diagonal Batching enables parallel inference in Recurrent Memory Transformers, significantly improving speed and efficiency for long-context tasks.

Transformer modelslong-context inferencequadratic time complexitylinear memory complexityRecurrent Memory TransformersRMTsmemory update mechanismsequential executionDiagonal Batchingrun-time computation reorderingparallelismGPU inferenceLLaMA-1B ARMT modelfull-attention LLaMA-1Binference costlatency

Abstract

Transformer models struggle with long-context inference due to their quadratic time and linear memory complexity. Recurrent Memory Transformers (RMTs) offer a solution by reducing the asymptotic cost to linear time and constant memory usage. However, their memory update mechanism leads to sequential execution, causing a performance bottleneck. We introduce Diagonal Batching, a scheduling scheme that unlocks parallelism across segments in RMTs while preserving exact recurrence. This approach eliminates the sequential constraint, enabling efficient GPU inference even for single long-context inputs without complex batching and pipelining techniques. Because the technique is purely a run-time computation reordering, existing RMT models adopt it with no retraining. Applied to a LLaMA-1B ARMT model, Diagonal Batching yields a 3.3x speedup over standard full-attention LLaMA-1B and a 1.8x speedup over the sequential RMT implementation on 131,072-token sequences. By removing sequential bottleneck, Diagonal Batching reduces inference cost and latency, thereby strengthening RMTs as a practical solution for real-world, long-context applications.

北京市昌平区探索星信息技术及软件开发工作室

京ICP备2026059466号
Diagonal Batching Unlocks Parallelism in Recurrent Memory Transformers for Long Contexts | TensorX