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 Q\mathbf{Q}, K\mathbf{K}, and V\mathbf{V} activations using contextual tensor-decompositions to achieve 10×10\times 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 di⋅n+jd_{i\cdot n+j} is the (i⋅n+j)(i\cdot n+j)-th element of d\mathbf{d}.

Scaled dot-product attention (Vaswani et al., 2017) determines how to focus on different parts of an input sequence by comparing queries (Q\mathbf{Q}) and keys (K\mathbf{K}). It produces a weighted combination of the values (V\mathbf{V}). Formally, the attention output is:

where each of Q,K,V\mathbf{Q},\mathbf{K},\mathbf{V} is an (n×dk)(n\times d_{k}) matrix for nn tokens and key dimension dkd_{k}. The division by dk\sqrt{d_{k}} 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 hh total heads into GG 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 g(i)g(i) maps a head i∈[h]i\in[h] to its group index g∈[G]g\in[G], then:

By adjusting GG between 11 and hh, 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 (t−s)(t-s) 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 tt-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 (t,s)(t,s) position pairs. Therefore, MLA adds the additional ktR\mathbf{k}_{t}^{R} 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 Qt,Kt,Vt\mathbf{Q}_{t},\mathbf{K}_{t},\mathbf{V}_{t} into a sum of (contextual) tensor products whose ranks are RqR_{q}, RkR_{k}, and RvR_{v}, respectively and may differ. Specifically, for each token tt, with a small abuse of notation, we define:

Latent Factor Maps. Each factor in the tensor product depends on the token’s hidden state xt\mathbf{x}_{t}. 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 Q,K,V\mathbf{Q},\mathbf{K},\mathbf{V} are factorized, multi-head attention proceeds as in standard Transformers. For each head i∈{1,…,h}i\in\{1,\dots,h\}:

Parameter Initialization. We initialize the weight matrices WraQ\bm{W}_{r}^{a^{Q}}, WraK\bm{W}_{r}^{a^{K}}, WraV\bm{W}_{r}^{a^{V}}, WrbQ\bm{W}_{r}^{b^{Q}}, WrbK\bm{W}_{r}^{b^{K}}, WrbV\bm{W}_{r}^{b^{V}} using Xavier initialization (Glorot & Bengio, 2010). Specifically, each entry of the weight matrix is drawn from a uniform distribution with bounds [−6/(nin+nout),6/(nin+nout)][-\sqrt{6/(n_{\text{in}}+n_{\text{out}})},\sqrt{6/(n_{\text{in}}+n_{\text{out}})}], where ninn_{\text{in}} and noutn_{\text{out}} 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 Kt\mathbf{K}_{t} 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 Qt\mathbf{Q}_{t} be factorized by TPA as

In addition, assume Qt\mathbf{Q}_{t} and Ks\mathbf{K}_{s} are factorized by TPA and then rotated by RoPE⁡t,RoPE⁡s\operatorname{RoPE}_{t},\operatorname{RoPE}_{s}. Let Q~t=RoPE⁡t(Qt)\widetilde{\mathbf{Q}}_{t}=\operatorname{RoPE}_{t}(\mathbf{Q}_{t}) and K~s=RoPE⁡s(Ks)\widetilde{\mathbf{K}}_{s}=\operatorname{RoPE}_{s}(\mathbf{K}_{s}). Then we have

Focusing on individual heads ii, 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, RoPE⁡t\operatorname{RoPE}_{t} acts as a block-diagonal orthogonal transform (i.e., a matrix Tt\mathbf{T}_{t}) on BQ(xt)\mathbf{B}_{Q}(\mathbf{x}_{t}). Consequently, AQ(xt)\mathbf{A}_{Q}(\mathbf{x}_{t}) remains unchanged, while each column of BQ(xt)\mathbf{B}_{Q}(\mathbf{x}_{t}) is rotated appropriately, preserving the TPA structure.

3 KV Caching and Memory Reduction

TPA Factorized KV Caching. Instead of storing the full Kt\mathbf{K}_{t} and Vt\mathbf{V}_{t}, TPA stores only their factorized ranks. Specifically, we keep

Compared to the standard caching cost of 2 h dh2\,h\,d_{h}, the ratio is:

For large hh and dhd_{h} (typically dh=64d_{h}=64 or 128128), setting RK,RV≪dhR_{K},R_{V}\ll d_{h} (e.g., 11 or 22) often yields 10×10\times or more reduction.

