Compare commits

...
13 Commits
Author SHA1 Message Date
ViperEkura e6be33aa53 doc: add SFT length-filter rationale — high PPL + high variance from A.3 IFD figure
- Add 15-token length floor to SFT samples, with explicit reference
  to Appendix A.3 / Figure 5 (ifd_length_grid)
- Short replies (<10 tokens) show both high per-token perplexity
  (L_uncond ~6-8, PPL ~400-3000 vs long replies ~2-3, PPL ~7-20)
  and wide variance in L_cond (span 0-17.5), which distorts
  downstream IFD-based difficulty estimates.
2026-07-24 06:37:35 +08:00
ViperEkura 0a1d0573ae doc: add SFT length filter (15-token floor) with IFD bias rationale
Per-field length filtering in SFT drops instruction--response pairs
with responses shorter than 15 tokens. Short replies exhibit high
per-token variance in both conditional and unconditional loss
(Appendix A.3 / Figure 5), which would distort downstream IFD-based
difficulty estimates.
2026-07-24 06:27:10 +08:00
ViperEkura 0c2bc916f2 fix: loss_compare caption token count (20B -> 5B)
The loss_compare.png x-axis only spans 0-5B tokens. The caption
incorrectly claimed ~20B tokens.
2026-07-24 06:21:52 +08:00
ViperEkura 4e70e827ff remove: drop ckpt_weight_density_per_run.png figure
The per-run weight density figure contained misleading legacy
iteration labels (500k/1M iter) that contradicted the paper's
stated token budgets (15B). Removing it avoids confusion; the
per-category density plot (ckpt_weight_density.png) and Table 7
remain as the primary evidence.
2026-07-24 06:20:08 +08:00
ViperEkura a7bbc7b29f fix: align SFT/DPO figures and text with actual training data
- SFT: 1,000 steps/WSD → ~3,800 steps/cosine; loss ~2.5→1.6 → ~2.1→1.5
- DPO: fix preference-loss narrative to match training-loss curve;
  correct initial grad-norm (~50 → ~200)
- Table 2: peak LR 1.5e-4 → 2.0e-4 (matches pt_metric.png)
- Clarify scheduling: pretraining=WSD, SFT+DPO=cosine
- Add disclaimer to ckpt_weight_density_per_run caption for legacy iter labels
2026-07-24 06:15:42 +08:00
ViperEkura 4b37f289c0 fix: attribute muon-25bt larger std to Muon characteristic, not training duration alone
The muon-25bt (Muon, 25B) vs kami/norm (AdamW, 15B) comparison confounds
optimizer choice with training duration. Reframe as "Muon produces larger
post-convergence weight variance" rather than "training duration drives
variance growth."
2026-07-20 11:59:09 +08:00
ViperEkura a83555f326 fix: align conclusion, abstract, introduction with body data and full paper scope
- Fix sigma=0.003 (init value) to consistently narrower than non-scaled
  (body 4.4, conclusion, abstract)
- Add SFT/DPO to conclusion, abstract, and introduction
- Fix checkpoint label: kami-15bt (Muon) -> (AdamW, GPT-2 residual scaling)
  in effective rank appendix
- Fix table weight_std caption: Muon -> residual scaling constrains drift
- Fix 4.2 math: 0.689/L -> 0.689/(2L)
- Fix length filter: per-sample -> per-text-field skip
2026-07-20 11:51:48 +08:00
ViperEkura d93ff48320 docs: fix storage backends (2->3) and per-field length filter
- Storage Backends: add JsonlStore (lazy on-the-fly tokenization)
- Describe memory-efficiency tradeoff across three backends
- Fix length filter: per-text-field skip, not per-sample global
- Abstract: HDF5/mmap -> tiered storage backends
- Rename metric figures (singular) and update references
2026-07-20 11:27:39 +08:00
ViperEkura da0f536526 Update DPO iteration count: 1,500 -> 3,000 2026-07-20 03:59:55 +08:00
ViperEkura c775a2b3e0 Update DPO metrics figure 2026-07-20 03:59:04 +08:00
ViperEkura 3f0ff911a8 Fix lingering cosine->WSD in Conclusion; update AGENTS.md with checkpoint labels and scheduler rules 2026-07-19 15:30:49 +08:00
ViperEkura e149997200 Add DPO data generation & training; fix checkpoint labels; switch to WSD scheduler; update optimizer ablations in abstract/conclusion 2026-07-19 15:25:42 +08:00
ViperEkura 6c2e04a86f Add DPO training subsection with dpo_metrics figure and Rafailov et al. citation 2026-07-19 15:01:23 +08:00
7 changed files with 204 additions and 91 deletions
Binary file not shown.

After

Width:  |  Height:  |  Size: 86 KiB

BIN
View File
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

+204 -91
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{}
@@ -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.