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}
|
||||
|
||||
\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}
|
||||
\date{}
|
||||
@@ -29,25 +29,25 @@
|
||||
|
||||
\begin{abstract}
|
||||
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
|
||||
JSON-driven BBPE preprocessing with multi-strategy packing, HDF5/mmap
|
||||
storage backends, and a companion SFT pipeline ({\sc Alembic}) with MinHash
|
||||
deduplication and LLM-as-Judge scoring. The 24-layer decoder uses GQA,
|
||||
SwiGLU, RoPE, and RMSNorm, trained with a hybrid Muon/AdamW optimizer and
|
||||
cosine scheduling under DDP/FSDP. A BF16 stability analysis shows that
|
||||
GPT-2 residual scaling substantially reduces per-block residual variance
|
||||
accumulation, keeping post-training variance well below the overflow
|
||||
threshold of standard initialization; empirically this yields a sustained
|
||||
loss advantage over Kaiming initialization throughout training.
|
||||
Post-training weight distribution analysis across three
|
||||
checkpoints---varying optimizer, initialization, and training
|
||||
budget---confirms that residual-scaled projections maintain narrow
|
||||
distributions throughout training, preserving the numerical stability
|
||||
established at initialization. An SVD-based effective rank analysis
|
||||
further reveals that the model operates near 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
|
||||
of a 1.2B-parameter Transformer on $\sim$25B tokens. The pipeline covers
|
||||
JSON-driven BBPE preprocessing with multi-strategy packing, tiered
|
||||
storage backends, and a companion SFT pipeline ({\sc Alembic}) with
|
||||
MinHash deduplication. The 24-layer GQA-SwiGLU decoder is trained with a
|
||||
hybrid Muon/AdamW optimizer and WSD scheduling under DDP/FSDP.
|
||||
Supervised fine-tuning on deduplicated bilingual instructions reduces loss
|
||||
from $\sim$2.1 to $\sim$1.5 over $\sim$3{,}800~steps; DPO alignment
|
||||
($\beta=0.1$, cosine schedule) on model-generated preference pairs
|
||||
shows stable training-loss convergence without over-optimisation. A BF16 stability analysis
|
||||
shows that GPT-2 residual scaling ($\sigma_0 = 0.02/\sqrt{2L}$) reduces
|
||||
per-block activation variance by a factor of 48, and post-training weight
|
||||
analysis across three checkpoints confirms that residual-scaled
|
||||
projections remain consistently narrower than non-scaled weights across
|
||||
optimizers, initializations, and training budgets. Optimizer ablations
|
||||
demonstrate that the hybrid Muon/AdamW outperforms pure AdamW, with
|
||||
2D weight matrices benefiting from Muon's orthogonalisation and 1D
|
||||
parameters from AdamW's second-moment adaptation. An SVD effective-rank
|
||||
analysis reveals near-capacity weight utilization, with Q/O projections
|
||||
showing consistently lower effective rank than K/V projections under
|
||||
grouped query attention.
|
||||
\end{abstract}
|
||||
|
||||
@@ -60,8 +60,12 @@ model architecture. Data must be preprocessed and stored efficiently, the
|
||||
training loop must handle distributed parallelism, gradient accumulation,
|
||||
checkpointing, and logging---and numerical pitfalls must be diagnosed and
|
||||
fixed. This paper describes the complete workflow using {\sc AstrAI}~\cite{astrai}, an
|
||||
open-source framework for Transformer training and inference, and highlights a
|
||||
BF16 precision issue encountered along the way.
|
||||
open-source framework for Transformer training and inference, from JSONL
|
||||
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}
|
||||
@@ -86,17 +90,30 @@ via a JSON specification that defines:
|
||||
\texttt{.bin} shards, auto-split at 100M tokens per shard.
|
||||
\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}
|
||||
|
||||
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]
|
||||
\item \textbf{H5Store}: HDF5-based, memory-loaded with shared-memory support
|
||||
for multi-worker access.
|
||||
\item \textbf{MmapStore}: Zero-copy memory-mapped \texttt{.bin} files shared
|
||||
via OS page cache.
|
||||
\item \textbf{H5Store}: HDF5-based, fully loaded into RAM at init
|
||||
with \texttt{share\_memory\_()} for cross-worker sharing.
|
||||
Fastest random access; requires the dataset to fit in memory.
|
||||
\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}
|
||||
|
||||
A resumable distributed sampler provides seed-based shuffle with
|
||||
@@ -137,13 +154,44 @@ pipeline proceeds as follows:
|
||||
previously kept sample $\mathbf{s}'$.
|
||||
\end{enumerate}
|
||||
|
||||
An optional LLM-as-Judge scoring module provides multi-dimensional
|
||||
quality scores that can be used to filter low-quality samples.
|
||||
A length filter is applied to SFT samples based on the IFD
|
||||
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
|
||||
Appendix~\ref{app:ifd}.
|
||||
\subsection{DPO Data Generation}
|
||||
|
||||
% ======================================================================
|
||||
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}
|
||||
% ======================================================================
|
||||
|
||||
@@ -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).
|
||||
\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
|
||||
scheduling (2\% warmup) and global L2 gradient clipping. The framework supports DDP and FSDP for multi-GPU distribution,
|
||||
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)
|
||||
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.
|
||||
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} \\
|
||||
\midrule
|
||||
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$ \\
|
||||
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 \\
|
||||
Sequence length & 2,048 tokens \\
|
||||
Total steps & 19,000 \\
|
||||
Total tokens & $\sim$25B ($\approx$23k steps) \\
|
||||
\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{table}
|
||||
|
||||
@@ -260,7 +310,7 @@ Total steps & 19,000 \\
|
||||
\centering
|
||||
\includegraphics[width=0.50\linewidth]{data/loss_compare.png}
|
||||
\caption{Training loss curves: GPT-2 residual scaling vs.~Kaiming
|
||||
initialization over $\sim$20B tokens.}
|
||||
initialization over $\sim$5B tokens.}
|
||||
\label{fig:loss}
|
||||
\end{figure}
|
||||
|
||||
@@ -275,44 +325,90 @@ Figure~\ref{fig:ckpt_comparison} compares four configurations. The left panel sh
|
||||
|
||||
\begin{figure}[H]
|
||||
\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
|
||||
rate schedule for the hybrid Muon/AdamW configuration across $\sim$25B
|
||||
tokens. The Muon optimizer is applied to all 2D weight matrices
|
||||
(attention projections and FFN layers), while AdamW handles 1D
|
||||
parameters (embeddings, biases, and normalization scales). The cosine
|
||||
schedule with 2\% warmup is visible in the lower panel.}
|
||||
\label{fig:muon_pt}
|
||||
parameters (embeddings, biases, and normalization scales). The WSD
|
||||
schedule (warmup, stable phase, and final decay) is visible in the lower panel.}
|
||||
\label{fig:pt_metric}
|
||||
\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
|
||||
exhibits the expected power-law decay in the early phase (0--5B tokens),
|
||||
followed by a gradual plateau as the model approaches convergence on
|
||||
the pretraining distribution. The cosine learning rate schedule
|
||||
reaches its minimum at the end of training, ensuring stable weight
|
||||
updates in the final phase.
|
||||
the pretraining distribution. The WSD schedule holds the learning rate
|
||||
constant during the long stable phase and then decays at the end of
|
||||
training, ensuring stable weight updates in the final phase.
|
||||
|
||||
\begin{figure}[H]
|
||||
\centering
|
||||
\includegraphics[width=0.85\linewidth]{data/sft_metrics.png}
|
||||
\caption{SFT training metrics over 1{,}000 fine-tuning steps on a
|
||||
\includegraphics[width=0.85\linewidth]{data/sft_metric.png}
|
||||
\caption{SFT training metrics over $\sim$3{,}800 fine-tuning steps on a
|
||||
mixed Chinese--English instruction dataset: training loss, learning
|
||||
rate, and gradient norm. The loss drops rapidly in the first 200
|
||||
steps and then enters a slower decay phase. The learning rate
|
||||
follows a cosine schedule with 2\% warmup. Gradient norms stabilize
|
||||
after approximately 300 steps, indicating that the fine-tuning
|
||||
process has reached a stable optimization regime.}
|
||||
\label{fig:sft_metrics}
|
||||
rate, and gradient norm. The smoothed loss decreases from $\sim$2.1 to
|
||||
$\sim$1.5, dropping rapidly in the first 500 steps and then entering a
|
||||
slower decay phase. The learning rate follows a cosine schedule with short linear warmup. Gradient norms stabilise after approximately 500 steps,
|
||||
indicating that the fine-tuning process has reached a stable optimization
|
||||
regime.}
|
||||
\label{fig:sft_metric}
|
||||
\end{figure}
|
||||
|
||||
Figure~\ref{fig:sft_metrics} shows the supervised fine-tuning metrics
|
||||
for the 1K-step SFT checkpoint used in the IFD analysis
|
||||
Figure~\ref{fig:sft_metric} shows the supervised fine-tuning metrics
|
||||
for the SFT checkpoint used in the IFD analysis
|
||||
(Appendix~\ref{app:ifd}). The training loss on the mixed
|
||||
Chinese--English instruction dataset decreases from $\sim$2.5 to
|
||||
$\sim$1.6 over 1{,}000 steps, with the gradient norm converging to a
|
||||
Chinese--English instruction dataset decreases from $\sim$2.1 to
|
||||
$\sim$1.5 over $\sim$3{,}800 steps, with the gradient norm converging to a
|
||||
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}
|
||||
\label{sec:num-stability}
|
||||
@@ -370,7 +466,7 @@ $1/\sqrt{2L}$:
|
||||
\end{equation}
|
||||
|
||||
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
|
||||
($7$-bit mantissa, ULP $= 0.0078$ at $w = 1.0$)~\cite{ieee754},
|
||||
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
|
||||
training---not just at initialization---we compare weight value
|
||||
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).
|
||||
|
||||
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}):
|
||||
|
||||
\begin{itemize}[nosep]
|
||||
\item \textbf{Training duration drives variance growth}: the
|
||||
\texttt{muon-25bt} checkpoint exhibits the largest weight std
|
||||
($\sim$0.024 for attention projections), exceeding both 15B
|
||||
checkpoints ($\sim$0.015--0.021), reflecting continued weight
|
||||
drift from $\sigma_0 = 0.02$ as training progresses.
|
||||
\item \textbf{Muon constrains early-stage drift}: at equal training
|
||||
budget (15B tokens), the Muon-trained \texttt{kami-15bt} shows
|
||||
\item \textbf{Muon produces larger post-convergence weight
|
||||
variance}: the \texttt{muon-25bt} checkpoint (Muon, 25B tokens)
|
||||
exhibits the largest weight std ($\sim$0.024), exceeding both
|
||||
15B AdamW checkpoints ($\sim$0.015--0.021), consistent with
|
||||
Muon allowing wider parameter distributions after convergence.
|
||||
\item \textbf{Residual scaling constrains early-stage drift}: at
|
||||
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
|
||||
\texttt{norm-15bt} ($\sim$0.021). Muon's orthogonalization
|
||||
step normalizes the update direction, constraining per-step
|
||||
weight movement. At 25B tokens, cumulative steps overtake
|
||||
this per-step constraint, producing the largest overall std.
|
||||
\texttt{norm-15bt} ($\sim$0.021), indicating that residual
|
||||
scaling itself limits weight drift beyond its initialization
|
||||
effect.
|
||||
\end{itemize}
|
||||
|
||||
Critically, the residual-scaled projections remain bounded at
|
||||
$\sigma \approx 0.003$ across all checkpoints regardless of training
|
||||
duration, confirming that the $1/\sqrt{2L}$ scaling continues to
|
||||
enforce its design constraint throughout training.
|
||||
Critically, the residual-scaled projections maintain consistently
|
||||
lower standard deviations than their non-scaled counterparts across
|
||||
all three checkpoints (Table~\ref{tab:weight_std}), confirming that
|
||||
the $1/\sqrt{2L}$ scaling continues to enforce its design constraint
|
||||
throughout training.
|
||||
|
||||
\begin{figure}[H]
|
||||
\centering
|
||||
@@ -458,12 +555,6 @@ spread.}
|
||||
\label{fig:ckpt_weight_density}
|
||||
\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}
|
||||
@@ -472,15 +563,31 @@ spread.}
|
||||
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,
|
||||
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
|
||||
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$
|
||||
without scaling. Post-training weight distribution analysis confirms that
|
||||
this scaling constraint persists throughout training, with residual-scaled
|
||||
projections maintaining narrow distributions ($\sigma \approx 0.003$)
|
||||
regardless of optimizer or training duration. An SVD-based effective rank
|
||||
analysis (Appendix~\ref{app:eff_rank}) further reveals that the model
|
||||
projections maintaining consistently lower standard deviations than
|
||||
non-scaled weights across all checkpoints
|
||||
(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
|
||||
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
|
||||
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).
|
||||
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}}$,
|
||||
@@ -734,9 +841,10 @@ distribution analysis in Section~\ref{sec:num-stability}.
|
||||
|
||||
\begin{table}[H]
|
||||
\centering
|
||||
\caption{Weight std by component across checkpoints. Non-scaled
|
||||
weights broaden with training duration; Muon constrains drift at
|
||||
equal token count (15B) but is overtaken by longer training (25B).
|
||||
\caption{Weight std by component across checkpoints. Non-scaled weights
|
||||
broaden with training; residual scaling constrains drift at equal
|
||||
token count (15B). Muon produces larger post-convergence weight
|
||||
variance than AdamW.
|
||||
Residual-scaled projections ($\mathbf{W}_o$, $\mathbf{W}_{\text{down}}$)
|
||||
remain bounded.}
|
||||
\label{tab:weight_std}
|
||||
@@ -744,7 +852,7 @@ remain bounded.}
|
||||
\begin{tabular}{@{}lccc@{}}
|
||||
\toprule
|
||||
\textbf{Component} & \textbf{kami-15bt} & \textbf{norm-15bt} & \textbf{muon-25bt} \\
|
||||
& (Muon, 15B) & (Normal, 15B) & (Muon, 25B) \\
|
||||
& (AdamW, 15B) & (AdamW, 15B) & (Muon, 25B) \\
|
||||
\midrule
|
||||
attn.q\_proj & 0.0154 & 0.0207 & 0.0230 \\
|
||||
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.
|
||||
\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}
|
||||
N.~Shazeer. GLU variants improve Transformer.
|
||||
\textit{arXiv:2002.05202}, 2020.
|
||||
|
||||
Reference in New Issue
Block a user