Skip to content

程式碼逐行解讀

Pre-training Pipeline(base_train.py)

模型建置

python
def build_model_meta(depth):
    base_dim = depth * args.aspect_ratio
    model_dim = ((base_dim + args.head_dim - 1) // args.head_dim) * args.head_dim
    # 計算 model_dim: nudge up to nearest multiple of head_dim
    config = GPTConfig(sequence_len=..., vocab_size=..., n_layer=depth, n_head=model_dim//head_dim, ...)
    with torch.device("meta"):
        model_meta = GPT(config)  # meta device: 只配置 shape,不分配實際記憶體
    return model_meta

三階段初始化:1) meta device 建構(shape only)→ 2) to_empty() 分配儲存 → 3) init_weights() 初始化權重。

Scaling Laws 自動化

python
target_tokens = int(target_param_data_ratio * num_scaling_params)
# token 數量 = 目標資料:參數比 × 可擴展參數數
# ─ 參考 d12 的實驗結果,推算更大 deep 的最佳配置

total_batch_size = B_REF * (target_tokens / D_REF) ** 0.383
# Power Lines 論文:B_opt ∝ D^0.383

batch_lr_scale = (total_batch_size / B_REF) ** 0.5
# AdamW: η ∝ √(B/B_ref);Muon 沿用相同 scaling

訓練迴圈(核心部分)

python
for micro_step in range(grad_accum_steps):
    loss = model(x, y)                           # forward
    train_loss = loss.detach()
    loss = loss / grad_accum_steps               # normalize for grad accumulation
    loss.backward()                              # backward
    x, y, _ = next(train_loader)                 # prefetch next batch during backward

# 優化器步驟前設定各組 learning rate
for group in optimizer.param_groups:
    group["lr"] = group["initial_lr"] * lrm      # lrm = get_lr_multiplier(step)
    if group['kind'] == 'muon':
        group["momentum"] = get_muon_momentum(step)
        group["weight_decay"] = get_weight_decay(step)
optimizer.step()
model.zero_grad(set_to_none=True)

注意力機制(gpt.py)

CausalSelfAttention

python
q = self.c_q(x).view(B, T, self.n_head, self.head_dim)
k = self.c_k(x).view(B, T, self.n_kv_head, self.head_dim)
v = self.c_v(x).view(B, T, self.n_kv_head, self.head_dim)
# 注意:Q 有 n_head 個 head,K/V 只有 n_kv_head(GQA)

# Value Residual(ResFormer)
if ve is not None:
    gate = 3 * torch.sigmoid(self.ve_gate(x[..., :12]))  # (0, 3) 範圍的 gate
    v = v + gate.unsqueeze(-1) * ve                        # 混合 value embedding

# Rotary Embeddings + QK Norm
q, k = apply_rotary_emb(q, cos, sin), apply_rotary_emb(k, cos, sin)
q, k = norm(q), norm(k)
q = q * 1.2  # Sharper attention

Tokenizer(tokenizer.py)

對話結構渲染:

python
def render_conversation(self, conversation, max_tokens=2048):
    # 1) 處理 system message(合併到第一個 user message)
    # 2) 遍歷 messages,user 用 <|user_start|>...<|user_end|> 包裹
    # 3) assistant 用 <|assistant_start|>...<|assistant_end|> 包裹
    # 4) mask = 1 的 token 才會計算 loss(僅 assistant 的回答)
    # 5) 支援 tool use:<|python_start|>、<|python_end|>、<|output_start|>、<|output_end|>

Dataloader(dataloader.py)

BOS-aligned best-fit packing:

python
# 每個 row 以 BOS token 開頭
# 從 buffer 中挑選「最大可完整放入」的文件
# 若無文件可完整放入 → 裁切最短的文件填滿剩餘空間
# 結果:100% 利用率,~35% tokens 被裁切(丟棄)