4 Unifying MHA, MQA, and GQA as Non-contextual TPA

To match MHA with TPA, let RQ=RK=RV=hR_{Q}=R_{K}=R_{V}=h. Focusing on Qt\mathbf{Q}_{t}:

so that ei⊗⋅\mathbf{e}_{i}\otimes\cdot corresponds to the ii-th head of Qt\mathbf{Q}_{t}.

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 ii-th row, reconstituting the usual MHA form of Qt\mathbf{Q}_{t}. Analogous constructions hold for Kt\mathbf{K}_{t} and Vt\mathbf{V}_{t} using WiK,WiV\bm{W}^{K}_{i},\bm{W}^{V}_{i}. 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 RK=RV=1R_{K}=R_{V}=1 along the head dimension. Concretely,

forces every head to use the same Kt,Vt\mathbf{K}_{t},\mathbf{V}_{t}. 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 B\mathbf{B}. Another variant is to share the token-dimension factors of keys and values:

lowering parameter counts and the KV cache footprint. While it constrains K\mathbf{K} and V\mathbf{V} 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 arQ,arK,arV\mathbf{a}^{Q}_{r},\mathbf{a}^{K}_{r},\mathbf{a}^{V}_{r}, one may introduce element-wise nonlinearities such as σ(⋅)\sigma(\cdot) or softmax⁡(⋅)\operatorname{softmax}(\cdot). 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 (RQ,RK,RV)(R_{Q},R_{K},R_{V}), 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 rr 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 Q\mathbf{Q} and K\mathbf{K}. Within TPA, we pre-rotate the factor btQ(xt)\mathbf{b}^{Q}_{t}(\mathbf{x}_{t}) and bsK(xs)\mathbf{b}^{K}_{s}(\mathbf{x}_{s}) directly, so that each Ks\mathbf{K}_{s} is already rotated prior to caching, see (3.6) and Theorem 1.

Attention Step and Output Projection. Once we have Q,K,V\mathbf{Q},\mathbf{K},\mathbf{V} factorized per token with RoPE applied on Q\mathbf{Q} and K\mathbf{K}, the attention step proceeds for each head i∈{1,…,h}i\in\{1,\dots,h\} using (3.4). Finally, concatenating these hh 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 σ\sigma is the SiLU (a.k.a., swish) nonlinearity, ⊙\odot is element-wise product, and W1,W2,W3\bm{W}_{1},\bm{W}_{2},\bm{W}_{3} 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 LL such blocks yields a T6 model architecture with LL 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 hh 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 4dmodel24d_{\text{model}}^{2} 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 (β1,β2)=(0.9,0.95)(\beta_{1},\beta_{2})=(0.9,0.95), a weight decay of 0.10.1, and gradient clipping at 1.01.0. We follow the same setting as nanoGPT that the learning rate is managed by a cosine annealing scheduler (Loshchilov & Hutter, 2016) with 2,0002{,}000 warmup steps and a (total) global batch size of 480480. For the small, medium, and large models, we set maximum learning rates of 6×10−46\times 10^{-4}, 3×10−43\times 10^{-4}, and 2×10−42\times 10^{-4} (respectively), and minimum learning rates of 3×10−53\times 10^{-5}, 3×10−53\times 10^{-5}, and 1×10−51\times 10^{-5} (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 4949B 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 5×5\times–10×10\times, 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-RR 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 Tt\mathbf{T}_{t} is the block-diagonal matrix encoding RoPE. This allows us to define

Similarly, for the key tensor Ks\mathbf{K}_{s}, we have

Now, consider the product of the rotated queries and keys:

Since Tt\mathbf{T}_{t} and Ts\mathbf{T}_{s} encode positional rotations, the product TtTs⊤\mathbf{T}_{t}\mathbf{T}_{s}^{\top} corresponds to a relative rotation Tt−s\mathbf{T}_{t-s}. Therefore, we can express the above as

Focusing on individual heads ii, 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 dh=64d_{h}=64 for all the models. Moreover, we fix the number of KV heads with 2 for GQA models; dhR=32d_{h}^{R}=32 for MLA models; and Rk=Rv=2R_{k}=R_{v}=2, Rq=6R_{q}=6 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 6×10−46\times 10^{-4}, 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.