You Only Cache Once: Decoder-Decoder Architectures for Language Models
Yutao Sun, Li Dong, Yi Zhu, Shaohan Huang, Wenhui Wang, Shuming Ma, Quanlu Zhang, Jianyong Wang, Furu Wei
Introduction
The decoder-only Transformer has become the de facto architecture for language models. Numerous efforts have continued to develop suitable architectures for language modeling. There have been main strands of explorations. First, encoder-only language models, such as BERT , bidirectionally encode the input sequence. Second, encoder-decoder models, such as T5 , use a bidirectional encoder to encode input and a unidirectional decoder to generate output. Both of the above layouts struggle with autoregressive generation due to bidirectionality. Specifically, encoders have to encode the whole input and output tokens again for the next generation step. Although encoder-decoder can use only decoder to generate, the output tokens do not fully leverage the parameters of encoder, especially for multi-turn conversation. Third, decoder-only language models, such as GPT , generate tokens autoregressively. By caching the previously computed key/value vectors, the model can reuse them for the current generation step. The key-value (KV) cache avoids encoding the history again for each token, greatly improving the inference speed. This compelling feature establishes the decoder-only language model as the standard option.
However, as the number of serving tokens increases, the KV caches occupy a lot of GPU memory, rendering the inference of large language models memory-bounded . For the example of a 65B-size language model (augmented with grouped-query attention and 8-bit KV quantization), 512K tokens occupy about 86GB GPU memory, which is even larger than the capacity of one H100-80GB GPU. In addition, the prefilling latency of long-sequence input is extremely high. For instance, using four H100 GPUs, the 7B language model (augmented with Flash-Decoding and kernel fusion) requires about 110 seconds to prefill 450K tokens, and 380 seconds for 1M length. The above bottlenecks make it difficult to deploy long-context language models in practice.
In this work, we propose a decoder-decoder architecture, YOCO, for large language models, which only caches KV pairs once. Specifically, we stack cross-decoder upon self-decoder. Given an input sequence, the self-decoder utilizes efficient self-attention to obtain KV caches. Then the cross-decoder layers employ cross-attention to reuse the shared KV caches. The decoder-decoder architecture is conceptually similar to encoder-decoder, but the whole model behaves more like a decoder-only model from the external view. So, it naturally fits into autoregressive generation tasks, such as language modeling. First, because YOCO only caches onceThe word “once” refers to global KV cache. Strictly, self-decoder also needs to store a certain number of caches. As the self-decoder utilizes an efficient attention module, the cache size is bounded to a constant, which can be ignored compared to global caches when the sequence length is large., the GPU memory consumption of KV caches is significantly reduced. Second, the computation flow of the decoder-decoder architecture enables prefilling to early exit before entering the self-decoder. The nice property speeds up the prefill stage dramatically, improving user experience for long-context language models. Third, YOCO allows for more efficient system design for distributed long-sequence training. In addition, we propose gated retention for self-decoder, which augments retention with a data-controlled gating mechanism.
We conduct extensive experiments to show that YOCO achieves favorable language modeling performance and has many advantages in terms of inference efficiency. Experimental results demonstrate that YOCO can be scaled up with more training tokens, larger model size, and longer context length. Specifically, we scale up the 3B YOCO model to trillions of training tokens, attaining results on par with prominent Transformer language models, such as StableLM . Moreover, the scaling curves ranging from 160M to 13B show that YOCO are competitive compared to Transformer. We also extend the context length of YOCO to 1M tokens, achieving near perfect needle retrieval accuracy. In the multi-needle test, YOCO obtains competitive results even compared to larger Transformers.
In addition to good performance on various tasks, the profiling results show that YOCO improves the GPU memory footprint, prefill latency, throughput, and serving capacity. In particular, the memory of KV caches can be reduced by about for 65B models. Even for a 3B model, the overall inference memory consumption can be reduced by two times for 32K tokens and by more than nine times for 1M tokens. The prefill stage is speeded up by for the 1M context and for the 32K input. For example, for a 512K context, YOCO reduces the Transformer prefilling latency from 180 seconds to less than six seconds. The results position YOCO as a strong candidate model architecture for future large language models with native long-sequence support.
You Only Cache Once (YOCO)
Both self- and cross-decoder follow a similar block layout (i.e., interleaved attention and feed-forward network) as in Transformer . We also include pre-RMSNorm , SwiGLU , and grouped-query attention as improvements. The difference between the two parts lies in attention modules. Self-decoder (Section 2.1) uses efficient self-attention (e.g., sliding-window attention). In comparison, cross-decoder (Section 2.2) uses global cross-attention to attend to the shared KV caches produced by the output of the self-decoder.
Self-decoder takes token embeddings as input and compute intermediate vector representation :
where represents efficient self-attention, , and RMSNorm is used for . Causal masking is used for efficient self-attention.
The key property of the efficient self-attention module is inference memory, i.e., constant number of KV caches. For example, the cache size of sliding-window attention depends on the window size instead of the input length. More design choices (e.g., gated retention) of the efficient self-attention module are detailed in Section 3.
2 Cross-Decoder
First, the output of the self-decoder generates global KV caches for cross-decoder:
3 Inference Advantages
In addition to competitive language modeling results, YOCO significantly reduces serving costs and improves inference performance. We report detailed inference comparisons in Section 4.4.
Saving GPU Memory and Serving More Tokens. Table 3 compares the memory complexity between Transformers and YOCO. Specifically, because global KV caches are reused and efficient self-attention needs constant caches, the number of caches is , where is the input length, is a constant (e.g., sliding window size), and is the number of layers. For long sequences, is much smaller than , so about caches are required, i.e., you only cache once.
In comparison, Transformer decoders have to store keys and values during inference. So YOCO roughly saves times GPU memory for caches compared to Transformer decoders. Because the inference capacity bottleneck becomes KV caches (Figure 6(b)), our method enables us to serve many more tokens without being out of GPU memory. The increased batch size is also beneficial to inference throughput.
Reducing Prefilling Time and Improving Throughput. As shown in Table 3, because the cross-decoder reuses the outputs of self-decoder, we can exit early before entering the cross-decoder during the prefill stage. The intriguing property of computation dependency greatly accelerates the prefilling speed.
First, only half the layers are needed for forward computation, i.e., at least half prefilling latency reduction. Second, the efficient attention modules of the self-decoder are usually fast. For the example of 512K context length, we can decrease the prefilling latency from 180 seconds (Transformer with optimized inference, such as Flash-Decoding and kernel fusion) to less than 6 seconds (Figure 8). Even for 32K length, YOCO has about three times speedup in terms of prefilling time. Table 3 compares prefilling time complexity of attention modules between Transformer and YOCO.
Design Choices of Self-Decoder
We can choose various efficient self-attention methods for self-decoder. As long as the module only requires constant inference memory, the cache memory complexity of the self-decoder depends on the number of layers. Moreover, a good module choice improves both training and deployment costs. In this work, we use gated retention (Section 3.1) or sliding-window attention (Section 3.2).
Gated retention (gRet, aka gRetNet or RetNet-3) augments retention with a data-dependent gating mechanism, which achieves training parallelism, good performance, and low inference cost simultaneously for sequence modeling. We use gRet as the default efficient self-attention module in the experiments. The method unifies the parallel, recurrent, and chunkwise recurrent computation paradigms. These three representations are equivalent and can obtain the same computation results. The training process usually uses the parallel or chunkwise recurrent paradigms, while the inference stage can employ the recurrent paradigm for constant KV memory. We describe the three representations as follows:
The Parallel Representation The gated retention is defined as:
The Recurrent Representation Being equivalent to Equation 4, the output of gated retention can be computed recurrently. For the -th timestep, the output is obtained via:
where are the same as in Equation 4. During auto-regressive inference, the self-decoder maintains as the intermediate state for an efficient generation.
The Chunkwise Recurrent Representation The chunk-wise representation is a unified formulation of recurrent and parallel representations. Given chunk size , the outputs are computed chunk by chunk. The computation is divided into inner-chunk and cross-chunk parts. Denote as the -th chunk, i.e., , we compute the -th chunk as:
where is the intermediate state of the -th chunk, and summarizes the data-controlled decay . The proof in Appendix B shows the equivalence between the computation paradigms. The chunkwise paradigm combines the best of parallelism and recurrence, i.e., saving FLOPs compared with fully parallel computation and reducing the iterations compared to recurrent computation. During the training and prefill stages, the chunk-wise representation increases throughput and reduces GPU memory consumption.
Multi-Head Gated Retention Similar to multi-head attention and multi-scale retention , we apply gated retention to each head and combine the outputs together:
2 Sliding-Window Attention
Sliding-window attention restricts the attention range into a fixed window size . In contrast, vanilla Transformer decoders attend to all previous tokens. During inference, the KV cache memory complexity can be reduced from to , i.e., the memory usage is constant rather than increasing with sequence length. Similar to multi-head self-attention , we compute the output of sliding-window attention via:
Experiments
We evaluate YOCO for large language models from the following perspectives. First, we follow the setting of StableLM-3B-4E1T to scale up training tokens (Section 4.1). Second, we present the scaling curves of the proposed architectures (Section 4.2). Third, we scale up the YOCO model to 1M context length and evaluate its long-sequence modeling capability (Section 4.3). Fourth, we analyze the deployment advantages, including GPU memory footprint, serving capacity, prefilling time, and throughput (Section 4.4). Experimental results show that YOCO achieves competitive performance across various evaluation metrics. More importantly, the proposed method significantly reduces the inference cost.
We train a 3B-size YOCO language models by scaling up the number of training tokens. Then we compare the checkpoints with strong Transformer-based language models.
Setup We use a similar training recipe as in StableLM-3B-4E1T . We adjust the head dimension to 128 instead of 80 as in StableLM for better kernel support. In order to keep the model size unchanged, we set the hidden size to 3072 and the number of layers to 26. Grouped-query attention is used, where the number of query heads is 24, and the number of key-value heads is 8. We train YOCO with gated retention (Section 3.1). The non-embedding parameter count is 2.8B. In comparison, StableLM-3B-4E1T is 2.7B and OpenLLaMA-v2-3B is 3.2B. The training sequence length is 4096. The batch size is 4M tokens. We use the AdamW optimizer with . The maximal learning rate is 3.2e-4 with 1000 warmup steps and linear decay to 1.28e-5. The total schedule is set to 5T tokens. We train the model with 400k steps (i.e., 1.6T tokens) given the resource budget. The curated training corpus is similar to . We use tiktoken-cl100k_base as the tokenizer. Detailed hyperparameters are described in Appendix C.
Results Table 4 compares the YOCO checkpoints with OpenLLaMA-v2-3B , StableLM-base-alpha-3B-v2 , and StableLM-3B-4E1T . We use LM Eval Harness to evaluate the zero-shot performance on various downstream tasks. OpenLLaMA-v2-3B and StableLM-base-alpha-3B-v2 are trained with 1T tokens. The intermediate numbers of StableLM-3B-4E1T are taken from its technical report . Experimental results across end tasks indicate that YOCO achieves comparable results with previous well-tuned Transformer language models. Both the checkpoints trained with 1T tokens and 1.6T tokens obtain consistent trend. Moreover, the results show that YOCO is scalable in terms of training tokens.
2 Scalability Compared with Transformers
We compare the scaling curves between Llama Transformer , YOCO with gated retention (YOCO; Section 3.1), and YOCO with sliding-window attention (YOCO; Section 3.2). We train language models of various sizes (i.e., 160M, 400M, 830M, 1.4B, 2.7B, 6.8B, and 13B) using the same training data and settings. The validation loss is used as the evaluation metric. The scaling law is supposed to extrapolate larger-size performance.
Setup We augment the Transformer architecture with Llama improvements, such as RMSNorm , SwiGLU , and removing bias. The sliding window size of YOCO is 1,024. We align the number of parameters by adjusting the FFN intermediate dimension. The training batch size is 0.25M tokens with a 2k sequence length. We train the models with 40k steps, i.e., 10B tokens. In practice, we find that the setting is effective for loss convergence, and the scaling laws can be well-fitted. More hyperparameters are detailed in Appendix D.
Results Figure 3 reports the validation loss with various parameter counts. We also fit the scaling curves as in . YOCO obtains comparable performance from 160M to 13B compared to the Llama-optimized transformer architecture. The findings demonstrate that YOCO scales effectively with respect to model size. Moreover, YOCO outperforms Transformer and YOCO. The gains come from hybrid architectures of attention and retention, whose inductive biases tend to be complementary to each other. We observed similar gains by interleaving the attention and retention modules (1:3). Recent hybrid architectures also confirm similar findings.
3 Long-Context Evaluation
We extend the context length of YOCO-3B (Section 4.1) to 1M tokens. We evaluate long-context models on needle retrieval and language modeling tasks.
We continue the model training with longer lengths progressively. The length schedule is 64K, 256K, and 1M tokens. The batch size is kept the same as before. The learning rate and RoPE are set as in Table 8. Training data is up-sampled according to sequence length . For a fair comparison, we do not use long-instruction tuning data. More training details are described in Appendix E. A chunk parallelism algorithm for YOCO is proposed in Appendix A, which reduces communication overhead and GPU memory fragmentation in our experiments of 1M length.
Needle In A Haystack The pressure test evaluates whether models can retrieve “needles” from a long document . We follow the evaluation setting of Gemini 1.5 and LWM . The needles are constructed as a city with a magic number. We run 10 times at the same depth and length. The averaged accuracy is reported. Figure 4 shows that YOCO-3B-1M passes the Needle-In-A-Haystack test with near perfect accuracy. The results indicate that YOCO has strong long-context modeling capability.
Multi-Needle Retrieval Besides the above single-needle retrieval, we conduct a multi-needle evaluation. We compare YOCO-3B-1M with previous long-context language models, including MiniCPM-128K , ChatGLM3-128K , YaRN-Mistral-128K , and LWM-1M-text . The evaluation is conducted in 128K sequence length, because most previous models are tuned with this length.
Table 5 reports the accuracy with needles. Among these models, LWM-1M-text and YOCO-3B-1M are trained with a 1M context length, while the others are in 128K length. Although LWM-1M-text continues training of Llama-2-7B, YOCO-3B-1M can still achieve comparable performance with half the model size. Moreover, the 7B-size YaRN-Mistral-128K obtained by postion interpolation lags behind the other models. Compared to MiniCPM-128K and ChatGLM3-128K, YOCO-3B-1M also outperforms these well-trained language models.
Perplexity over Long Sequences Figure 5 shows the cumulative average negative log-likelihood (NLL) as a function of context length. We evaluate both book and repository-level code data. We follow the setting of and filter validation data that are longer than 1M tokens. NLL decreases consistently with longer sequence length. The results indicate that YOCO can effectively utilize long-distance dependency for language modeling. We also observe that the NLL-length curves tend to fit the power law, where the gaps are affected by the noise within the validation examples.
4 Inference Advantages
We analyze inference efficiency from various perspectives, such as GPU memory footprint, prefilling latency, throughput, and serving capacity. We demonstrate that YOCO reduces the deployment cost by orders of magnitude, especially for long-sequence inference. More importantly, the user experience (such as latency) is improved while maintaining good performance and reducing expenses.
We compare YOCO with Transformer. The default model configuration follows Section 4.1. Notice that Transformer uses grouped-query attention , Flash-Decoding , and kernel fusion for a fair comparison. As described in Section 3.1, gated retention uses the chunk-recurrent representation in the prefill stage, and the recurrent representation in the generation stage. The chunk size is set to 256. We implement a Triton kernel for gated retention. The evaluation sequence length is ranging from 32K to 1M. The last 1,024 tokens are supposed to be generated, while the previous tokens are given input context. The experiments are conducted with H100-80GB GPU cards.
GPU Memory The inference memory consumption is made up of three parts, namely model weights, intermediate activation, and KV cache. Figure 6(b) presents the breakdown memory profiling results. Along with an increase in context length, the main memory bottleneck becomes KV caches, while model weights consume constant memory. The results show that YOCO alleviates the activation cost and KV cache memory footprint.
As shown in Figure 6(a), the memory cost is significantly reduced using YOCO. Moreover, the memory consumption of YOCO increases slowly along the sequence length. For example of 1M length, the overall inference memory usage is only 12.4GB, while Transformers occupy GPU memory. YOCO makes it feasible to deploy long-sequence modeling on customer-level GPUs. Even with a 32K sequence length, YOCO requires about less memory than Transformer. Although we compare 3B-size models here, the reduction ratio becomes larger as the number of layers increases.
Figure 7 reports the GPU memory consumption of KV cache for each token. As YOCO only caches one layer of global key-value pairs, it needs roughly times fewer memory compared to Transformer. For example, YOCO can serve 128K tokens with 1GB GPU memory, while Transformer with GQA can only support 1.6K tokens at 65B model size.
Prefilling Latency In the prefill stage, the model encodes input tokens in parallel. As shown in Figure 8, the prefilling latency is a pain point of user experience for long-context models. For 512K- and 1M-length input sequences, Transformer needs about 180 seconds and 300 seconds, respectively. The computational complexity of Transformer is , which requires a large number of FLOPs for long context. In contrast, YOCO’s prefilling time is , growing linearly (Section 2.3) along the sequence length.
Figure 8 shows that YOCO reduces the Transformer prefilling time from 180 seconds to less than 6 seconds for 512K context. As described in Section 2.3, the prefill stage can early exit before entering cross-decoder. So, there is at least two times speedup of prefilling latency even for short context. For example, YOCO is faster than Transformer for 32K length.
Throughput The throughput indicates how many tokens the model can process per second, involving both pre-filling and generation time. Figure 9 shows that YOCO achieves higher throughput across context lengths compared to Transformer. For the example of 512K queries, Transformer’s throughput is 4.5 token/s while YOCO reaches 43.1 token/s, i.e, achieving speedup. The throughput is improved for the following reasons. First, YOCO decreases the time required for prefilling as previously demonstrated. Second, as the memory consumption is reduced, we can use larger batch size for inference, which also contributes to the throughput improvement.
Conclusion
In this work, we propose a decoder-decoder architecture (YOCO) for large language modeling. YOCO achieves significantly better inference efficiency and competitive performance compared with Transformers. Experimental results demonstrate that YOCO achieves favorable results for large language models under various settings, i.e., scaling up number of training tokens, scaling up model size, and scaling up context length to 1M tokens. Profiling results also show that YOCO improves inference efficiency by orders of magnitude, especially for long-sequence modeling.
The work can be advanced from the following perspectives:
YOCO + BitNet + Groq. Groq achieves very high throughput by putting all things within SRAM. However, the memory capacity bottleneck limits the model size and input token count. Now, hundreds of chips are connected to host just one model. As a solution, YOCO reduces KV cache memory, and BitNet reduces model weight memory. The LLM deployment cost is expected to be reduced by orders of magnitude using the above combination.
YOCO for Multimodal Large Language Models. The YOCO layout is general to the use of multiple self-decoders. The cross-attention layers are natural for multimodal fusion . The causal dependency of self-decoders also perfectly fits in streaming video. The async multimodal large language models can avoid different data steams block each other, which is critical for real-time applications, such as robotics.
Optimized Mechanism for KV Cache Module. Figure 2 explicitly highlights KV cache, which opens up new opportunities to develop native memory mechanisms. First, we can integrate a cache compression mechanism to obtain more compact memory. Second, we can build an index for efficient key-value retrieval. As YOCO reuses caches, it enables us to maintain only one index rather than creating an index for each layer. Third, the disentangled modeling supports pre-caching context, which is potentially useful for native RAG and LLM-native search engines.
Acknowledgement
We would like to acknowledge Ben Huntley for maintaining the GPU cluster. The long-sequence training utilizes CUBE, which is an internal version of . We implement the Triton kernel of gated retention based on FLA .
References
Appendix A Chunk Parallelism for Long-Sequence Training of YOCO
We introduce chunk parallelism for YOCO to reduce the communication frequency, accelerating long-sequence training. Dividing long sequences into different devices is essential when the training length is extremely long . However, the overall throughput tends to be bounded by GPU communication . Cross-decoder disentangles self-attention dependency while preserving modeling capability, bringing intriguing advantages to distributed long-sequence training.
In self-decoder, the dependency only exists in the adjacent devices. For example, gated retention only requires the hidden state in LABEL:eq:gret:recurrent, and sliding-window attention attends to tokens within the context window. Therefore, the communication amount of self-decoder is relatively small. In the cross-decoder, the all-gather operation is only triggered once for the KV cache, rather than communicating in each layer. The hardware-friendly architecture gives more flexibility to distributed long-sequence training.
Appendix B Chunk-wise Representation of Gated Retention
We illustrate the equivalence between recurrent representation and chunkwise recurrent representation of gated retention. For the output , can be split as where is the chunk size:
where , , , indicates the -th chunk, i.e., . is written as a recurrent function:
Denote as the -th chunk, i.e., , , , We concatenate the output in a block together:
Finally, we show that the chunkwise recurrent representation of gated retention is equivalent to the other two representations.
Appendix C Hyperparameters for YOCO-3B
We describe the hyperparameters used for Section 4.1. The hidden dimension is set to 3072. The number of layers is 26. The number of query heads is 24, and the number of key/value heads is 8 with grouped-query attention . The total number of parameters without embedding is 2.83B. The training batch size is 4M tokens. We use 4096 training length. The optimizer is AdamW with . The learning rate is with 1000 warmup steps. We set a 5T-token learning rate schedule with linear decay to .
Appendix D Hyperparameters for Scaling Curves
We describe the hyperparameters used for Section 4.2. Table 7 reports the hidden dimension, number of layers, and number of heads used for different model sizes. The head dimension of gated retention is set to 256. To align the number of parameters, the FFN size for Transformer is while the FFN size for YOCO is . The training length is set to 2048. The batch size is set to 0.25M tokens. We use the AdamW optimizer with . The learning rate is for 160M to 1.4B sizes and for 2.7B to 13B sizes. The warmup step is 375 with linear rate decay. The weight decay is set to 0.05. We train the models with 40k steps, i.e., 10B tokens.
Appendix E Hyperparameters for Length Extension
We progressively extend the context length to 1M tokens in Section 4.3. The length schedule is 64K, 256K, and 1M. We up-sample the documents that are longer than the training length. Table 8 shows that we use different RoPE and learning rate for each stage.
Appendix F Pseudo Code of Gated Retention
We present pseudocode for the three computation paradigms of gated retention (Section 3.1). Parallel implementation enables training parallelism to fully utilize GPUs. The recurrent paradigm enables low-cost inference. Chunkwise retention combines the above advantages (i.e., parallel within each chunk and recurrent across chunks), which has linear memory complexity for long sequences.
Appendix G Comparisons with Transformer Variants
Table 9 reports the validation perplexity for language modeling. Following Zoology , we divide the perplexity into Ar-Hit, where the predicted token is a bigram previously seen in the previous context, and First-Occur, where the predicted token cannot be recalled from the context.
G.2 Long-Context Evaluation
We evaluate the long-context modeling for the above architectures on four tasks of the ZeroSCROLLS benchmark. We continue training the 160M models in Table 9 as long-context models. Specifically, we further train the models with 2B tokens in 16,384 length. The rotation base scaling is also used for length extension. For sparse Transformer, we keep the 2,048 context window and do not change the rotation base (i.e., RoPE ).
Figure 11 reports the perplexity of the answers with different input lengths. Among all these architectures, YOCO and Transformer consistently perform better than others across tasks and lengths.