🤖 AI Summary
Autoregressive language models generate text token-by-token, which limits parallelization and hinders efficient output of highly predictable subsequent tokens. This work proposes MARS, a method that enables instruction-tuned models to predict multiple tokens in a single forward pass through lightweight continued training, without altering the model architecture or increasing parameter count. MARS maintains compatibility with the original inference interface and allows runtime adjustment of the trade-off between generation speed and quality. It integrates block-level KV caching, confidence-threshold-based control, and multi-token prediction training. Experiments show that MARS matches or exceeds baseline performance in single-token generation while preserving accuracy in multi-token prediction, achieving 1.5–1.7× higher throughput and up to 1.71× faster inference on Qwen2.5-7B.
📝 Abstract
Autoregressive (AR) language models generate text one token at a time, even when consecutive tokens are highly predictable given earlier context. We introduce MARS (Mask AutoRegreSsion), a lightweight fine-tuning method that teaches an instruction-tuned AR model to predict multiple tokens per forward pass. MARS adds no architectural modifications, no extra parameters, and produces a single model that can still be called exactly like the original AR model with no performance degradation. Unlike speculative decoding, which maintains a separate draft model alongside the target, or multi-head approaches such as Medusa, which attach additional prediction heads, MARS requires only continued training on existing instruction data. When generating one token per forward pass, MARS matches or exceeds the AR baseline on six standard benchmarks. When allowed to accept multiple tokens per step, it maintains baseline-level accuracy while achieving 1.5-1.7x throughput. We further develop a block-level KV caching strategy for batch inference, achieving up to 1.71x wall-clock speedup over AR with KV cache on Qwen2.5-7B. Finally, MARS supports real-time speed adjustment via confidence thresholding: under high request load, the serving system can increase throughput on the fly without swapping models or restarting, providing a practical latency-quality knob for deployment.