Tensor Product Attention Is All You Need
Yifan Zhang, Yifeng Liu, Huizhuo Yuan, Zhen Qin, Yang Yuan, Quanquan Gu, Andrew Chi-Chih Yao
Introduction
Large language models (LLMs) have revolutionized natural language processing, demonstrating exceptional performance across tasks (Brown et al., 2020; Chowdhery et al., 2023; Touvron et al., 2023; Bubeck et al., 2023). As these models evolve, their ability to process longer contexts becomes increasingly important for sophisticated applications such as document analysis, complex reasoning, and code completions. However, managing longer sequences during inference poses significant computational and memory challenges, particularly due to the storage of key-value (KV) caches (Zhang et al., 2023c; Liu et al., 2024c). Because memory consumption grows linearly with sequence length, the maximum context window is limited by practical hardware constraints.
A variety of solutions have been explored to address this memory bottleneck. Some approaches compress or selectively prune cached states through sparse attention patterns (Child et al., 2019) or token eviction strategies (Zhang et al., 2023c; Xiao et al., 2024; Ribar et al., 2024), though such methods risk discarding tokens that may later prove important. Other work proposes off-chip storage of key-value states (He & Zhai, 2024), at the expense of increased I/O latency. Attention variants like multi-query attention (MQA) (Shazeer, 2019) and grouped-query attention (GQA) (Ainslie et al., 2023) reduce per-token cache requirements by sharing keys and values across heads, but often compromise flexibility or require significant architectural modifications. Meanwhile, low-rank weight factorization methods such as LoRA (Hu et al., 2022) effectively reduce fine-tuning memory, yet do not address the KV cache overhead that dominates runtime. The recently introduced Multi-head Latent Attention (MLA) in Deepseek-V2 (Liu et al., 2024a) caches compressed key-value representations but needs additional position-encoded parameters per head due to incompatibility with Rotary Position Embedding (RoPE) efficiently (Su et al., 2024b).
In order to overcome the limitations of existing approaches, we introduce Tensor Product Attention (TPA), as illustrated in Figure 1, a novel architecture that uses higher-order tensors to factorize queries (Q), keys (K), and values (V) during attention computation. By dynamically factorizing activations rather than static weights (e.g., LoRA), TPA constructs low-rank, contextual representations that substantially reduce KV cache memory usage with improved representational capacity. In practice, TPA can reduce the memory overhead by an order of magnitude compared to standard multi-head attention (MHA) with lower pretraining validation loss (perplexity) and improved downstream performance.
A key advantage of TPA is its native compatibility with rotary positional embeddings (RoPE) (Su et al., 2024b), enabling a straightforward drop-in replacement for multi-head attention (MHA) layers in modern LLM architectures such as LLaMA (Touvron et al., 2023) and Gemma (Team et al., 2024).
Our primary contributions are summarized as follows:
We propose Tensor Product Attention (TPA), A mechanism that factorizes , , and activations using contextual tensor-decompositions to achieve or more reduction in inference-time KV cache size relative to standard attention mechanism (Vaswani et al., 2017) with improved performance compared to previous methods such as MHA, MQA, GQA, and MLA. In addition, we unify existing attention mechanisms by revealing that MHA, MQA, and GQA all arise naturally as non-contextual variants of TPA.
We propose Tensor ProducT ATTenTion Transformer (T6), a new TPA-based model architecture for sequence modeling. On language modeling experiments, T6 consistently improves validation perplexity and downstream evaluation performance with reduced KV cache size.
We show TPA integrates seamlessly with RoPE (Su et al., 2024b), facilitating easy adoption in popular foundation model architectures such as LLaMA and Gemma.
Background
In this section, we review several classical forms of attention: Scaled Dot-Product Attention, Multi-Head Attention (MHA) (Vaswani et al., 2017), Multi-Query Attention (MQA) (Shazeer, 2019), and Grouped Query Attention (GQA) (Ainslie et al., 2023), as well as Rotary Position Embedding (RoPE, Su et al. (2024b)). We also introduce a recent method called Multi-head Latent Attention (MLA) used in DeepSeek-V2 (Liu et al., 2024a) and DeepSeek-V3 (Liu et al., 2024b).
where is the -th element of .
Scaled dot-product attention (Vaswani et al., 2017) determines how to focus on different parts of an input sequence by comparing queries () and keys (). It produces a weighted combination of the values (). Formally, the attention output is:
where each of is an matrix for tokens and key dimension . The division by stabilizes training by controlling the scale of the inner products.
2 Multi-Head Attention (MHA)
MHA can capture a rich set of dependencies while each head focuses on different subspaces.
3 Multi-Query Attention (MQA)
By sharing these key and value projections, MQA cuts down on memory usage (especially for the key-value cache in autoregressive inference) but loses some expressivity since all heads must rely on the same key/value representations.
4 Grouped Query Attention (GQA)
Grouped Query Attention (GQA) (Ainslie et al., 2023) generalizes MHA and MQA by grouping heads. Specifically, we partition the total heads into groups. Each group has a single set of keys and values, but each individual head within that group still retains its own query projection. Formally, if maps a head to its group index , then:
By adjusting between and , GQA can interpolate between sharing all key/value projections across heads (i.e., MQA) and having one set of projections per head (i.e., MHA).
5 Rotary Position Embedding (RoPE)
which ensures that relative positions are preserved, thereby providing a form of translation invariance in the rotary position embedding.
6 Multi-head Latent Attention (MLA)
Below, we briefly outline the Multi-head Latent Attention (MLA) approach used by DeepSeek-V2 (Liu et al., 2024a) and DeepSeek-V3 (Liu et al., 2024b). MLA introduces a low-rank compression of the keys and values to reduce the Key-Value (KV) caching cost at inference.
MLA also compresses the queries, lowering their training-time memory footprint:
Given compressed queries, keys, and values, the final attention output for the -th token is:
Different from (2.2), acceleration by pre-computing [\bm{W}^{DQ}_{i}\bm{W}^{UQ}_{i}{\color[rgb]{0,0,1}\definecolor[named]{pgfstrokecolor}{rgb}{0,0,1}\mathbf{T}_{t-s}}(\bm{W}^{UK}_{i})^{\top}] fails since it varies for different position pairs. Therefore, MLA adds the additional part with a relatively smaller size for RoPE compatibility. In Section 3.2, we will show that TPA addresses the issue of RoPE-incompatibility by applying tensor product.
Tensor Product Attention
In this section, we provide a detailed description of our proposed Tensor Product Attention (TPA), which allows contextual low-rank factorization for queries, keys, and values. First, we explain how TPA factorizes queries, keys, and values with explicit tensor shapes. Next, we describe how TPA can be integrated into the multi-head attention framework and how it reduces memory consumption in KV caching at inference time. Finally, we show how RoPE can seamlessly integrate with TPA (including a pre-rotated variant).
Contextual Factorization (CF). Instead of forming each head’s query, key, or value via a single linear map, TPA factorizes each into a sum of (contextual) tensor products whose ranks are , , and , respectively and may differ. Specifically, for each token , with a small abuse of notation, we define:
Latent Factor Maps. Each factor in the tensor product depends on the token’s hidden state . For example, for queries, we can write:
One often merges the rank index into a single output dimension. For instance, for queries:
Scaled Dot-Product Attention. Once are factorized, multi-head attention proceeds as in standard Transformers. For each head :
Parameter Initialization. We initialize the weight matrices , , , , , using Xavier initialization (Glorot & Bengio, 2010). Specifically, each entry of the weight matrix is drawn from a uniform distribution with bounds , where and are the input and output dimensions of the respective weight matrices. This initialization strategy helps maintain the variance of activations and gradients across the network.
2 RoPE Compatibility and Acceleration
Direct Integration. A useful optimization is to integrate RoPE directly into the TPA factorization. For example, one can pre-rotate the token-dimension factors:
yielding a pre-rotated key representation:
Thus, each is already rotated before caching, removing the need for explicit rotation at the decoding time and accelerating autoregressive inference. Depending on hardware and performance requirements, one can also adopt different RoPE integration approaches for training and inference.
Let be factorized by TPA as
In addition, assume and are factorized by TPA and then rotated by . Let and . Then we have
Focusing on individual heads , the above matrix equality implies:
Theorem 1 indicates that TPA does not break RoPE’s relative translational property. We prove Theorem 1 in Appendix A. In short, acts as a block-diagonal orthogonal transform (i.e., a matrix ) on . Consequently, remains unchanged, while each column of is rotated appropriately, preserving the TPA structure.
3 KV Caching and Memory Reduction
TPA Factorized KV Caching. Instead of storing the full and , TPA stores only their factorized ranks. Specifically, we keep
Compared to the standard caching cost of , the ratio is:
For large and (typically or ), setting (e.g., or ) often yields or more reduction.
4 Unifying MHA, MQA, and GQA as Non-contextual TPA
To match MHA with TPA, let . Focusing on :
so that corresponds to the -th head of .
Substituting (3.8)–(3.9) into (3.1) gives:
Each term \mathbf{e}_{i}\otimes\bigl{(}({\bm{W}^{Q}_{i}})^{\top}\mathbf{x}_{t}\bigr{)} in (3.10) contributes only to the -th row, reconstituting the usual MHA form of . Analogous constructions hold for and using . Thus, MHA is a non-contextual, full-rank variant of TPA.
and similarly for queries/values. This reduces per-token computations and can be effective when head-dimension relationships are relatively stable across all tokens.
MQA and GQA as Non-Contextual TPA. Multi-Query Attention (MQA) (Shazeer, 2019) and Grouped Query Attention (GQA) (Ainslie et al., 2023) also emerge naturally from TPA by restricting the head-dimension factors to be non-contextual and low-rank:
MQA as Rank-1 TPA. In MQA, all heads share a single set of keys/values, corresponding to along the head dimension. Concretely,
forces every head to use the same . Each head retains a distinct query projection, matching the MQA design.
Hence, by constraining TPA’s head-dimension factors to be constant masks (one for MQA; multiple for GQA), these popular variants are recovered as special cases.
5 Other Variants of TPA
and similarly for keys/values. This arrangement is effective if the token-dimension structure remains mostly uniform across the sequence, while the head-dimension factors capture context.
TPA KV Only. One can preserve a standard query mapping,
and factorize only the keys and values. This leaves the query projection as the original linear transformation while reducing memory usage via factorized KV caching.
TPA KV with Shared . Another variant is to share the token-dimension factors of keys and values:
lowering parameter counts and the KV cache footprint. While it constrains and to be formed from the same token basis, it can still perform well and provide additional memory savings.
Nonlinear Head Factors. Rather than applying purely linear mappings to the head-dimension factors , one may introduce element-wise nonlinearities such as or . This effectively yields a Mixture of Heads Attention (MoH Attention), where each component becomes a learned mixture weight modulated by the nonlinearity.
Discussion. These variants illustrate TPA’s versatility in balancing memory cost, computational overhead, and representation power. By choosing which dimensions (heads or tokens) remain contextual and adjusting ranks , TPA unifies multiple existing attention mechanisms—such as MHA, MQA, and GQA—under one framework, while potentially reducing the KV cache size by an order of magnitude during autoregressive inference.
6 Model Architectures
In practice, we merge all ranks into a single dimension of the output, reshape, and sum over rank indices; see Section 3.1 for details. The factorization for K and V follows the same pattern.
Rotary Positional Embedding (RoPE). As discussed in Section 3.2, RoPE (Su et al., 2024b) is applied to the and . Within TPA, we pre-rotate the factor and directly, so that each is already rotated prior to caching, see (3.6) and Theorem 1.
Attention Step and Output Projection. Once we have factorized per token with RoPE applied on and , the attention step proceeds for each head using (3.4). Finally, concatenating these heads and then projecting them back using an output weight matrix gives the final attention result, as shown in (3.5).
SwiGLU Feed-Forward Network. Following Shazeer (2020); Touvron et al. (2023), our T6 uses a SwiGLU-based Feed-Forward Network (FFN):
where is the SiLU (a.k.a., swish) nonlinearity, is element-wise product, and are learnable parameters. Note that other activation functions can also be used.
Overall T6 Block Structure. Putting everything together, one T6 block consists of:
We place norm layers (e.g., RMSNorm) before each sub-layer. Stacking such blocks yields a T6 model architecture with layers.
Experiments
All experiments reported in this paper are implemented on the nanoGPT code base (Karpathy, 2022), using the FineWeb-Edu 100B dataset (Lozhkov et al., 2024). The dataset contains 100 billion tokens for training and 0.1 billion tokens for validation. We compare T6 against the baseline Llama architecture (Touvron et al., 2023) with SwiGLU activation (Shazeer, 2020) and RoPE embeddings (Su et al., 2024a), as well as Llama variants that replace Multi-Head Attention (MHA; Vaswani et al., 2017) with Multi-Query Attention (MQA; Shazeer, 2019), Grouped Query Attention (GQA; Ainslie et al., 2023), or Multi-head Latent Attention (MLA; Liu et al., 2024a). In our experiments, the number of heads is adjusted for each attention mechanism to ensure that all attention mechanisms have the same number of parameters as the standard Multi-Head Attention (MHA), which has parameters per attention layer. We train models at four scales: small (124M parameters), medium (353M), and large (773M). Details on architecture hyperparameters and training hardware appear in Appendix B.1.
Training Setup. We follow the nanoGPT training configuration. In particular, we use the AdamW (Loshchilov, 2017) optimizer with , a weight decay of , and gradient clipping at . We follow the same setting as nanoGPT that the learning rate is managed by a cosine annealing scheduler (Loshchilov & Hutter, 2016) with warmup steps and a (total) global batch size of . For the small, medium, and large models, we set maximum learning rates of , , and (respectively), and minimum learning rates of , , and (respectively).
Training & Validation Curves. Figures 2 and 3 compare training and validation loss curves for the large (773M) and medium (353M) models on FineWeb-Edu-100B. Overall, TPA (red curves) and its simpler variant TPA-KVonly (pink curves) converge as fast as or faster than the baselines (MHA, MQA, GQA, MLA) while also achieving visibly lower final losses. For instance, in Figure 2(b), TPA and TPA-KVonly remain below the MHA baseline in terms of validation loss at nearly all training stages. Meanwhile, Multi-Head Latent Attention (MLA) (Liu et al., 2024a) (blue curves) generally trains more slowly and yields higher losses.
Validation Perplexity. Figure 4 shows the validation perplexities of the medium- and large-scale models. Mirroring the loss curves, TPA and TPA-KVonly steadily outperform MHA, MQA, GQA, and MLA over the course of training. By the end of pretraining (around B tokens), TPA-based approaches achieve the lowest perplexities in most configurations.
Downstream Evaluation. We evaluate zero-shot and two-shot performance on standard benchmarks, including ARC (Yadav et al., 2019), BoolQ (Clark et al., 2019), HellaSwag (Zellers et al., 2019), OBQA (Mihaylov et al., 2018), PIQA (Bisk et al., 2020), WinoGrande (Sakaguchi et al., 2020) and MMLU (Hendrycks et al., 2021), using the lm-evaluation-harness codebase (Gao et al., 2024). For ARC-E, ARC-C, HellaSwag, OBQA, PIQA, and SciQ, we report accuracy norm; for other tasks, we report standard accuracy. Tables 8–9 in the appendix present results for small models; Tables 2–3 for medium models; Tables 4–5 for large models;
For the medium-size (353M) models (Tables 2–3), TPA generally ties or outperforms all competing methods, achieving, for example, an average of 51.41% in zero-shot mode versus MHA’s 50.11%, MQA’s 50.44%, and MLA’s 48.96%. When given two-shot prompts, TPA again leads with 53.12% average accuracy. A similar trend appears for the large-size (773M) models (Tables 4–5), where TPA-KVonly attains the highest average (53.52% zero-shot, 55.33% two-shot), closely followed by full TPA.
Our experiments confirm that TPA consistently matches or exceeds the performance of established attention mechanisms (MHA, MQA, GQA, MLA) across medium and large model scales. The fully factorized TPA excels on mid-scale models, while TPA-KVonly can rival or surpass it at larger scales. In both cases, factorizing the attention activations shrinks autoregressive KV cache requirements by up to –, thus enabling much longer context windows under fixed memory budgets. In summary, tensor product attention provides a flexible, memory-efficient alternative to standard multi-head attention, advancing the scalability of modern language models.
Related Work
Transformers and Attention. As a sequence-to-sequence architecture Transformer (Vaswani et al., 2017) introduced Multi-Head Attention (MHA), enabling more effective capture of long-range dependencies. Subsequent work has explored a variety of attention mechanisms aimed at improving scalability and efficiency, including sparse patterns (Child et al., 2019; Shi et al., 2023; Han et al., 2024; Liang et al., 2024a; Li et al., 2024; Liang et al., 2024b), kernel-based projections (Choromanski et al., 2021), and linearized transformers (Tsai et al., 2019; Katharopoulos et al., 2020; Schlag et al., 2021; Zhang et al., 2023b; Sun et al., 2023; Zhang et al., 2024). To decrease memory usage and circumvent the limitation of memory bandwidth in training, Shazeer (2019) proposed Multi-Query Attention (MQA) where multiple query heads share the same key head and value head. To tackle with the issue of quality degradation and instability in training, Grouped-Query Attention (GQA) (Ainslie et al., 2023) divides queries into several groups, and each group of queries shares a single key head and value head. Recently, DeepSeek-V2 (Liu et al., 2024a) applied multihead latent attention (MLA) to achieve better performance than MHA while reducing KV cache in inference time by sharing the same low-rank representation of key and value. In comparison to the approaches above, TPA applied a low-rank tensor product to compute the queries, keys, and values where the cached representations of keys and values are much smaller than those in MHA, achieving better reduction on memory assumption of KV cache in inference time.
Low-Rank Factorizations. Low-rank approximations have been applied to compress model parameters and reduce complexity including LoRA (Hu et al., 2022), which factorizes weight updates during fine-tuning, and its derivatives for other training scenarios such as efficient pretraining (ReLoRA (Lialin et al., 2023), MoRA (Jiang et al., 2024)), long-context training (LongLoRA (Chen et al., 2024), SinkLoRA (Zhang, 2024)), as well as continual training (InfLoRA (Liang & Li, 2024), GS-LoRA (Zhao et al., 2024), I-LoRA (Ren et al., 2024)). These approaches typically produce static low-rank expansions that do not explicitly depend on the input context. And Malladi et al. (2023); Zeng & Lee (2024) provided theoretical proof of the expressiveness of low-rank approximation. For the initialization of factorization matrices, OLoRA (Büyükakyüz, 2024) applied QR-decomposition of pretrained weight to achieve better performance of language models while LoLDU (Shi et al., 2024) used LDU-decomposition to accelerate training of LoRA. Moreover, AdaLoRA (Zhang et al., 2023a) utilized Singular Value Decomposition (SVD) of the pretrained weight and introduced importance score for each parameter as a measurement to achieve dynamic adjustment of rank. TPA, by contrast, constructs Q, K, and V as contextually factorized tensors, enabling dynamic adaptation.
KV Cache Optimization. During the inference time of Transformers, key and value tensors of the previous tokens are repeatedly computed due to their auto-regressive nature. To enhance efficiency, firstly proposed by Ott et al. (2019), these tensors can be cached in memory for future decoding, referred to as the KV cache. However, the KV cache requires additional memory usage and may add to more latencies due to the bandwidth limitation (Adnan et al., 2024). Therefore, previous studies have explored diverse approaches to mitigate these issues, including KV cache eviction to discard less significant tokens (Zhang et al., 2023c; Xiao et al., 2024; Cai et al., 2024; Adnan et al., 2024), dynamic sparse attention among selected keys and values (Ribar et al., 2024; Tang et al., 2024; Singhania et al., 2024), KV cache offloading to CPU (He & Zhai, 2024; Lee et al., 2024; Sun et al., 2024), as well as quantization of KV cache (Xiao et al., 2023; Liu et al., 2024c; Hooper et al., 2024). Besides these methods, it is also effective to reduce the amount of KV cache for each token, by approaches such as reducing the number of KV heads (Ren et al., 2024; Ainslie et al., 2023), cross-layer KV re-usage (Xiao et al., 2019; Mu et al., 2024; Wu et al., 2024), and low-rank KV representation (Saxena et al., 2024). Different from the methods above, TPA reduces the size of the KV cache by using tensor-decomposed keys and values.
Conclusion
We introduced Tensor Product Attention (TPA), which factorizes query, key, and value matrices into rank- tensor products dependent on the token’s hidden state. Storing only the factorized key/value components during autoregressive decoding substantially decreases the kv memory size with improved performance compared with MHA, MQA, GQA, and MLA. The approach is fully compatible with RoPE (and can store pre-rotated keys). Variants of TPA include factorizing only the key/value or sharing basis vectors across tokens. Overall, TPA offers a powerful mechanism for compressing KV storage while improving the model performance, thereby enabling longer sequence contexts under constrained memory.
References
Appendix A Proofs of Theorems
Because RoPE is a linear orthogonal transform, we can write
where is the block-diagonal matrix encoding RoPE. This allows us to define
Similarly, for the key tensor , we have
Now, consider the product of the rotated queries and keys:
Since and encode positional rotations, the product corresponds to a relative rotation . Therefore, we can express the above as
Focusing on individual heads , the above matrix equality implies:
This equality confirms that the relative positional encoding between queries and keys is preserved under TPA’s factorization and RoPE’s rotation. Thus, TPA maintains compatibility with RoPE. This completes the proof of Theorem 1. ∎
Appendix B More on Experiments
We list the main architecture hyper-parameters and training devices in Table 6. We fix for all the models. Moreover, we fix the number of KV heads with 2 for GQA models; for MLA models; and , for TPA and TPA-KV only models. Other hyper-parameters are listed in Table 7.
B.2 Additional Experimental Results
We display the evaluation results for small-size (124M) models in Tables 8-9.
B.3 Ablation Studies on Learning Rates
We implement a set of parallel experiments for medium models with learning rate , and the curves for training loss, validation loss and validation perplexity are displayed in Figure 5. We also show the performance of these models on the benchmarks described in Section 4 in Tables 10-11. The results show that TPA and TPA-KVonly models can also outperform other types of attention with different learning rates.