Compare commits
13
Commits
411354eeb1
..
master
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e6be33aa53 | ||
|
|
0a1d0573ae | ||
|
|
0c2bc916f2 | ||
|
|
4e70e827ff | ||
|
|
a7bbc7b29f | ||
|
|
4b37f289c0 | ||
|
|
a83555f326 | ||
|
|
d93ff48320 | ||
|
|
da0f536526 | ||
|
|
c775a2b3e0 | ||
|
|
3f0ff911a8 | ||
|
|
e149997200 | ||
|
|
6c2e04a86f |
Binary file not shown.
|
After Width: | Height: | Size: 86 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 128 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 68 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 84 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 172 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 100 KiB |
@@ -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{}
|
||||||
@@ -29,25 +29,25 @@
|
|||||||
|
|
||||||
\begin{abstract}
|
\begin{abstract}
|
||||||
We present {\sc AstrAI}, an open-source framework for end-to-end training
|
We present {\sc AstrAI}, an open-source framework for end-to-end training
|
||||||
of a 1.2B-parameter Transformer on $\sim$20B tokens. The pipeline covers
|
of a 1.2B-parameter Transformer on $\sim$25B tokens. The pipeline covers
|
||||||
JSON-driven BBPE preprocessing with multi-strategy packing, HDF5/mmap
|
JSON-driven BBPE preprocessing with multi-strategy packing, tiered
|
||||||
storage backends, and a companion SFT pipeline ({\sc Alembic}) with MinHash
|
storage backends, and a companion SFT pipeline ({\sc Alembic}) with
|
||||||
deduplication and LLM-as-Judge scoring. The 24-layer decoder uses GQA,
|
MinHash deduplication. The 24-layer GQA-SwiGLU decoder is trained with a
|
||||||
SwiGLU, RoPE, and RMSNorm, trained with a hybrid Muon/AdamW optimizer and
|
hybrid Muon/AdamW optimizer and WSD scheduling under DDP/FSDP.
|
||||||
cosine scheduling under DDP/FSDP. A BF16 stability analysis shows that
|
Supervised fine-tuning on deduplicated bilingual instructions reduces loss
|
||||||
GPT-2 residual scaling substantially reduces per-block residual variance
|
from $\sim$2.1 to $\sim$1.5 over $\sim$3{,}800~steps; DPO alignment
|
||||||
accumulation, keeping post-training variance well below the overflow
|
($\beta=0.1$, cosine schedule) on model-generated preference pairs
|
||||||
threshold of standard initialization; empirically this yields a sustained
|
shows stable training-loss convergence without over-optimisation. A BF16 stability analysis
|
||||||
loss advantage over Kaiming initialization throughout training.
|
shows that GPT-2 residual scaling ($\sigma_0 = 0.02/\sqrt{2L}$) reduces
|
||||||
Post-training weight distribution analysis across three
|
per-block activation variance by a factor of 48, and post-training weight
|
||||||
checkpoints---varying optimizer, initialization, and training
|
analysis across three checkpoints confirms that residual-scaled
|
||||||
budget---confirms that residual-scaled projections maintain narrow
|
projections remain consistently narrower than non-scaled weights across
|
||||||
distributions throughout training, preserving the numerical stability
|
optimizers, initializations, and training budgets. Optimizer ablations
|
||||||
established at initialization. An SVD-based effective rank analysis
|
demonstrate that the hybrid Muon/AdamW outperforms pure AdamW, with
|
||||||
further reveals that the model operates near its representational
|
2D weight matrices benefiting from Muon's orthogonalisation and 1D
|
||||||
capacity, with attention Q/O projections consistently showing lower
|
parameters from AdamW's second-moment adaptation. An SVD effective-rank
|
||||||
utilization than K/V projections, a pattern stable across all
|
analysis reveals near-capacity weight utilization, with Q/O projections
|
||||||
configurations and consistent with the low-rank structure induced by
|
showing consistently lower effective rank than K/V projections under
|
||||||
grouped query attention.
|
grouped query attention.
|
||||||
\end{abstract}
|
\end{abstract}
|
||||||
|
|
||||||
@@ -60,8 +60,12 @@ 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, from JSONL
|
||||||
BF16 precision issue encountered along the way.
|
ingestion through pretraining, supervised fine-tuning (SFT) on deduplicated
|
||||||
|
bilingual instructions, and direct preference optimization (DPO) alignment.
|
||||||
|
We also 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}
|
||||||
@@ -86,17 +90,30 @@ via a JSON specification that defines:
|
|||||||
\texttt{.bin} shards, auto-split at 100M tokens per shard.
|
\texttt{.bin} shards, auto-split at 100M tokens per shard.
|
||||||
\end{itemize}
|
\end{itemize}
|
||||||
|
|
||||||
Samples shorter than 50~chars or longer than 2M~chars are filtered out.
|
In text-mode sections, individual fields shorter than 50~chars or
|
||||||
|
longer than 2M~chars are skipped during tokenization.
|
||||||
|
|
||||||
\subsection{Storage Backends}
|
\subsection{Storage Backends}
|
||||||
|
|
||||||
Two storage backends serve the DataLoader:
|
Three storage backends serve the DataLoader, trading off memory
|
||||||
|
footprint against access speed for datasets ranging from fine-tuning
|
||||||
|
scale to TB-level pretraining:
|
||||||
|
|
||||||
\begin{itemize}[nosep]
|
\begin{itemize}[nosep]
|
||||||
\item \textbf{H5Store}: HDF5-based, memory-loaded with shared-memory support
|
\item \textbf{H5Store}: HDF5-based, fully loaded into RAM at init
|
||||||
for multi-worker access.
|
with \texttt{share\_memory\_()} for cross-worker sharing.
|
||||||
\item \textbf{MmapStore}: Zero-copy memory-mapped \texttt{.bin} files shared
|
Fastest random access; requires the dataset to fit in memory.
|
||||||
via OS page cache.
|
\item \textbf{MmapStore}: Zero-copy memory-mapped \texttt{.bin} files
|
||||||
|
via \texttt{np.memmap}. Data stays on disk, managed by the OS
|
||||||
|
page cache; multiple workers share physical pages without
|
||||||
|
duplication. Suited for large pretraining corpora that exceed
|
||||||
|
RAM.
|
||||||
|
\item \textbf{JsonlStore}: Reads raw \texttt{.jsonl} directly. Lazy
|
||||||
|
mode (\texttt{processor=fn}) keeps only raw text records in
|
||||||
|
memory and defers per-sample tokenization to
|
||||||
|
\texttt{fetch\_record()}, avoiding a pre-tokenized copy
|
||||||
|
entirely---used by DPO/GRPO for on-the-fly training from
|
||||||
|
source files.
|
||||||
\end{itemize}
|
\end{itemize}
|
||||||
|
|
||||||
A resumable distributed sampler provides seed-based shuffle with
|
A resumable distributed sampler provides seed-based shuffle with
|
||||||
@@ -137,13 +154,44 @@ pipeline proceeds as follows:
|
|||||||
previously kept sample $\mathbf{s}'$.
|
previously kept sample $\mathbf{s}'$.
|
||||||
\end{enumerate}
|
\end{enumerate}
|
||||||
|
|
||||||
An optional LLM-as-Judge scoring module provides multi-dimensional
|
A length filter is applied to SFT samples based on the IFD
|
||||||
quality scores that can be used to filter low-quality samples.
|
length-bias analysis in Appendix~\ref{sec:ifd_bias}
|
||||||
|
(Figure~\ref{fig:length_bias}): instruction--response pairs
|
||||||
|
whose response contains fewer than 15 tokens are discarded,
|
||||||
|
because short replies exhibit both high per-token perplexity
|
||||||
|
($L_{\text{uncond}} \approx 6\text{--}8$, PPL~$\approx 400\text{--}3000$)
|
||||||
|
and wide variance in both $L_{\text{cond}}$ and $L_{\text{uncond}}$,
|
||||||
|
which would distort downstream IFD-based difficulty estimates.
|
||||||
|
The threshold is applied per-field, analogous to the pretraining
|
||||||
|
filter described above.
|
||||||
|
|
||||||
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 +279,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 +293,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=2.0\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}
|
||||||
|
|
||||||
@@ -260,7 +310,7 @@ Total steps & 19,000 \\
|
|||||||
\centering
|
\centering
|
||||||
\includegraphics[width=0.50\linewidth]{data/loss_compare.png}
|
\includegraphics[width=0.50\linewidth]{data/loss_compare.png}
|
||||||
\caption{Training loss curves: GPT-2 residual scaling vs.~Kaiming
|
\caption{Training loss curves: GPT-2 residual scaling vs.~Kaiming
|
||||||
initialization over $\sim$20B tokens.}
|
initialization over $\sim$5B tokens.}
|
||||||
\label{fig:loss}
|
\label{fig:loss}
|
||||||
\end{figure}
|
\end{figure}
|
||||||
|
|
||||||
@@ -275,44 +325,90 @@ Figure~\ref{fig:ckpt_comparison} compares four configurations. The left panel sh
|
|||||||
|
|
||||||
\begin{figure}[H]
|
\begin{figure}[H]
|
||||||
\centering
|
\centering
|
||||||
\includegraphics[width=0.85\linewidth]{data/muon_pt.png}
|
\includegraphics[width=0.85\linewidth]{data/pt_metric.png}
|
||||||
\caption{Muon optimizer training dynamics: training loss and learning
|
\caption{Muon optimizer training dynamics: training loss and learning
|
||||||
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:pt_metric}
|
||||||
\end{figure}
|
\end{figure}
|
||||||
|
|
||||||
Figure~\ref{fig:muon_pt} shows the extended training dynamics of the
|
Figure~\ref{fig:pt_metric} 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
|
||||||
\includegraphics[width=0.85\linewidth]{data/sft_metrics.png}
|
\includegraphics[width=0.85\linewidth]{data/sft_metric.png}
|
||||||
\caption{SFT training metrics over 1{,}000 fine-tuning steps on a
|
\caption{SFT training metrics over $\sim$3{,}800 fine-tuning steps on a
|
||||||
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 smoothed loss decreases from $\sim$2.1 to
|
||||||
steps and then enters a slower decay phase. The learning rate
|
$\sim$1.5, dropping rapidly in the first 500 steps and then entering a
|
||||||
follows a cosine schedule with 2\% warmup. Gradient norms stabilize
|
slower decay phase. The learning rate follows a cosine schedule with short linear warmup. Gradient norms stabilise after approximately 500 steps,
|
||||||
after approximately 300 steps, indicating that the fine-tuning
|
indicating that the fine-tuning process has reached a stable optimization
|
||||||
process has reached a stable optimization regime.}
|
regime.}
|
||||||
\label{fig:sft_metrics}
|
\label{fig:sft_metric}
|
||||||
\end{figure}
|
\end{figure}
|
||||||
|
|
||||||
Figure~\ref{fig:sft_metrics} shows the supervised fine-tuning metrics
|
Figure~\ref{fig:sft_metric} shows the supervised fine-tuning metrics
|
||||||
for the 1K-step SFT checkpoint used in the IFD analysis
|
for the SFT checkpoint used in the IFD analysis
|
||||||
(Appendix~\ref{app:ifd}). The training loss on the mixed
|
(Appendix~\ref{app:ifd}). The training loss on the mixed
|
||||||
Chinese--English instruction dataset decreases from $\sim$2.5 to
|
Chinese--English instruction dataset decreases from $\sim$2.1 to
|
||||||
$\sim$1.6 over 1{,}000 steps, with the gradient norm converging to a
|
$\sim$1.5 over $\sim$3{,}800 steps, with the gradient norm converging to a
|
||||||
stable range after the warmup phase.
|
stable range after the warmup phase.
|
||||||
|
|
||||||
|
\subsection{Direct Preference Optimization}
|
||||||
|
\label{sec:dpo}
|
||||||
|
|
||||||
|
Following supervised fine-tuning, we align the model with pairwise
|
||||||
|
human preferences via Direct Preference Optimization
|
||||||
|
(DPO)~\cite{rafailov2023dpo}. DPO avoids explicit reward-model training
|
||||||
|
by optimizing the policy $\pi$ directly against a frozen reference
|
||||||
|
policy $\pi_{\text{ref}}$ (the SFT checkpoint). For a prompt $x$ with
|
||||||
|
preferred response $y_w$ and dispreferred response $y_l$, the loss is:
|
||||||
|
\begin{equation}
|
||||||
|
\mathcal{L}_{\text{DPO}} = -\log \sigma\!\left(
|
||||||
|
\beta \Bigl[
|
||||||
|
\log\tfrac{\pi(y_w\mid x)}{\pi_{\text{ref}}(y_w\mid x)}
|
||||||
|
-\log\tfrac{\pi(y_l\mid x)}{\pi_{\text{ref}}(y_l\mid x)}
|
||||||
|
\Bigr]\right),
|
||||||
|
\end{equation}
|
||||||
|
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
|
||||||
|
same hybrid Muon/AdamW optimizer. The pretraining phase employs WSD (warmup--stable--decay) scheduling
|
||||||
|
to maintain a long high-learning-rate plateau. Both SFT and DPO
|
||||||
|
alignment use cosine schedules with short linear warmup. The shorter ~3{,}000-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_metric}.
|
||||||
|
|
||||||
|
\begin{figure}[H]
|
||||||
|
\centering
|
||||||
|
\includegraphics[width=0.95\linewidth]{data/dpo_metric.png}
|
||||||
|
\caption{DPO training metrics over $\sim$3{,}000 alignment steps:
|
||||||
|
preference loss (raw and 100-step moving average), learning rate
|
||||||
|
schedule, and gradient norm.}
|
||||||
|
\label{fig:dpo_metric}
|
||||||
|
\end{figure}
|
||||||
|
|
||||||
|
Figure~\ref{fig:dpo_metric} summarises the DPO training dynamics.
|
||||||
|
The raw training loss (left panel) starts near $0.7$ and is visibly
|
||||||
|
noisy; the smoothed curve reveals a steady downward trend that reaches
|
||||||
|
$\sim$0.1--0.15 by step 3{,}000 without rebound, indicating stable
|
||||||
|
convergence. The learning-rate schedule (centre panel) peaks at
|
||||||
|
$5\times10^{-6}$ after a short linear warmup and then follows cosine
|
||||||
|
decay to a floor of $\sim$0.5\,$\times\,$10$^{-6}$. Gradient norms (right panel) start
|
||||||
|
near 200 with occasional spikes above 250, then gradually decline and
|
||||||
|
stabilise in the 35--50 range after step 1{,}000, indicating consistent
|
||||||
|
gradient magnitudes throughout alignment.
|
||||||
|
|
||||||
% ======================================================================
|
% ======================================================================
|
||||||
\section{Numerical Stability via Residual Scaling}
|
\section{Numerical Stability via Residual Scaling}
|
||||||
\label{sec:num-stability}
|
\label{sec:num-stability}
|
||||||
@@ -370,7 +466,7 @@ $1/\sqrt{2L}$:
|
|||||||
\end{equation}
|
\end{equation}
|
||||||
|
|
||||||
This reduces per-block residual variance contribution from $0.689$ to
|
This reduces per-block residual variance contribution from $0.689$ to
|
||||||
$0.689/L \approx 0.014$, a factor of $2L = 48$. The post-24-block variance
|
$0.689/(2L) \approx 0.014$, a factor of $2L = 48$. The post-24-block variance
|
||||||
drops from $17.5$ to $1.34$, a $13.1\times$ improvement. In BF16
|
drops from $17.5$ to $1.34$, a $13.1\times$ improvement. In BF16
|
||||||
($7$-bit mantissa, ULP $= 0.0078$ at $w = 1.0$)~\cite{ieee754},
|
($7$-bit mantissa, ULP $= 0.0078$ at $w = 1.0$)~\cite{ieee754},
|
||||||
this keeps weight magnitudes within stable precision bounds. We further
|
this keeps weight magnitudes within stable precision bounds. We further
|
||||||
@@ -416,7 +512,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
|
||||||
@@ -428,24 +524,25 @@ are visible in the per-component weight std
|
|||||||
(Table~\ref{tab:weight_std}, Appendix~\ref{app:weight_std}):
|
(Table~\ref{tab:weight_std}, Appendix~\ref{app:weight_std}):
|
||||||
|
|
||||||
\begin{itemize}[nosep]
|
\begin{itemize}[nosep]
|
||||||
\item \textbf{Training duration drives variance growth}: the
|
\item \textbf{Muon produces larger post-convergence weight
|
||||||
\texttt{muon-25bt} checkpoint exhibits the largest weight std
|
variance}: the \texttt{muon-25bt} checkpoint (Muon, 25B tokens)
|
||||||
($\sim$0.024 for attention projections), exceeding both 15B
|
exhibits the largest weight std ($\sim$0.024), exceeding both
|
||||||
checkpoints ($\sim$0.015--0.021), reflecting continued weight
|
15B AdamW checkpoints ($\sim$0.015--0.021), consistent with
|
||||||
drift from $\sigma_0 = 0.02$ as training progresses.
|
Muon allowing wider parameter distributions after convergence.
|
||||||
\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.
|
||||||
this per-step constraint, producing the largest overall std.
|
|
||||||
\end{itemize}
|
\end{itemize}
|
||||||
|
|
||||||
Critically, the residual-scaled projections remain bounded at
|
Critically, the residual-scaled projections maintain consistently
|
||||||
$\sigma \approx 0.003$ across all checkpoints regardless of training
|
lower standard deviations than their non-scaled counterparts across
|
||||||
duration, confirming that the $1/\sqrt{2L}$ scaling continues to
|
all three checkpoints (Table~\ref{tab:weight_std}), confirming that
|
||||||
enforce its design constraint throughout training.
|
the $1/\sqrt{2L}$ scaling continues to enforce its design constraint
|
||||||
|
throughout training.
|
||||||
|
|
||||||
\begin{figure}[H]
|
\begin{figure}[H]
|
||||||
\centering
|
\centering
|
||||||
@@ -458,12 +555,6 @@ spread.}
|
|||||||
\label{fig:ckpt_weight_density}
|
\label{fig:ckpt_weight_density}
|
||||||
\end{figure}
|
\end{figure}
|
||||||
|
|
||||||
\begin{figure}[H]
|
|
||||||
\centering
|
|
||||||
\includegraphics[width=0.95\linewidth]{data/ckpt_weight_density_per_run.png}
|
|
||||||
\caption{Per-checkpoint weight density breakdowns.}
|
|
||||||
\label{fig:ckpt_weight_density_per_run}
|
|
||||||
\end{figure}
|
|
||||||
|
|
||||||
% ======================================================================
|
% ======================================================================
|
||||||
\section{Conclusion}
|
\section{Conclusion}
|
||||||
@@ -472,15 +563,31 @@ spread.}
|
|||||||
We have described the end-to-end pipeline for training a 1.2B Transformer with
|
We have described the end-to-end pipeline for training a 1.2B Transformer with
|
||||||
{\sc AstrAI}: data preprocessing with JSON-driven tokenization and packing,
|
{\sc AstrAI}: data preprocessing with JSON-driven tokenization and packing,
|
||||||
a 24-layer GQA-SwiGLU architecture, callback-based training with a hybrid
|
a 24-layer GQA-SwiGLU architecture, callback-based training with a hybrid
|
||||||
Muon/AdamW optimizer under DDP/FSDP executors, and cosine scheduling. We
|
Muon/AdamW optimizer under DDP/FSDP executors, and WSD scheduling. We
|
||||||
further analyzed numerical stability under BF16, showing that GPT-2 residual
|
further analyzed numerical stability under BF16, showing that GPT-2 residual
|
||||||
scaling ($\sigma_o = 0.02/\sqrt{2L}$) reduces per-block residual variance
|
scaling ($\sigma_o = 0.02/\sqrt{2L}$) reduces per-block residual variance
|
||||||
by a factor of 48, keeping post-24-layer variance at $1.34$ versus $17.5$
|
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 consistently lower standard deviations than
|
||||||
regardless of optimizer or training duration. An SVD-based effective rank
|
non-scaled weights across all checkpoints
|
||||||
analysis (Appendix~\ref{app:eff_rank}) further reveals that the model
|
(Table~\ref{tab:weight_std}).
|
||||||
|
|
||||||
|
Supervised fine-tuning on deduplicated bilingual instructions (processed
|
||||||
|
by the companion {\sc Alembic} pipeline with MinHash deduplication)
|
||||||
|
reduces training loss from $\sim$2.1 to $\sim$1.5 over $\sim$3{,}800~cosine-scheduled
|
||||||
|
steps. Subsequent DPO alignment on preference pairs ($\beta=0.1$, cosine
|
||||||
|
schedule) shows stable training-loss convergence without over-optimisation.
|
||||||
|
|
||||||
|
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}.
|
||||||
|
|
||||||
@@ -668,7 +775,7 @@ selection signal without re-evaluating after fine-tuning.
|
|||||||
|
|
||||||
To assess how well the trained parameters utilize their allocated
|
To assess how well the trained parameters utilize their allocated
|
||||||
capacity, we perform an SVD-based effective rank analysis on three
|
capacity, we perform an SVD-based effective rank analysis on three
|
||||||
checkpoints: \texttt{kami-15bt} (Muon, 15B tokens), \texttt{norm-15bt}
|
checkpoints: \texttt{kami-15bt} (AdamW, GPT-2 residual scaling, 15B tokens), \texttt{norm-15bt}
|
||||||
(Normal init, 15B tokens), and \texttt{muon-25bt} (Muon, 25B tokens).
|
(Normal init, 15B tokens), and \texttt{muon-25bt} (Muon, 25B tokens).
|
||||||
For each 2D weight matrix $\mathbf{W} \in \mathbb{R}^{m\times n}$ with
|
For each 2D weight matrix $\mathbf{W} \in \mathbb{R}^{m\times n}$ with
|
||||||
SVD $\mathbf{W} = \mathbf{U}\boldsymbol{\Sigma}\mathbf{V}^{\mkern-1mu\mathsf{T}}$,
|
SVD $\mathbf{W} = \mathbf{U}\boldsymbol{\Sigma}\mathbf{V}^{\mkern-1mu\mathsf{T}}$,
|
||||||
@@ -734,9 +841,10 @@ distribution analysis in Section~\ref{sec:num-stability}.
|
|||||||
|
|
||||||
\begin{table}[H]
|
\begin{table}[H]
|
||||||
\centering
|
\centering
|
||||||
\caption{Weight std by component across checkpoints. Non-scaled
|
\caption{Weight std by component across checkpoints. Non-scaled weights
|
||||||
weights broaden with training duration; Muon constrains drift at
|
broaden with training; residual scaling constrains drift at equal
|
||||||
equal token count (15B) but is overtaken by longer training (25B).
|
token count (15B). Muon produces larger post-convergence weight
|
||||||
|
variance than AdamW.
|
||||||
Residual-scaled projections ($\mathbf{W}_o$, $\mathbf{W}_{\text{down}}$)
|
Residual-scaled projections ($\mathbf{W}_o$, $\mathbf{W}_{\text{down}}$)
|
||||||
remain bounded.}
|
remain bounded.}
|
||||||
\label{tab:weight_std}
|
\label{tab:weight_std}
|
||||||
@@ -744,7 +852,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 \\
|
||||||
@@ -808,6 +916,11 @@ A.~Radford, J.~Wu, R.~Child, D.~Luan, D.~Amodei, I.~Sutskever.
|
|||||||
Language models are unsupervised multitask learners.
|
Language models are unsupervised multitask learners.
|
||||||
\textit{OpenAI Blog}, 2019.
|
\textit{OpenAI Blog}, 2019.
|
||||||
|
|
||||||
|
\bibitem{rafailov2023dpo}
|
||||||
|
R.~Rafailov, A.~Sharma, E.~Mitchell, C.~D.~Manning, S.~Ermon, C.~Finn.
|
||||||
|
Direct Preference Optimization: Your language model is secretly a reward model.
|
||||||
|
\textit{NeurIPS}, 2023.
|
||||||
|
|
||||||
\bibitem{shazeer2020glu}
|
\bibitem{shazeer2020glu}
|
||||||
N.~Shazeer. GLU variants improve Transformer.
|
N.~Shazeer. GLU variants improve Transformer.
|
||||||
\textit{arXiv:2002.05202}, 2020.
|
\textit{arXiv:2002.05202}, 2020.
|
||||||
|
|||||||
Reference in New Issue
Block a user