| % Main results: ppl-evals/results_mlra_10m/ and ppl-evals/results_pg19/. | |
| % Main-results references: docs/main_results.bib; preamble: booktabs, natbib. | |
| \section{Experiments} | |
| \label{sec:experiments} | |
| \subsection{Implementation Details} | |
| \label{sec:implementation-details} | |
| We train KVPath-Concat and KVPath-Residual on FineWeb-Edu using the GPT-2 | |
| BPE tokenizer. Training proceeds in three stages: we first train from scratch | |
| at a context length of 2K on approximately 10 billion tokens, then extend | |
| the context to 8K and subsequently to 32K, using approximately 1 billion | |
| additional tokens at each stage. Each extension is initialized from the | |
| final checkpoint of the preceding stage. We keep the RoPE base fixed at | |
| $10^6$ throughout training. | |
| We use AdamW with $\beta_1=0.9$, $\beta_2=0.95$, $\epsilon=10^{-8}$, | |
| weight decay 0.1, and gradient clipping at 1.0, with a global batch size | |
| of 256. The peak learning rate is $10^{-3}$ for the initial 2K stage and | |
| $10^{-4}$ for the 8K and 32K extensions. At each stage, the learning rate | |
| is warmed up for 100 steps and then follows a cosine decay schedule. | |
| \subsection{Main Results} | |
| \label{sec:main-results} | |
| We compare GQA, CLA, KVPath-Concat, and KVPath-Residual under matched | |
| parameter and KV-cache budgets. | |
| \paragraph{GQA.} | |
| Our GQA baseline~\citep{ainslie2023gqa} has 12 layers with hidden size 512, | |
| eight query heads, and four KV heads. The head dimension is 64, and the | |
| SwiGLU intermediate dimension is 2048. Each layer independently generates | |
| and caches its keys and values. | |
| \paragraph{CLA.} | |
| Our CLA baseline~\citep{brandon2024cla} also uses 12 layers, hidden size 512, | |
| eight query heads, and four KV heads. It shares the full K/V states within | |
| each group of four consecutive layers, while retaining independent query | |
| and output projections at each layer. The head dimension is 256, and the | |
| SwiGLU intermediate dimension is 1024. | |
| \paragraph{KVPath-Concat.} | |
| KVPath-Concat uses 12 layers with hidden size 512, eight query heads, and | |
| four KV heads, organized into four-layer paths. Each key and value head | |
| concatenates a shared 128-dimensional alpha component with a | |
| 32-dimensional delta component computed at each layer, giving a head | |
| dimension of 160. The SwiGLU intermediate dimension is 1536. | |
| \paragraph{KVPath-Residual.} | |
| KVPath-Residual uses the same 12-layer backbone, hidden size 512, eight | |
| query heads, four KV heads, and four-layer paths. The first layer of | |
| each path initializes full 160-dimensional K/V states as the alpha | |
| component. Each subsequent layer produces a 32-dimensional delta | |
| component per key or value head, expands it to 160 dimensions, and adds | |
| it to the preceding state. The SwiGLU intermediate dimension is 1516. | |
| All four configurations have exactly 98,710,016 parameters and the same | |
| compressed KV-cache budget of 6,144 elements per token across the full | |
| model (12~KiB per token in BF16). The intermediate dimensions are chosen | |
| to match parameter counts while accommodating the different attention | |
| projections. | |
| % Cache audit, including both K and V: | |
| % GQA: 2*12*4*64 = 6144; CLA: 2*3*4*256 = 6144; | |
| % Concat: 2*4*(3*128 + 12*32) = 6144; | |
| % Residual: 2*4*(3*160 + 9*32) = 6144. | |
| % These are compressed representation budgets, not peak runtime allocations. | |
| \paragraph{Evaluation.} | |
| We evaluate on C4~\citep{raffel2020t5}, | |
| FineWeb-Edu~\citep{penedo2024fineweb}, | |
| Wikipedia~\citep{wikimedia_dumps}, and PG-19~\citep{rae2020compressive}. | |
| For each model, we evaluate the final 32K checkpoint at context lengths | |
| of 2K, 8K, and 32K. All models use identical evaluation tokens and | |
| sequence boundaries at each context length. We use approximately 10M | |
| tokens per benchmark and hold the evaluated token region fixed across | |
| context lengths, partitioning it into non-overlapping sequences. | |
| Texts are concatenated in a fixed order, with end-of-document tokens | |
| between PG-19 books. We report GPT-2 token-level perplexity, including | |
| for PG-19, and compute the average as the unweighted arithmetic mean | |
| over the four benchmarks. | |
| % Requires \usepackage{booktabs} in the paper preamble. | |
| \begin{table}[t] | |
| \centering | |
| \small | |
| \setlength{\tabcolsep}{4pt} | |
| \begin{tabular}{@{}llrrrrr@{}} | |
| \toprule | |
| Context & Model & C4 & FineWeb-Edu & Wikipedia & PG-19 & Avg. \\ | |
| \midrule | |
| 2K & GQA & 35.67 & 21.94 & 32.56 & 33.58 & 30.94 \\ | |
| & CLA & 36.22 & 22.41 & 33.13 & 35.11 & 31.72 \\ | |
| & KVPath-Concat & \textbf{33.81} & \textbf{20.89} & \textbf{32.35} & \textbf{32.87} & \textbf{29.98} \\ | |
| & KVPath-Residual & 34.21 & 21.05 & 32.61 & 33.71 & 30.40 \\ | |
| \midrule | |
| 8K & GQA & 34.42 & 21.14 & 30.82 & 31.10 & 29.37 \\ | |
| & CLA & 34.91 & 21.56 & 31.20 & 32.67 & 30.08 \\ | |
| & KVPath-Concat & \textbf{32.50} & \textbf{20.04} & \textbf{30.34} & \textbf{30.30} & \textbf{28.29} \\ | |
| & KVPath-Residual & 32.93 & 20.24 & 30.74 & 31.11 & 28.76 \\ | |
| \midrule | |
| 32K & GQA & 34.15 & 20.95 & 30.38 & 30.68 & 29.04 \\ | |
| & CLA & 34.57 & 21.33 & 30.67 & 32.22 & 29.70 \\ | |
| & KVPath-Concat & \textbf{32.23} & \textbf{19.80} & \textbf{29.59} & \textbf{29.36} & \textbf{27.75} \\ | |
| & KVPath-Residual & 32.65 & 20.02 & 30.19 & 30.41 & 28.32 \\ | |
| \bottomrule | |
| \end{tabular} | |
| \caption{Perplexity ($\downarrow$) of the final 32K checkpoints evaluated | |
| at three context lengths. All models have 98.71M parameters and a | |
| compressed KV-cache budget of 6,144 elements per token. | |
| Avg.\ is the arithmetic mean over the four benchmarks. | |
| Bold indicates the best result within each context length.} | |
| \label{tab:main-results} | |
| \end{table} | |
| \paragraph{Results.} | |
| Table~\ref{tab:main-results} shows that KVPath-Concat achieves the lowest | |
| perplexity on all four benchmarks at every context length. Relative to | |
| GQA, it reduces average perplexity by 3.09\%, 3.65\%, and 4.45\% at 2K, | |
| 8K, and 32K, respectively. KVPath-Residual also improves average | |
| perplexity over GQA, with reductions of 1.75\%, 2.08\%, and 2.48\%. | |
| Its gains are less uniform: it is slightly worse than GQA on Wikipedia | |
| at 2K and on PG-19 at 2K and 8K, but improves on all four benchmarks | |
| at 32K. CLA has higher perplexity than GQA throughout this evaluation. | |
| The increasing average gains of both KVPath variants at longer | |
| contexts support the usefulness of compact layer-dependent KV | |
| representations under a fixed cache budget. | |
| \subsection{Compatibility with MLA} | |
| \label{sec:mla-compatibility} | |
| \paragraph{MLA.} | |
| Multi-head latent attention (MLA)~\citep{deepseek2024v2} reduces KV-cache | |
| storage by compressing each layer's keys and values into a low-dimensional | |
| latent. Head-specific content keys and values are reconstructed from this | |
| latent when attention is computed, while a small positional key branch is | |
| cached separately and receives RoPE. Our MLA baseline has 12 layers, hidden | |
| size 512, and eight attention heads. At each layer, it caches a | |
| 128-dimensional KV latent and a 32-dimensional positional key, and | |
| reconstructs 64-dimensional content keys and values for each head. | |
| \paragraph{KVPath-MLA.} | |
| We instantiate KVPath on MLA using the KVPath-Concat design, which performs | |
| best in Section~\ref{sec:main-results}. For each four-layer path, its first | |
| layer produces a shared 384-dimensional alpha latent $c_\alpha$. Layer | |
| $\ell$ then uses its own normalization and up-projection to reconstruct | |
| 64-dimensional alpha components, while independently generating | |
| 32-dimensional key and value delta components. Suppressing the head | |
| index for clarity, this construction is | |
| \begin{equation} | |
| \begin{aligned} | |
| c_\alpha&=W_\alpha h^{(0)},\\ | |
| (k_\alpha^{(\ell)},v_\alpha^{(\ell)}) | |
| &=W_U^{(\ell)}\operatorname{RMSNorm}^{(\ell)}(c_\alpha),\\ | |
| (k_\delta^{(\ell)},v_\delta^{(\ell)}) | |
| &=W_\delta^{(\ell)}h^{(\ell)}. | |
| \end{aligned} | |
| \label{eq:kvpath-mla-components} | |
| \end{equation} | |
| The alpha latent is therefore shared across the path, although each layer | |
| can reconstruct different alpha components. The delta components are | |
| layer-specific and shared across heads. Following KVPath-Concat, each layer | |
| concatenates the two components and computes attention as | |
| \begin{equation} | |
| \begin{aligned} | |
| q^{(\ell)} | |
| &=[q_\alpha^{(\ell)}\,\Vert\, | |
| \operatorname{RoPE}(q_\delta^{(\ell)})],\\ | |
| k^{(\ell)} | |
| &=[k_\alpha^{(\ell)}\,\Vert\, | |
| \operatorname{RoPE}(k_\delta^{(\ell)})],\\ | |
| v^{(\ell)} | |
| &=[v_\alpha^{(\ell)}\,\Vert\,v_\delta^{(\ell)}],\\ | |
| O^{(\ell)} | |
| &=\operatorname{softmax}\!\left( | |
| \frac{Q^{(\ell)}(K^{(\ell)})^\top}{\sqrt{96}}+M | |
| \right)V^{(\ell)}. | |
| \end{aligned} | |
| \label{eq:kvpath-mla-attention} | |
| \end{equation} | |
| Here $Q^{(\ell)},K^{(\ell)},V^{(\ell)}$ stack the corresponding | |
| token-wise representations. Thus, the 64-dimensional alpha components and | |
| 32-dimensional delta components form a single 96-dimensional attention | |
| head and a single attention distribution. | |
| MLA and KVPath-MLA both have exactly 99,694,592 parameters and a compressed | |
| cache budget of 1,920 elements per token. MLA stores a 128-dimensional | |
| latent and a 32-dimensional positional key at every layer, whereas | |
| KVPath-MLA stores one 384-dimensional alpha latent per four-layer path and | |
| 32-dimensional key and value delta components at every layer. Their | |
| SwiGLU widths are adjusted to match the parameter count. Following the | |
| same training setup described above, both models are trained on approximately | |
| 10B tokens at 2K and then continued on approximately 1B tokens at 8K and | |
| 1B tokens at 32K, with all settings matched between them. | |
| % Theoretical compressed cache budgets, not measured runtime allocations: | |
| % MLA: 12*(128+32)=1920; KVPath-MLA: 3*384+12*(32+32)=1920. | |
| % Exact FFNs: MLA 2048 in all layers; KVPath-MLA 1792 in layers 1--10 and | |
| % 1791 in layers 11--12. | |
| \begin{table}[t] | |
| \centering | |
| \small | |
| \setlength{\tabcolsep}{4pt} | |
| \begin{tabular}{@{}llrrrrr@{}} | |
| \toprule | |
| Context & Model & C4 & FineWeb-Edu & Wikipedia & PG-19 & Avg. \\ | |
| \midrule | |
| 2K & MLA & 36.27 & 22.22 & 32.22 & 34.88 & 31.40 \\ | |
| & KVPath-MLA & \textbf{35.14} & \textbf{21.63} & \textbf{30.74} & \textbf{32.80} & \textbf{30.08} \\ | |
| \midrule | |
| 8K & MLA & 34.91 & 21.37 & 30.48 & 32.18 & 29.74 \\ | |
| & KVPath-MLA & \textbf{33.82} & \textbf{20.78} & \textbf{28.86} & \textbf{30.05} & \textbf{28.38} \\ | |
| \midrule | |
| 32K & MLA & 34.63 & 21.15 & 30.08 & 31.28 & 29.28 \\ | |
| & KVPath-MLA & \textbf{33.52} & \textbf{20.55} & \textbf{28.36} & \textbf{29.34} & \textbf{27.95} \\ | |
| \bottomrule | |
| \end{tabular} | |
| \caption{Perplexity ($\downarrow$) of the final 32K MLA and KVPath-MLA | |
| checkpoints. Both models have 99.69M parameters and a compressed cache | |
| budget of 1,920 elements per token. Avg.\ is the arithmetic mean over | |
| the four benchmarks.} | |
| \label{tab:mla-results} | |
| \end{table} | |
| \paragraph{Results.} | |
| Following the protocol in Section~\ref{sec:main-results}, we evaluate each | |
| final 32K checkpoint at 2K, 8K, and 32K. Table~\ref{tab:mla-results} shows | |
| that KVPath-MLA improves over MLA on every benchmark and context length. | |
| It reduces average perplexity by 4.20\%, 4.57\%, and 4.57\% at 2K, 8K, | |
| and 32K, respectively. These consistent gains show that the | |
| KVPath-Concat principle remains effective when the shared cross-layer | |
| representation is an MLA latent rather than an explicit KV state. | |
| \subsection{Scaling} | |
| \label{sec:scaling} | |
| To examine whether the benefits persist at a larger scale, we plan to compare | |
| GQA and KVPath-I at approximately 1.0B parameters, training each model on | |
| 100B tokens. The comparison will use the same training data and optimization | |
| schedule, with matched parameter and KV-cache budgets. | |
| Table~\ref{tab:scaling-results} reserves space for the final validation results. | |
| % TODO: Add the model configurations, context length, and results once available. | |
| \begin{table}[t] | |
| \centering | |
| \small | |
| \begin{tabular}{lrrrr} | |
| \hline | |
| Model & Params & Tokens & KV/token & PPL $\downarrow$ \\ | |
| \hline | |
| GQA & $\sim$1.0B & 100B & -- & -- \\ | |
| KVPath-I & $\sim$1.0B & 100B & -- & -- \\ | |
| \hline | |
| \end{tabular} | |
| \caption{Planned comparison at 1.0B parameters and 100B training tokens. | |
| Cache budgets and validation results remain to be filled in.} | |
| \label{tab:scaling-results} | |
| \end{table} | |
| \subsection{Ablation Study} | |
| \label{sec:ablation} | |
| We focus on two design choices in KVPath-I: the number of layers connected | |
| by a path and the dimension of the layer-specific KV. The reference is the | |
| 12-layer configuration in Section~\ref{sec:main-results}, with $g=4$, | |
| $d_p=128$, and $d_s=32$. Each ablation varies one of these choices while | |
| keeping the remaining architecture fixed, and follows the same training | |
| recipe. Table~\ref{tab:ablation-results} includes parameter counts and cache | |
| sizes to make the resulting budget changes explicit. | |
| \paragraph{Group size.} | |
| We consider $g\in\{2,4,8\}$ and an ungrouped path spanning all 12 layers | |
| ($g=12$). This comparison examines how far a shared path can extend across | |
| depth before refreshing its state becomes useful. For $g=8$, the final | |
| group contains the remaining four layers. Increasing $g$ reduces the number | |
| of independently initialized paths and their storage cost, while retaining | |
| the same layer-specific KV at every layer. | |
| \paragraph{Layer-specific dimension.} | |
| We vary $d_s\in\{16,32,64\}$ with $d_p=128$ and $g=4$ fixed. This comparison | |
| examines the effect of allocating more dimensions to information computed | |
| independently at each layer. The path representation remains unchanged, | |
| while the total KV dimension per head is $128+d_s$. | |
| \begin{table*}[t] | |
| \centering | |
| \small | |
| \begin{tabular}{lrrrrr} | |
| \hline | |
| Configuration & Params & KV/token & 2K PPL $\downarrow$ & 8K PPL $\downarrow$ & 32K PPL $\downarrow$ \\ | |
| \hline | |
| Reference ($g=4$, $d_s=32$) & 98,710,016 & 6,144 & 21.1444 & 20.3029 & 20.4338 \\ | |
| \hline | |
| $g=2$ & 100,282,880 & 9,216 & -- & -- & -- \\ | |
| $g=8$ & 98,185,728 & 5,120 & -- & -- & -- \\ | |
| No grouping ($g=12$) & 97,661,440 & 4,096 & -- & -- & -- \\ | |
| \hline | |
| $d_s=16$ & 96,350,720 & 4,608 & -- & -- & -- \\ | |
| $d_s=64$ & 103,428,608 & 9,216 & -- & -- & -- \\ | |
| \hline | |
| \end{tabular} | |
| \caption{KVPath-I ablations. Each row changes only the indicated setting | |
| relative to the reference. Parameter counts and compressed cache elements | |
| per token are calculated from the configurations. Reference perplexities | |
| are reproduced from Table~\ref{tab:main-results}; uncompleted results are | |
| left blank (--).} | |
| \label{tab:ablation-results} | |
| \end{table*} | |
| % TODO: Discuss the group-size and layer-specific-dimension results after | |
| % the ablations finish; no conclusions are assumed here. | |
Xet Storage Details
- Size:
- 14.4 kB
- Xet hash:
- b0e5fd61cddc56632e501a2786f89a28179d164aa985cfa527c5898ed588be94
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.