Add DPO data generation & training; fix checkpoint labels; switch to WSD scheduler; update optimizer ablations in abstract/conclusion

This commit is contained in:
2026-07-19 15:25:42 +08:00
parent 6c2e04a86f
commit e149997200
2 changed files with 88 additions and 39 deletions
Binary file not shown.

Before

Width:  |  Height:  |  Size: 100 KiB

+88 -39
View File
@@ -18,7 +18,7 @@
\DeclareMathOperator{\Var}{Var} \DeclareMathOperator{\Var}{Var}
\title{End-to-End Training of a 1.2B Transformer with AstrAI \\ \title{End-to-End Training of a 1.2B Transformer with AstrAI \\
\large Data Pipeline, Distributed Training, and BF16 Numerical Stability via Residual Scaling} \large Data Pipeline, Distributed Training, and Ablations on Optimizer, Initialization, and BF16 Numerical Stability}
\author{AstrAI Contributors} \author{AstrAI Contributors}
\date{} \date{}
@@ -34,21 +34,24 @@ JSON-driven BBPE preprocessing with multi-strategy packing, HDF5/mmap
storage backends, and a companion SFT pipeline ({\sc Alembic}) with MinHash storage backends, and a companion SFT pipeline ({\sc Alembic}) with MinHash
deduplication and LLM-as-Judge scoring. The 24-layer decoder uses GQA, deduplication and LLM-as-Judge scoring. The 24-layer decoder uses GQA,
SwiGLU, RoPE, and RMSNorm, trained with a hybrid Muon/AdamW optimizer and SwiGLU, RoPE, and RMSNorm, trained with a hybrid Muon/AdamW optimizer and
cosine scheduling under DDP/FSDP. A BF16 stability analysis shows that WSD (Warmup--Stable--Decay) scheduling under DDP/FSDP. A BF16 stability analysis shows that
GPT-2 residual scaling substantially reduces per-block residual variance GPT-2 residual scaling substantially reduces per-block residual variance
accumulation, keeping post-training variance well below the overflow accumulation, keeping post-training variance well below the overflow
threshold of standard initialization; empirically this yields a sustained threshold of standard initialization; empirically this yields a sustained
loss advantage over Kaiming initialization throughout training. loss advantage over Kaiming initialization throughout training.
Post-training weight distribution analysis across three Systematic ablations show that the hybrid Muon/AdamW optimizer
checkpoints---varying optimizer, initialization, and training outperforms pure AdamW on 2D weight matrices, and that GPT-2 residual
budget---confirms that residual-scaled projections maintain narrow scaling yields a sustained loss advantage over both Kaiming and Normal
distributions throughout training, preserving the numerical stability initialization throughout training. Post-training weight distribution
established at initialization. An SVD-based effective rank analysis analysis across three checkpoints---varying optimizer, initialization,
further reveals that the model operates near its representational and training budget---confirms that residual-scaled projections
capacity, with attention Q/O projections consistently showing lower maintain narrow distributions throughout training, preserving the
utilization than K/V projections, a pattern stable across all numerical stability established at initialization. An SVD-based
configurations and consistent with the low-rank structure induced by effective rank analysis further reveals that the model operates near
grouped query attention. its representational capacity, with attention Q/O projections
consistently showing lower utilization than K/V projections, a pattern
stable across all configurations and consistent with the low-rank
structure induced by grouped query attention.
\end{abstract} \end{abstract}
% ====================================================================== % ======================================================================
@@ -60,8 +63,10 @@ model architecture. Data must be preprocessed and stored efficiently, the
training loop must handle distributed parallelism, gradient accumulation, training loop must handle distributed parallelism, gradient accumulation,
checkpointing, and logging---and numerical pitfalls must be diagnosed and checkpointing, and logging---and numerical pitfalls must be diagnosed and
fixed. This paper describes the complete workflow using {\sc AstrAI}~\cite{astrai}, an fixed. This paper describes the complete workflow using {\sc AstrAI}~\cite{astrai}, an
open-source framework for Transformer training and inference, and highlights a open-source framework for Transformer training and inference. Beyond the
BF16 precision issue encountered along the way. pipeline, we conduct systematic ablations on optimizers (hybrid Muon/AdamW
versus pure AdamW) and initializations (GPT-2 residual scaling, Kaiming,
and Normal), and analyse a BF16 precision issue encountered along the way.
% ====================================================================== % ======================================================================
\section{Data Pipeline} \section{Data Pipeline}
@@ -140,10 +145,33 @@ pipeline proceeds as follows:
An optional LLM-as-Judge scoring module provides multi-dimensional An optional LLM-as-Judge scoring module provides multi-dimensional
quality scores that can be used to filter low-quality samples. quality scores that can be used to filter low-quality samples.
An IFD (Instruction Fulfillment Difficulty) analysis is provided in \subsection{DPO Data Generation}
Appendix~\ref{app:ifd}.
% ====================================================================== To construct pairwise preference data for Direct Preference
Optimization, we start from a bilingual (Chinese--English) instruction
set curated by the same MinHash pipeline described above. For each
prompt $x$ in the instruction set, we generate two responses:
\begin{itemize}[nosep]
\item \textbf{Chosen} $y_w$: generated by the reference model
$\pi_{\text{ref}}$, i.e.~the SFT checkpoint after supervised
fine-tuning on the curated instruction data.
\item \textbf{Rejected} $y_l$: generated by the base model
$\pi_{\text{base}}$, i.e.~the model at the end of pretraining
before any instruction tuning.
\end{itemize}
Both generations use greedy decoding to eliminate sampling variance and
ensure that the preference signal reflects model capability rather than
decoding randomness. The resulting preference pairs
$(x, y_w, y_l)$ are stored in the same JSONL format used for SFT and
fed directly into the DPO training loop
(Section~\ref{sec:dpo}). Because the base model and the reference
model share the same architecture but differ in instruction-following
ability, the contrast between $y_w$ and $y_l$ is sharp and consistent,
which stabilises the DPO gradient updates. All instructions are
balanced across Chinese and English domains to preserve bilingual
alignment capability during preference optimisation.
\section{Model Architecture} \section{Model Architecture}
% ====================================================================== % ======================================================================
@@ -231,8 +259,8 @@ The model is trained on next-token cross-entropy loss:
\mathcal{L} = -\sum_{t=1}^{T} \log P(x_t \mid x_{<t}; \theta). \mathcal{L} = -\sum_{t=1}^{T} \log P(x_t \mid x_{<t}; \theta).
\end{equation} \end{equation}
Training uses a hybrid optimizer: Muon for 2D weight matrices and AdamW~\cite{loshchilov2019adamw} for 1D parameters (embeddings, biases, LayerNorm), with cosine learning rate Training uses a hybrid optimizer: Muon for 2D weight matrices and AdamW~\cite{loshchilov2019adamw} for 1D parameters (embeddings, biases, LayerNorm), with WSD (Warmup--Stable--Decay)
scheduling (2\% warmup) and global L2 gradient clipping. The framework supports DDP and FSDP for multi-GPU distribution, learning rate scheduling (2\% warmup) and global L2 gradient clipping. The framework supports DDP and FSDP for multi-GPU distribution,
with gradient accumulation to manage memory. with gradient accumulation to manage memory.
Table~\ref{tab:train_params} lists the key hyperparameters. Table~\ref{tab:train_params} lists the key hyperparameters.
@@ -245,14 +273,16 @@ Table~\ref{tab:train_params} lists the key hyperparameters.
\textbf{Hyperparameter} & \textbf{Value} \\ \textbf{Hyperparameter} & \textbf{Value} \\
\midrule \midrule
Precision & BF16 (weights + optimizer states) \\ Precision & BF16 (weights + optimizer states) \\
Optimizer & AdamW, $\eta=1.5\times10^{-4}$ \\ Optimizer & Hybrid Muon/AdamW$^a$, $\eta=1.5\times10^{-4}$ \\
Betas & $(0.9, 0.95)$, weight decay $0.1$ \\ Betas & $(0.9, 0.95)$, weight decay $0.1$ \\
Gradient clip & Global L2, max norm $1.0$ \\ Gradient clip & Global L2, max norm $1.0$ \\
Scheduler & Cosine, warmup ratio $0.02$ \\ Scheduler & WSD (warmup 2\%, stable, decay) \\
Batch size & 4 per device $\times$ 4 GPUs $\times$ 32 accumulation \\ Batch size & 4 per device $\times$ 4 GPUs $\times$ 32 accumulation \\
Sequence length & 2,048 tokens \\ Sequence length & 2,048 tokens \\
Total steps & 19,000 \\ Total tokens & $\sim$25B ($\approx$23k steps) \\
\bottomrule \bottomrule
\multicolumn{2}{@{}l@{}}{\footnotesize $^a$Muon applied to 2D weight matrices (attention projections and FFN layers);}\\
\multicolumn{2}{@{}l@{}}{\footnotesize \phantom{$^a$}AdamW applied to 1D parameters (embeddings, biases, LayerNorm scales).}\\
\end{tabular} \end{tabular}
\end{table} \end{table}
@@ -280,8 +310,8 @@ Figure~\ref{fig:ckpt_comparison} compares four configurations. The left panel sh
rate schedule for the hybrid Muon/AdamW configuration across $\sim$25B rate schedule for the hybrid Muon/AdamW configuration across $\sim$25B
tokens. The Muon optimizer is applied to all 2D weight matrices tokens. The Muon optimizer is applied to all 2D weight matrices
(attention projections and FFN layers), while AdamW handles 1D (attention projections and FFN layers), while AdamW handles 1D
parameters (embeddings, biases, and normalization scales). The cosine parameters (embeddings, biases, and normalization scales). The WSD
schedule with 2\% warmup is visible in the lower panel.} schedule (warmup, stable phase, and final decay) is visible in the lower panel.}
\label{fig:muon_pt} \label{fig:muon_pt}
\end{figure} \end{figure}
@@ -289,9 +319,9 @@ Figure~\ref{fig:muon_pt} shows the extended training dynamics of the
Muon optimizer configuration over $\sim$25B tokens. The loss curve Muon optimizer configuration over $\sim$25B tokens. The loss curve
exhibits the expected power-law decay in the early phase (0--5B tokens), exhibits the expected power-law decay in the early phase (0--5B tokens),
followed by a gradual plateau as the model approaches convergence on followed by a gradual plateau as the model approaches convergence on
the pretraining distribution. The cosine learning rate schedule the pretraining distribution. The WSD schedule holds the learning rate
reaches its minimum at the end of training, ensuring stable weight constant during the long stable phase and then decays at the end of
updates in the final phase. training, ensuring stable weight updates in the final phase.
\begin{figure}[H] \begin{figure}[H]
\centering \centering
@@ -300,7 +330,7 @@ updates in the final phase.
mixed Chinese--English instruction dataset: training loss, learning mixed Chinese--English instruction dataset: training loss, learning
rate, and gradient norm. The loss drops rapidly in the first 200 rate, and gradient norm. The loss drops rapidly in the first 200
steps and then enters a slower decay phase. The learning rate steps and then enters a slower decay phase. The learning rate
follows a cosine schedule with 2\% warmup. Gradient norms stabilize follows a WSD schedule (2\% warmup, stable, decay). Gradient norms stabilize
after approximately 300 steps, indicating that the fine-tuning after approximately 300 steps, indicating that the fine-tuning
process has reached a stable optimization regime.} process has reached a stable optimization regime.}
\label{fig:sft_metrics} \label{fig:sft_metrics}
@@ -331,7 +361,14 @@ preferred response $y_w$ and dispreferred response $y_l$, the loss is:
\end{equation} \end{equation}
where $\beta$ controls the KL-divergence penalty against the reference. where $\beta$ controls the KL-divergence penalty against the reference.
We use $\beta=0.1$, a batch of 64 preference pairs per step, and the We use $\beta=0.1$, a batch of 64 preference pairs per step, and the
same hybrid Muon/AdamW optimizer with cosine scheduling. same hybrid Muon/AdamW optimizer. Unlike the pretraining and SFT
phases, which both employ WSD (warmup--stable--decay) scheduling to
maintain a long high-learning-rate plateau, DPO alignment uses a
cosine schedule with short linear warmup. The shorter ~1{,}500-step
alignment run does not benefit from an extended stable phase; instead,
cosine decay lowers the learning rate steadily, which discourages
over-optimisation away from the reference distribution and matches the
observed convergence pattern in Figure~\ref{fig:dpo_metrics}.
\begin{figure}[H] \begin{figure}[H]
\centering \centering
@@ -459,7 +496,7 @@ identified in the theoretical analysis (Section~\ref{sec:num-stability}).
To verify that the residual-scaling constraint persists throughout To verify that the residual-scaling constraint persists throughout
training---not just at initialization---we compare weight value training---not just at initialization---we compare weight value
distributions across three checkpoints: distributions across three checkpoints:
\texttt{kami-15bt} (Muon, 15B tokens), \texttt{norm-15bt} (Normal \texttt{kami-15bt} (AdamW, GPT-2 residual scaling, 15B tokens), \texttt{norm-15bt} (AdamW, Normal
init, 15B tokens), and \texttt{muon-25bt} (Muon, 25B tokens). init, 15B tokens), and \texttt{muon-25bt} (Muon, 25B tokens).
Figure~\ref{fig:ckpt_weight_density} shows the kernel density Figure~\ref{fig:ckpt_weight_density} shows the kernel density
@@ -476,13 +513,16 @@ are visible in the per-component weight std
($\sim$0.024 for attention projections), exceeding both 15B ($\sim$0.024 for attention projections), exceeding both 15B
checkpoints ($\sim$0.015--0.021), reflecting continued weight checkpoints ($\sim$0.015--0.021), reflecting continued weight
drift from $\sigma_0 = 0.02$ as training progresses. drift from $\sigma_0 = 0.02$ as training progresses.
\item \textbf{Muon constrains early-stage drift}: at equal training \item \textbf{Residual scaling constrains early-stage drift}: at
budget (15B tokens), the Muon-trained \texttt{kami-15bt} shows equal training budget (15B tokens) and with the same AdamW
optimizer, the GPT-2-scaled \texttt{kami-15bt} shows
\emph{smaller} std ($\sim$0.015) than the Normal-init \emph{smaller} std ($\sim$0.015) than the Normal-init
\texttt{norm-15bt} ($\sim$0.021). Muon's orthogonalization \texttt{norm-15bt} ($\sim$0.021), indicating that residual
step normalizes the update direction, constraining per-step scaling itself limits weight drift beyond its initialization
weight movement. At 25B tokens, cumulative steps overtake effect. The Muon-trained \texttt{muon-25bt} (25B tokens)
this per-step constraint, producing the largest overall std. exhibits the largest std because cumulative training steps
eventually overtake both initialization and per-step optimizer
constraints.
\end{itemize} \end{itemize}
Critically, the residual-scaled projections remain bounded at Critically, the residual-scaled projections remain bounded at
@@ -522,8 +562,17 @@ by a factor of 48, keeping post-24-layer variance at $1.34$ versus $17.5$
without scaling. Post-training weight distribution analysis confirms that without scaling. Post-training weight distribution analysis confirms that
this scaling constraint persists throughout training, with residual-scaled this scaling constraint persists throughout training, with residual-scaled
projections maintaining narrow distributions ($\sigma \approx 0.003$) projections maintaining narrow distributions ($\sigma \approx 0.003$)
regardless of optimizer or training duration. An SVD-based effective rank regardless of optimizer or training duration.
analysis (Appendix~\ref{app:eff_rank}) further reveals that the model
Optimizer ablations (Figure~\ref{fig:ckpt_comparison}) demonstrate that
the hybrid Muon/AdamW configuration consistently outperforms pure AdamW
with both Kaiming and Normal initializations, achieving lower training loss
and faster convergence. Within the Muon family, applying AdamW to the
embedding layer rather than Muon yields a further loss reduction and more
stable gradient norms, confirming that 1D parameters benefit from
second-moment adaptation while 2D weight matrices gain from Muon's
orthogonalisation step. An SVD-based effective rank analysis
(Appendix~\ref{app:eff_rank}) further reveals that the model
operates near its representational capacity. The complete framework and operates near its representational capacity. The complete framework and
model weights are available at \url{https://github.com/ViperEkura/AstrAI}. model weights are available at \url{https://github.com/ViperEkura/AstrAI}.
@@ -787,7 +836,7 @@ remain bounded.}
\begin{tabular}{@{}lccc@{}} \begin{tabular}{@{}lccc@{}}
\toprule \toprule
\textbf{Component} & \textbf{kami-15bt} & \textbf{norm-15bt} & \textbf{muon-25bt} \\ \textbf{Component} & \textbf{kami-15bt} & \textbf{norm-15bt} & \textbf{muon-25bt} \\
& (Muon, 15B) & (Normal, 15B) & (Muon, 25B) \\ & (AdamW, 15B) & (AdamW, 15B) & (Muon, 25B) \\
\midrule \midrule
attn.q\_proj & 0.0154 & 0.0207 & 0.0230 \\ attn.q\_proj & 0.0154 & 0.0207 & 0.0230 \\
attn.k\_proj & 0.0153 & 0.0206 & 0.0238 \\ attn.k\_proj & 0.0153 & 0.0206 & 0.0238 \\