xuan-luo/temp / docs /experiments.tex
xuan-luo's picture
download
raw
14.4 kB
% 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.