Compare commits

..
8 Commits
Author SHA1 Message Date
dela a2c4217dae Document LatentMoE sigmoid routing, sparse dispatch, and K3 block figures
Ledger C10–C12 match the permute-pad-bmm path and Switch aux/z-loss.
Section 8 adds overview and component TikZ; MoE capacity is C_moe so it
does not collide with KDA chunk size.
2026-08-26 14:43:58 +08:00
dela ea7167b3f7 Size vocab by max token id so duplicate-piece vocabs (Yi-6B) don't overflow embedding 2026-08-26 14:42:12 +08:00
dela 94d0f2ff6a Do not let HF tokenizers truncate wiki articles at 4096
Yi-6B sets model_max_length=4096. encode() would clip long Wikipedia
pages before we pack seq_len chunks. Raise the cap so only our
chunker limits context.
2026-08-26 14:31:07 +08:00
dela 24c9d56b72 Keep attnres on resume, fix final chunk_index, default 0.5b to Yi-6B
CLI default attnres=off was overwriting block checkpoints on resume
so later loads hit Unexpected key(s). Only apply flags the user
passed. Track next_chunk so a budget-exit save does not skip the
untrained yield. 0.5b now uses 01-ai/Yi-6B (64k); refuse resume
when the ckpt tokenizer does not match.
2026-08-26 14:23:29 +08:00
dela 5cc0555563 Fix chrF fallback so English wiki garbage fails translation_success
Sacrebleu errors used to fall back to set-overlap unigrams, so any
English hyp scored ~70–80 against English refs and SFT reported
success 1.0. Use count-based char n-grams, keep BLEU failures from
clobbering chrF, and print a few hyps during eval.
2026-08-26 14:23:22 +08:00
dela 5a7d949b01 Skip the cold-start SFT best ckpt and free CUDA cache after eval
Step 0 generate was writing a 5GB success=0 snapshot and leaving the
32GB card fragmented, so the next Adam step OOM'd after batch-32 eval.
Only promote _best after step 0 and empty_cache when eval returns.
2026-08-26 10:08:31 +08:00
dela 9652a9a7eb Save SFT last/best checkpoints during training and on interrupt
train_sft used to torch.save only after the full epoch budget, so Ctrl+C
dropped all translation weights. Write _last every --ckpt-every steps
and on KeyboardInterrupt; write _best when frozen eval (success, chrF)
improves; --resume continues from _last.
2026-08-26 10:00:11 +08:00
dela 071dfaf42c Sanitize SwanLab env before login so 0.9 nested project does not crash
OpenBayes sets SWANLAB_PROJECT as a string; swanlab>=0.9 parses that as
ProjectSettings and raises QuoteAwareEnvSettingsSource. Drop it, keep
SWANLAB_PROJ_NAME, and share run-id extraction with train_k3.
2026-08-26 10:00:06 +08:00
19 changed files with 777 additions and 169 deletions
+2 -2
View File
@@ -138,7 +138,7 @@ PYTHONPATH=. python train.py
### 复现训练 ### 复现训练
目标是 **~0.5B zh↔en 指令翻译模型**(`K3Config.preset("0.5b")` = 482M,tied Qwen3 词表)。本机 RTX 3060 6GB 只跑 8M 全流程孪生;0.5B 预训练需要 **32–40GB Ampere bf16**。 目标是 **~0.5B zh↔en 指令翻译模型**(`K3Config.preset("0.5b")` ≈ 415M,tied Yi-6B 64k 词表)。本机 RTX 3060 6GB 只跑 8M 全流程孪生;0.5B 预训练需要 **32–40GB Ampere bf16**。
成功标准是冻结集上的 `translation_success()`,**不是** wiki train loss。wiki 预训练没见过 `Translate to English:\n...`,预训练阶段 `eval_mt` 的 success_rate 预期 ≈0。 成功标准是冻结集上的 `translation_success()`,**不是** wiki train loss。wiki 预训练没见过 `Translate to English:\n...`,预训练阶段 `eval_mt` 的 success_rate 预期 ≈0。
@@ -160,7 +160,7 @@ uv run python train_k3.py --preset 0.5b --attnres block \
--max-tokens 1000000000 --warmup 2000 --max-tokens 1000000000 --warmup 2000
``` ```
`0.5b` 预设:`d=768`,`L=24`(6×3 KDA + 1 MLA),`H=12`,`head=64`,`chunk=64`,LatentMoE `ℓ=384` / 16 routed / Top-2 / shared 2,tied Qwen3 embedding,activation checkpoint 默认开。训练默认 seq 2048、micro-batch 2、grad-acc 8、lr 3e-4、warmup **64 optimizer steps**。checkpoint:`ckpts/k3_0.5b.pt`,另写 `_last` / `_best`(best 按 held-out CE)。 `0.5b` 预设:`d=768`,`L=24`(6×3 KDA + 1 MLA),`H=12`,`head=64`,`chunk=64`,LatentMoE `ℓ=384` / 16 routed / Top-2 / shared 2,tied Yi-6B embedding(64k,有 EOS),activation checkpoint 默认开。训练默认 seq 2048、micro-batch 2、grad-acc 8、lr 3e-4、warmup **64 optimizer steps**。checkpoint:`ckpts/k3_0.5b.pt`,另写 `_last` / `_best`(best 按 held-out CE)。**不能**从 Qwen3 词表的旧 ckpt `--resume`。
Token 会计:`step` 仍是 micro-batch;有效 token = `batch × seq_len × micro_steps`。默认 0.5b 冒烟是 **8.2M token ≈ 0.017 tok/param**。翻译前置 LM 的最低有意义预算是 **1B token**(`--max-tokens`),不是 2000 step。 Token 会计:`step` 仍是 micro-batch;有效 token = `batch × seq_len × micro_steps`。默认 0.5b 冒烟是 **8.2M token ≈ 0.017 tok/param**。翻译前置 LM 的最低有意义预算是 **1B token**(`--max-tokens`),不是 2000 step。
+7 -5
View File
@@ -9,13 +9,15 @@ Hybrid Attention (K3): 每 4 层 1 次 Gated MLA, 末层强制 MLA.
Presets: Presets:
toy — ~8M, 自训 8k SP, 本地过拟合 toy — ~8M, 自训 8k SP, 本地过拟合
0.5b — ~482M, Qwen3 词表, 32–40GB bf16;默认 step 是冒烟,翻译前置用 --max-tokens 0.5b — ~415M, Yi-6B 词表 (64k), 32–40GB bf16;默认 step 是冒烟,翻译前置用 --max-tokens
""" """
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
# Qwen3 config.json; train_k3 overrides with len(tokenizer). # 01-ai/Yi-6B config.json; train_k3 overrides with len(tokenizer).
YI6B_VOCAB_SIZE = 64000
# Kept for old Qwen3 checkpoints / docs.
QWEN3_VOCAB_SIZE = 151936 QWEN3_VOCAB_SIZE = 151936
@@ -24,7 +26,7 @@ class K3Config:
# 主干 # 主干
hidden_size: int = 256 hidden_size: int = 256
num_hidden_layers: int = 4 num_hidden_layers: int = 4
vocab_size: int = 8192 # toy: data/spm_4k; 0.5b: Qwen3 vocab_size: int = 8192 # toy: data/spm_4k; 0.5b: Yi-6B 64k
initializer_range: float = 0.02 initializer_range: float = 0.02
norm_eps: float = 1e-6 norm_eps: float = 1e-6
tie_word_embeddings: bool = False tie_word_embeddings: bool = False
@@ -76,11 +78,11 @@ class K3Config:
return cls() return cls()
if name in {"0.5b", "500m"}: if name in {"0.5b", "500m"}:
# H * head_dim == hidden. Routed 16 Top-2; LatentMoE padded bmm. # H * head_dim == hidden. Routed 16 Top-2; LatentMoE padded bmm.
# ~482M with tied Qwen3 embeddings. 6×(3 KDA + 1 MLA). # ~415M with tied Yi-6B embeddings. 6×(3 KDA + 1 MLA).
return cls( return cls(
hidden_size=768, hidden_size=768,
num_hidden_layers=24, num_hidden_layers=24,
vocab_size=QWEN3_VOCAB_SIZE, vocab_size=YI6B_VOCAB_SIZE,
tie_word_embeddings=True, tie_word_embeddings=True,
max_position_embeddings=2048, max_position_embeddings=2048,
num_heads=12, num_heads=12,
+7 -1
View File
@@ -57,7 +57,11 @@ class HuggingFaceTokenizer:
@property @property
def vocab_size(self) -> int: def vocab_size(self) -> int:
return int(len(self._tok)) # len(tok) 数的是去重后的 surface form; 词表有重复 piece 时 (如 Yi-6B
# 63992 vs 最大 id 63999) 会小于真实 id 范围, embedding 越界触发
# device-side assert. 以最大 id + 1 为准.
max_id = max(self._tok.get_vocab().values())
return max(int(len(self._tok)), max_id + 1)
def encode(self, text: str) -> list[int]: def encode(self, text: str) -> list[int]:
return list(self._tok.encode(text, add_special_tokens=False)) return list(self._tok.encode(text, add_special_tokens=False))
@@ -75,6 +79,8 @@ def load_tokenizer(source: str) -> Tokenizer:
from transformers import AutoTokenizer from transformers import AutoTokenizer
tok = AutoTokenizer.from_pretrained(source, trust_remote_code=True) tok = AutoTokenizer.from_pretrained(source, trust_remote_code=True)
# 只借词表分词, 语料随后按 seq_len 切块, 不受原模型 4096 上限约束
tok.model_max_length = 10**9
return HuggingFaceTokenizer(tok) return HuggingFaceTokenizer(tok)
+9 -3
View File
@@ -67,14 +67,18 @@ def evaluate_pairs(
want = "zh" if target_lang.startswith("zh") else "en" want = "zh" if target_lang.startswith("zh") else "en"
lang_ok += int(_detect_lang(hyp) == want) lang_ok += int(_detect_lang(hyp) == want)
chrf_sum += _chrf(hyp, ref) chrf_sum += _chrf(hyp, ref)
corpus = {} corpus = {"chrf": chrf_sum / max(n, 1), "bleu": None}
try: try:
from sacrebleu.metrics import BLEU, CHRF from sacrebleu.metrics import CHRF
corpus["chrf"] = float(CHRF(word_order=2).corpus_score(hyps, [refs[:n]]).score) corpus["chrf"] = float(CHRF(word_order=2).corpus_score(hyps, [refs[:n]]).score)
except Exception:
pass
try:
from sacrebleu.metrics import BLEU
corpus["bleu"] = float(BLEU().corpus_score(hyps, [refs[:n]]).score) corpus["bleu"] = float(BLEU().corpus_score(hyps, [refs[:n]]).score)
except Exception: except Exception:
corpus["chrf"] = chrf_sum / max(n, 1)
corpus["bleu"] = None corpus["bleu"] = None
return { return {
"n": n, "n": n,
@@ -131,6 +135,8 @@ def main() -> None:
) )
printable = {k: v for k, v in out.items() if k != "hyps"} printable = {k: v for k, v in out.items() if k != "hyps"}
print(json.dumps(printable, ensure_ascii=False, indent=2)) print(json.dumps(printable, ensure_ascii=False, indent=2))
for i, hyp in enumerate(out["hyps"][: min(5, out["n"])]):
print(f" [{i}] {hyp}")
elif not args.prefix: elif not args.prefix:
raise SystemExit("pass --prefix and/or --src + --ref") raise SystemExit("pass --prefix and/or --src + --ref")
+27 -11
View File
@@ -28,8 +28,33 @@ def _detect_lang(text: str) -> str | None:
return tag[:2] return tag[:2]
def _chrf_ngram(hyp: str, ref: str, max_n: int = 4) -> float:
"""Count-based char n-gram F (β=2), 0–100. Not set-overlap unigrams."""
from collections import Counter
hyp, ref = hyp.strip(), ref.strip()
if not hyp or not ref:
return 0.0
scores: list[float] = []
for n in range(1, max_n + 1):
if len(hyp) < n or len(ref) < n:
scores.append(0.0)
continue
hc = Counter(hyp[i : i + n] for i in range(len(hyp) - n + 1))
rc = Counter(ref[i : i + n] for i in range(len(ref) - n + 1))
overlap = sum((hc & rc).values())
prec = overlap / max(sum(hc.values()), 1)
rec = overlap / max(sum(rc.values()), 1)
if prec + rec == 0:
scores.append(0.0)
continue
beta2 = 4.0
scores.append((1.0 + beta2) * prec * rec / (beta2 * prec + rec))
return 100.0 * sum(scores) / max(len(scores), 1)
def _chrf(hyp: str, ref: str) -> float: def _chrf(hyp: str, ref: str) -> float:
"""chrF++ in 0–100. Falls back to char unigram F if sacrebleu is missing.""" """chrF++ in 0–100. Falls back to count-based char n-grams if sacrebleu is missing."""
if not hyp or not ref: if not hyp or not ref:
return 0.0 return 0.0
try: try:
@@ -37,16 +62,7 @@ def _chrf(hyp: str, ref: str) -> float:
return float(CHRF(word_order=2).sentence_score(hyp, [ref]).score) return float(CHRF(word_order=2).sentence_score(hyp, [ref]).score)
except Exception: except Exception:
hyp_c, ref_c = list(hyp), list(ref) return _chrf_ngram(hyp, ref)
if not hyp_c:
return 0.0
ref_set = set(ref_c)
overlap = sum(1 for c in hyp_c if c in ref_set)
prec = overlap / len(hyp_c)
rec = overlap / max(len(ref_c), 1)
if prec + rec == 0:
return 0.0
return 100.0 * 2 * prec * rec / (prec + rec)
def translation_success( def translation_success(
+41
View File
@@ -0,0 +1,41 @@
"""Sanitize SwanLab env before import/init.
swanlab>=0.9 ``Settings.project`` is a nested model. A string
``SWANLAB_PROJECT`` (OpenBayes and older docs) makes pydantic raise
``error parsing value for field "project" from source
_QuoteAwareEnvSettingsSource``. Project name belongs in
``SWANLAB_PROJ_NAME`` / ``init(project=...)``.
"""
from __future__ import annotations
import os
def prepare_swanlab_env(default_project: str = "kda") -> str:
"""Drop nested ``SWANLAB_PROJECT``, keep a plain project name.
Must run before ``import swanlab`` / ``swanlab.login`` / ``init``.
"""
raw = os.environ.pop("SWANLAB_PROJECT", None)
name = os.environ.get("SWANLAB_PROJ_NAME") or raw or default_project
name = str(name).strip().strip("\"'")
if not name or name[0] in "{[":
name = default_project
os.environ.pop("SWANLAB_PROJECT", None)
os.environ["SWANLAB_PROJ_NAME"] = name
return name
def swanlab_run_id(run) -> str | None:
for attr in ("id", "run_id"):
val = getattr(run, attr, None)
if isinstance(val, str) and val:
return val
public = getattr(run, "public", None)
if public is not None:
for attr in ("cloud_run_id", "run_id", "id"):
val = getattr(public, attr, None)
if isinstance(val, str) and val:
return val
return None
+46
View File
@@ -45,6 +45,8 @@ questions:
text: "Block AttnRes 的两阶段算法为什么和 naive 逐层实现数值等价?" text: "Block AttnRes 的两阶段算法为什么和 naive 逐层实现数值等价?"
- id: Q9 - id: Q9
text: "深度残差接入 CausalLM 时怎样避免参数被重复注册?" text: "深度残差接入 CausalLM 时怎样避免参数被重复注册?"
- id: Q10
text: "为什么不用 stack([e(z) for e in experts]) 稠密计算全部专家?稀疏 permute-dispatch 如何让每个 token 只算 k 个专家?"
claims: claims:
- id: C1 - id: C1
@@ -83,6 +85,18 @@ claims:
text: "BorrowedSubLayer 用普通 tuple 持有 norm/fn,不注册为子模块,保证参数与 state_dict 键不重复" text: "BorrowedSubLayer 用普通 tuple 持有 norm/fn,不注册为子模块,保证参数与 state_dict 键不重复"
kind: methodological kind: methodological
status: supporting status: supporting
- id: C10
text: "LatentMoE 稀疏执行 = permute-dispatch + pad 到 [R, C, ℓ] + 三次 bmm + scatter-add,每个 token 只算 k 个专家(FLOPs R·C 而非 R·N)"
kind: methodological
status: core
- id: C11
text: "K3 路由 = s=σ(W_r x)、Top-k(s+b)、p_i = s_i/Σ_{j∈T}s_j;expert_bias 只进 TopK 选择、不进归一化权重"
kind: methodological
status: core
- id: C12
text: "负载均衡:Switch/GShard aux = n_r·Σ f_e·P_e 与 router z-loss = mean (logsumexp logits)^2,训练时加到 CE 上,只更新 router"
kind: methodological
status: core
symbols: symbols:
- {name: B, latex: "B", meaning: "batch size", kind: "shape parameter"} - {name: B, latex: "B", meaning: "batch size", kind: "shape parameter"}
@@ -115,6 +129,14 @@ symbols:
- {name: h_l, latex: "h_l", meaning: "深度注意力聚合出的层输入", domain: "[B, T, D]", kind: value} - {name: h_l, latex: "h_l", meaning: "深度注意力聚合出的层输入", domain: "[B, T, D]", kind: value}
- {name: b_j, latex: "b_j", meaning: "Block AttnRes 第 j 块的输出", domain: "[B, T, D]", kind: value} - {name: b_j, latex: "b_j", meaning: "Block AttnRes 第 j 块的输出", domain: "[B, T, D]", kind: value}
- {name: p, latex: "p", meaning: "块内 running partial", domain: "[B, T, D]", kind: value} - {name: p, latex: "p", meaning: "块内 running partial", domain: "[B, T, D]", kind: value}
- {name: s_moe, latex: "s", meaning: "router sigmoid 分数 σ(W_r x)", domain: "[B, T, n_r]", kind: value}
- {name: b, latex: "b", meaning: "expert bias(非持久 buffer,只进 TopK)", domain: "[n_r]", kind: value}
- {name: p_i, latex: "p_i", meaning: "sigmoid-L1 路由权重", domain: "[B, T, k]", kind: value}
- {name: C_moe, latex: "C_{\\mathrm{moe}}", meaning: "MoE 专家容量 = max 负载(pad 宽度)", kind: "shape parameter"}
- {name: f_e, latex: "f_e", meaning: "专家 e 被路由到的 token 占比", kind: value}
- {name: P_e, latex: "P_e", meaning: "专家 e 的平均 sigmoid 分数", kind: value}
- {name: L_aux, latex: "\\mathcal{L}_{aux}", meaning: "Switch/GShard 负载均衡损失", kind: value}
- {name: L_z, latex: "\\mathcal{L}_z", meaning: "router z-loss", kind: value}
terms: terms:
- {canonical: "KDA", aliases: ["Key-Decayed Attention", "键衰减注意力"]} - {canonical: "KDA", aliases: ["Key-Decayed Attention", "键衰减注意力"]}
@@ -129,6 +151,10 @@ terms:
- {canonical: "depth residual", aliases: ["DepthResidual", "深度维残差"]} - {canonical: "depth residual", aliases: ["DepthResidual", "深度维残差"]}
- {canonical: "online softmax", aliases: ["在线 softmax", "增量 softmax"]} - {canonical: "online softmax", aliases: ["在线 softmax", "增量 softmax"]}
- {canonical: "atomic layer", aliases: ["原子层", "atomic sublayer"]} - {canonical: "atomic layer", aliases: ["原子层", "atomic sublayer"]}
- {canonical: "permute-dispatch", aliases: ["置换-分发", "专家分发", "dispatch"]}
- {canonical: "grouped GEMM", aliases: ["padded bmm", "分组矩阵乘", "batched GEMM"]}
- {canonical: "load balancing loss", aliases: ["负载均衡损失", "aux loss", "Switch/GShard aux"]}
- {canonical: "z-loss", aliases: ["router z-loss", "logit 正则"]}
derivations: derivations:
- id: DER1 - id: DER1
@@ -162,6 +188,26 @@ derivations:
- {id: "3", from: "单源 partial p", to: "(m, n, d) = (s_p, p, 1),因为 e^{s_p - m} = 1", rule: definition} - {id: "3", from: "单源 partial p", to: "(m, n, d) = (s_p, p, 1),因为 e^{s_p - m} = 1", rule: definition}
- {id: "4", from: "(m_a,n_a,d_a), (m_b,n_b,d_b)", to: "m = \\max(m_a,m_b);\\ n = e^{m_a-m} n_a + e^{m_b-m} n_b;\\ d = e^{m_a-m} d_a + e^{m_b-m} d_b", rule: scale} - {id: "4", from: "(m_a,n_a,d_a), (m_b,n_b,d_b)", to: "m = \\max(m_a,m_b);\\ n = e^{m_a-m} n_a + e^{m_b-m} n_b;\\ d = e^{m_a-m} d_a + e^{m_b-m} d_b", rule: scale}
- {id: "5", from: "(m, n, d)", to: "h_l = n / d,与 forward_naive 逐位一致", rule: definition} - {id: "5", from: "(m, n, d)", to: "h_l = n / d,与 forward_naive 逐位一致", rule: definition}
- id: DER4
claim: C11
title: "K3 sigmoid-TopK 路由推导"
expand: true
figure: null
steps:
- {id: "1", from: "l = W_r x", to: "s = \\sigma(l) \\in [B,T,n_r]", rule: definition}
- {id: "2", from: "s + b", to: "T = \\mathrm{TopK}(s+b, k)", rule: selection}
- {id: "3", from: "T, s", to: "p_i = s_i / \\sum_{j \\in T} s_j", rule: normalize}
- {id: "4", from: "p, z", to: "u = \\sum_{i \\in T} p_i E_i^{rt}(z)", rule: definition}
- id: DER5
claim: C10
title: "稀疏 dispatch 执行流推导"
expand: true
figure: null
steps:
- {id: "1", from: "tok 重复 k 次 + eid 扁平化", to: "order = argsort(eid),同专家 token 连续", rule: permute}
- {id: "2", from: "counts = bincount(eid)", to: "C = max(counts);padded = index_put(zeros[R,C,ℓ], (eid, local_pos), z[tok])", rule: pad}
- {id: "3", from: "padded + 堆叠权重 [R,...]", to: "三次 bmm 得 [R,C,ff] → [R,C,ℓ](grouped GEMM)", rule: substitute}
- {id: "4", from: "out[eid,local_pos] 加权", to: "u = index_add(0, tok, p ⊙ out),FLOPs R·C 而非 R·N", rule: scatter-add}
figures: figures:
- id: F1 - id: F1
+1
View File
@@ -10,6 +10,7 @@
\usepackage{subcaption} \usepackage{subcaption}
\usepackage{float} \usepackage{float}
\usepackage{tikz} \usepackage{tikz}
\usetikzlibrary{positioning, arrows.meta, decorations.pathreplacing, calc}
\usepackage{hyperref} \usepackage{hyperref}
\usepackage{xcolor} \usepackage{xcolor}
\usepackage{multicol} \usepackage{multicol}
BIN
View File
Binary file not shown.
+15 -14
View File
@@ -75,7 +75,7 @@ class SiTU(nn.Module):
p_i = \frac{s_i}{\sum_{j\in T} s_j} p_i = \frac{s_i}{\sum_{j\in T} s_j}
\] \]
\item \textbf{Routed 专家}(在 latent 空间 $\ell$ 上,稀疏执行,见 \S7.4): \item \textbf{Routed 专家}(在 latent 空间 $\ell$ 上,稀疏执行,见 \ref{sec:sparse-dispatch} 节):
\[ \[
u = \sum_{i \in \mathrm{Top\text{-}k}} p_i \cdot E_i^{\mathrm{rt}}(z) u = \sum_{i \in \mathrm{Top\text{-}k}} p_i \cdot E_i^{\mathrm{rt}}(z)
\qquad \shape{B, T, \ell} \qquad \shape{B, T, \ell}
@@ -107,10 +107,10 @@ Shared 专家保持全宽 $d$,提供基础表达能力。
\begin{codemathtop}{layers/latent\_moe.py — \_route(K3 eq.13)} \begin{codemathtop}{layers/latent\_moe.py — \_route(K3 eq.13)}
\begin{lstlisting} \begin{lstlisting}
def _route(self, logits): # logits: [B, T, n_routed] def _route(self, logits): # logits: [B, T, n_routed]
scores = sigmoid(logits) # s = σ(W_r x) scores = sigmoid(logits) # s = sigma(W_r x)
ids = topk(scores + self.expert_bias, k).indices # T = TopK(s+b) ids = topk(scores + self.expert_bias, k).indices # T = TopK(s+b)
selected = scores.gather(-1, ids) selected = scores.gather(-1, ids)
probs = selected / selected.sum(-1).clamp_min(1e-9) # p_i = s_i / Σ_{j∈T} s_j probs = selected / selected.sum(-1).clamp_min(1e-9) # p_i = s_i / sum_{j in T} s_j
return ids, probs return ids, probs
\end{lstlisting} \end{lstlisting}
\end{codemathtop} \end{codemathtop}
@@ -138,7 +138,7 @@ def _routed_u(self, z, ids, probs): # z: [B,T,ell] ids/probs: [B,T,
tok = arange(N).unsqueeze(1).expand(N, k).reshape(-1) # 每个 token 重复 k 次 tok = arange(N).unsqueeze(1).expand(N, k).reshape(-1) # 每个 token 重复 k 次
eid, pw = ids.reshape(-1), probs.reshape(-1) eid, pw = ids.reshape(-1), probs.reshape(-1)
order = eid.argsort(stable=True) # 按专家 id 排序 → 同专家连续 order = eid.argsort(stable=True) # 按专家 id 排序 -> 同专家连续
tok, eid, pw = tok[order], eid[order], pw[order] tok, eid, pw = tok[order], eid[order], pw[order]
counts = bincount(eid, minlength=R) # 每个专家的 token 数 counts = bincount(eid, minlength=R) # 每个专家的 token 数
@@ -164,9 +164,9 @@ def _routed_u(self, z, ids, probs): # z: [B,T,ell] ids/probs: [B,T,
\end{lstlisting} \end{lstlisting}
\end{codemathtop} \end{codemathtop}
三步走:\textbf{① permute-dispatch}(按专家排序 + pad 到 $[R, C, \ell]$)→ 三步走:\textbf{(1) permute-dispatch}(按专家排序 + pad 到 $[R, C, \ell]$)$\to$
\textbf{② padded bmm}(专家参数堆成 batch 维,三次 batched GEMM 一次算完 $R$ 个专家)→ \textbf{(2) padded bmm}(专家参数堆成 batch 维,三次 batched GEMM 一次算完 $R$ 个专家)$\to$
\textbf{③ scatter-add}(\texttt{index\_add} 把加权输出按 \texttt{tok} 累加回 $u$)。 \textbf{(3) scatter-add}(\texttt{index\_add} 把加权输出按 \texttt{tok} 累加回 $u$)。
\begin{importantbox}{为什么不用 dense stack?} \begin{importantbox}{为什么不用 dense stack?}
朴素写法 \texttt{stack([e(z) for e in experts])} 会让每个专家都算全部 $B\cdot T$ 个 token, 朴素写法 \texttt{stack([e(z) for e in experts])} 会让每个专家都算全部 $B\cdot T$ 个 token,
@@ -197,7 +197,7 @@ Top-k 路由容易"塌缩"到少数专家(router 学出永远选某几个专
\] \]
\begin{center} \begin{center}
\begin{tabular}{ll} \begin{tabular}{lp{11.5cm}}
\toprule \toprule
项 & 作用 \\ 项 & 作用 \\
\midrule \midrule
@@ -215,18 +215,19 @@ def _balancing_losses(self, logits, ids):
counts = bincount(ids.reshape(-1), minlength=self.n_routed).float() counts = bincount(ids.reshape(-1), minlength=self.n_routed).float()
frac = counts / counts.sum().clamp_min(1.0) # f_e frac = counts / counts.sum().clamp_min(1.0) # f_e
prob_mean = scores.mean(dim=0) # P_e prob_mean = scores.mean(dim=0) # P_e
aux = self.n_routed * (frac * prob_mean).sum() # N Σ f_e P_e aux = self.n_routed * (frac * prob_mean).sum() # N * sum_e f_e * P_e
z_loss = logsumexp(flat, dim=-1).square().mean() # mean (logsumexp)^2 z_loss = logsumexp(flat, dim=-1).square().mean() # mean (logsumexp)^2
return self.aux_loss_coef * aux, self.z_loss_coef * z_loss return self.aux_loss_coef * aux, self.z_loss_coef * z_loss
\end{lstlisting} \end{lstlisting}
\end{codemathtop} \end{codemathtop}
两个损失只在 \texttt{self.training} 且系数非零时计算;系数默认 两个损失只在 \texttt{self.training} 且系数非零时计算;系数默认
$\alpha_{\mathrm{aux}} = 10^{-2}$、$\alpha_z = 10^{-3}$(\texttt{K3Config.moe\_aux\_loss\_coef} / $\alpha_{\mathrm{aux}} = 10^{-2}$、$\alpha_z = 10^{-3}$(即 \texttt{K3Config} 的
\texttt{moe\_z\_loss\_coef},可用 \texttt{--moe-aux-coef} / \texttt{--moe-z-coef} 覆盖)。 \texttt{moe\_aux\_loss\_coef} / \texttt{moe\_z\_loss\_coef},可用 \texttt{--moe-aux-coef}
train loop 里 \texttt{moe\_router\_losses(model)} 把所有 LatentMoE 层的损失求和, / \texttt{--moe-z-coef} 覆盖)。train loop 里 \texttt{moe\_router\_losses(model)}
\texttt{loss = task + aux + z\_loss} 一起反传。aux/z 只更新 router 参数, 把所有 LatentMoE 层的损失求和,\texttt{loss = task + aux + z\_loss} 一起反传。
不碰专家权重(\texttt{ids} 已 \texttt{.detach()})。 两个损失只依赖 router 输出 \texttt{logits} 与不可微的索引 \texttt{ids},
所以梯度只流回 router 的 $W_r$,不碰专家权重。
\subsection{形状总览} \subsection{形状总览}
+166
View File
@@ -153,6 +153,172 @@ class DecoderBlock(nn.Module):
\end{lstlisting} \end{lstlisting}
\end{codemathtop} \end{codemathtop}
\subsection{架构图}
\begin{figure}[H]
\centering
\begin{subfigure}[t]{0.44\textwidth}
\centering
\begin{tikzpicture}[>=Stealth, node distance=4mm,
blk/.style={draw, rounded corners=2pt, minimum width=32mm, minimum height=6mm,
align=center, font=\small},
io/.style={font=\small\itshape}]
\node[io] (in) {Input tokens};
\node[blk, fill=gray!8, below=5mm of in] (emb) {Embedding};
\node[blk, fill=blue!10, draw=blue!40, below=5mm of emb] (l0) {KDA + MoE};
\node[blk, fill=blue!10, draw=blue!40, below=2mm of l0] (l1) {KDA + MoE};
\node[blk, fill=blue!10, draw=blue!40, below=2mm of l1] (l2) {KDA + MoE};
\node[blk, fill=orange!12, draw=orange!50, below=2mm of l2] (l3) {MLA + MoE};
\node[below=1mm of l3, font=\normalsize] (dots) {$\vdots$};
\node[blk, fill=orange!12, draw=orange!50, below=1mm of dots] (lL) {MLA + MoE};
\draw[decorate, decoration={brace, amplitude=5pt, mirror}]
([xshift=2mm]l0.north east) -- ([xshift=2mm]lL.south east)
node[midway, right=6pt, font=\small] {$\times L$};
\node[blk, fill=gray!8, below=5mm of lL] (fnorm) {RMSNorm};
\node[blk, fill=gray!8, below=of fnorm] (head) {LM Head};
\node[io, below=of head] (out) {Logits};
\foreach \a/\b in {in/emb, emb/l0, l0/l1, l1/l2, l2/l3, l3/dots, dots/lL,
lL/fnorm, fnorm/head, head/out}
\draw[->] (\a) -- (\b);
\node[left=1mm of l0, font=\scriptsize, text=gray] {0};
\node[left=1mm of l1, font=\scriptsize, text=gray] {1};
\node[left=1mm of l2, font=\scriptsize, text=gray] {2};
\node[left=1mm of l3, font=\scriptsize, text=gray] {3};
\node[left=1mm of lL, font=\scriptsize, text=gray] {$L{-}1$};
\node[right=3mm of lL, font=\tiny, text=orange!60!black] {(强制)};
\end{tikzpicture}
\caption{整体模型}
\end{subfigure}
\hfill
\begin{subfigure}[t]{0.44\textwidth}
\centering
\begin{tikzpicture}[>=Stealth, node distance=5mm,
blk/.style={draw, rounded corners=2pt, minimum width=26mm, minimum height=6mm,
align=center, font=\small},
add/.style={circle, draw, thick, inner sep=0pt, minimum size=5.5mm,
font=\small\bfseries},
io/.style={font=\small\itshape}]
\node[io] (x) {$x$};
\node[blk, fill=gray!8, below=8mm of x] (n1) {RMSNorm};
\node[blk, fill=blue!10, draw=blue!40, below=of n1] (attn) {Attention};
\node[add, below=8mm of attn] (a1) {$+$};
\node[blk, fill=gray!8, below=8mm of a1] (n2) {RMSNorm};
\node[blk, fill=green!10, draw=green!40, below=of n2] (ffn) {FFN};
\node[add, below=8mm of ffn] (a2) {$+$};
\node[io, below=8mm of a2] (y) {$y$};
\foreach \a/\b in {x/n1, n1/attn, attn/a1, a1/n2, n2/ffn, ffn/a2, a2/y}
\draw[->] (\a) -- (\b);
\draw[->, gray!50, rounded corners=3pt]
(x.east) -- ++(14mm,0) |- (a1.east);
\draw[->, gray!50, rounded corners=3pt]
(a1.west) -- ++(-14mm,0) |- (a2.west);
\node[right=9mm of attn, font=\tiny, text=blue!60!black, align=left]
{KDA\\[-1pt]or MLA};
\node[left=9mm of ffn, font=\tiny, text=green!50!black, align=right]
{LatentMoE\\[-1pt]or SwiGLU};
\end{tikzpicture}
\caption{DecoderBlock}
\end{subfigure}
\caption{K3 混合架构。(a)~整体模型:每 4 层 1 次 MLA(层 3, 7, 11, \ldots),末层强制 MLA,
所有 FFN 均为 LatentMoE。(b)~DecoderBlock:Pre-Norm 残差,两个子块各含
RMSNorm $\to$ 子层 $\to$ 残差加。}
\label{fig:k3-overview}
\end{figure}
\begin{figure}[H]
\centering
%% ---------- (a) KDA ----------
\begin{subfigure}[t]{0.28\textwidth}
\centering
\begin{tikzpicture}[>=Stealth, node distance=5mm,
blk/.style={draw, rounded corners=2pt, minimum width=24mm, minimum height=6mm,
align=center, font=\footnotesize},
io/.style={font=\footnotesize\itshape}]
\node[io] (x) {$x$};
\node[blk, fill=blue!8, below=5mm of x] (proj)
{5 投影\\[-1pt]{\tiny $q, k, v, g, \beta$}};
\node[blk, fill=blue!12, draw=blue!40, below=of proj] (gate)
{Gate 激活};
\node[blk, fill=blue!20, draw=blue!50, below=of gate, minimum height=9mm]
(kda) {\texttt{chunk\_kda}\\[-1pt]{\tiny decay $+$ delta rule}};
\node[blk, fill=blue!8, below=of kda] (op) {$W_o$};
\node[io, below=5mm of op] (y) {$y$};
\foreach \a/\b in {x/proj, proj/gate, gate/kda, kda/op, op/y}
\draw[->] (\a) -- (\b);
\end{tikzpicture}
\caption{KDA Attention}
\end{subfigure}
\hfill
%% ---------- (b) Gated MLA ----------
\begin{subfigure}[t]{0.35\textwidth}
\centering
\begin{tikzpicture}[>=Stealth, node distance=5mm,
blk/.style={draw, rounded corners=2pt, minimum width=24mm, minimum height=6mm,
align=center, font=\footnotesize},
mul/.style={circle, draw, inner sep=0pt, minimum size=5mm, font=\tiny},
io/.style={font=\footnotesize\itshape}]
\node[io] (x) {$x$};
\node[blk, fill=orange!8, below=5mm of x] (lr)
{Q / KV 低秩压缩\\[-1pt]{\tiny $q_\downarrow\!\!\to\!\mathrm{norm}\!\to\!q_\uparrow$\;;\;
$c\!=\!\mathrm{norm}(W_\downarrow x)$}};
\node[blk, fill=orange!15, draw=orange!50, below=of lr] (abs)
{矩阵吸收 + 打分\\[-1pt]{\tiny $q_{\mathrm{abs}}\!=\!q\!\cdot\!W_{UK}$\;;\;
$\mathrm{score}\!=\!q_{\mathrm{abs}}\!\cdot\!c^T$}};
\node[blk, fill=orange!10, below=of abs] (sm)
{Causal Softmax};
\node[blk, fill=orange!12, draw=orange!40, below=of sm] (wuv)
{$\mathrm{attn}\!\cdot\!c \;\to\; W_{UV}^T$};
\node[mul, below=6mm of wuv] (m) {$\odot$};
\node[blk, fill=orange!6, right=4mm of m, minimum width=13mm, minimum height=5mm]
(g) {\tiny $\sigma(W_g x)$};
\draw[->] (g) -- (m);
\node[blk, fill=orange!8, below=6mm of m, minimum width=16mm] (op) {$W_o$};
\node[io, below=5mm of op] (y) {$y$};
\foreach \a/\b in {x/lr, lr/abs, abs/sm, sm/wuv, wuv/m, m/op, op/y}
\draw[->] (\a) -- (\b);
\end{tikzpicture}
\caption{Gated MLA}
\end{subfigure}
\hfill
%% ---------- (c) LatentMoE ----------
\begin{subfigure}[t]{0.30\textwidth}
\centering
\begin{tikzpicture}[>=Stealth, node distance=5mm,
blk/.style={draw, rounded corners=2pt, minimum width=16mm, minimum height=6mm,
align=center, font=\footnotesize},
add/.style={circle, draw, inner sep=0pt, minimum size=5mm,
font=\scriptsize\bfseries},
io/.style={font=\footnotesize\itshape}]
\node[io] (x) at (0,0) {$x$};
\node[blk, fill=green!10] (sh) at (-1.1,-1.3)
{Shared\\[-1pt]{\tiny SiTU, $d\!\to\!d$}};
\node[blk, fill=green!8, minimum width=20mm] (dr) at (1.1,-1.3)
{$W_\downarrow$ + Router\\[-1pt]{\tiny $\sigma$-TopK}};
\draw[->] (x) -- (sh);
\draw[->] (x) -- (dr);
\node[blk, fill=green!15, draw=green!40, minimum width=20mm] (re) at (1.1,-2.7)
{Routed 专家\\[-1pt]{\tiny SiTU, $\ell\!\to\!\ell$}};
\draw[->] (dr) -- (re);
\node[blk, fill=green!8, minimum width=20mm] (up) at (1.1,-4.0)
{RMSNorm $\to$ $W_\uparrow$};
\draw[->] (re) -- (up);
\node[add] (a) at (0,-5.2) {$+$};
\draw[->, rounded corners=3pt] (sh.south) -- ++(0,-3mm) -| (a);
\draw[->, rounded corners=3pt] (up.south) -- ++(0,-3mm) -| (a);
\node[io] (y) at (0,-6.0) {$y$};
\draw[->] (a) -- (y);
\end{tikzpicture}
\caption{LatentMoE}
\end{subfigure}
\caption{K3 三大组件。
(a)~KDA:5 路投影 $\to$ gate 激活 $\to$ \texttt{chunk\_kda}(decay $+$ delta rule)
$\to$ 输出投影。
(b)~Gated MLA:$q$ 吸收 $W_{UK}$ 后在 latent $c$ 上打分(NoPE);输出经 sigmoid 门控。
(c)~LatentMoE:shared 全宽 $d$ + routed 半宽 $\ell\!=\!d/2$;sigmoid-TopK 路由,
padded bmm 稀疏执行。}
\label{fig:k3-components}
\end{figure}
\subsection{本章小结} \subsection{本章小结}
K3 架构 = Hybrid Attention(3 KDA + 1 MLA,末层强制 MLA)+ LatentMoE。 K3 架构 = Hybrid Attention(3 KDA + 1 MLA,末层强制 MLA)+ LatentMoE。
+11 -4
View File
@@ -30,7 +30,7 @@ $n_r$ & routed 专家数 & 16 \\
$k$ & Top-$k$ & 2 \\ $k$ & Top-$k$ & 2 \\
$n_s$ & shared 专家数 & 2 \\ $n_s$ & shared 专家数 & 2 \\
$d_{\mathrm{ff}}$ & 专家中间维度 & 96 \\ $d_{\mathrm{ff}}$ & 专家中间维度 & 96 \\
$C$ & MoE 专家容量(pad 宽度) & 动态 \\ $C_{\mathrm{moe}}$ & MoE 专家容量(pad 宽度) & 动态 \\
$\alpha_{\mathrm{aux}}$ & Switch/GShard aux 系数 & $10^{-2}$ \\ $\alpha_{\mathrm{aux}}$ & Switch/GShard aux 系数 & $10^{-2}$ \\
$\alpha_z$ & router z-loss 系数 & $10^{-3}$ \\ $\alpha_z$ & router z-loss 系数 & $10^{-3}$ \\
$N$ & AttnRes 原子层数 ($= 2L$) & 8 \\ $N$ & AttnRes 原子层数 ($= 2L$) & 8 \\
@@ -113,8 +113,8 @@ $s$ & \shape{B, T, n_r} & sigmoid 分数 $\sigma(\mathrm{logits})$ \\
$b$ & \shape{n_r} & expert bias(非持久,只进 TopK) \\ $b$ & \shape{n_r} & expert bias(非持久,只进 TopK) \\
ids & \shape{B, T, k} & Top-$k$ 专家索引 \\ ids & \shape{B, T, k} & Top-$k$ 专家索引 \\
$p_i$ & \shape{B, T, k} & sigmoid-L1 权重 $s_i/\sum_{j\in T}s_j$ \\ $p_i$ & \shape{B, T, k} & sigmoid-L1 权重 $s_i/\sum_{j\in T}s_j$ \\
padded & \shape{n_r, C, \ell} & dispatch 后 pad 到容量 $C$ \\ padded & \shape{n_r, C_{\mathrm{moe}}, \ell} & dispatch 后 pad 到容量 $C_{\mathrm{moe}}$ \\
$C$ & 标量 & 最大专家负载(pad 宽度) \\ $C_{\mathrm{moe}}$ & 标量 & 最大专家负载(pad 宽度) \\
$u$ & \shape{B, T, \ell} & routed 加权输出 \\ $u$ & \shape{B, T, \ell} & routed 加权输出 \\
$s_{\mathrm{sh}}$ & \shape{B, T, D} & shared 专家求和 \\ $s_{\mathrm{sh}}$ & \shape{B, T, D} & shared 专家求和 \\
$y$ & \shape{B, T, D} & $s_{\mathrm{sh}} + W_\uparrow \mathrm{RMSNorm}(u)$ \\ $y$ & \shape{B, T, D} & $s_{\mathrm{sh}} + W_\uparrow \mathrm{RMSNorm}(u)$ \\
@@ -164,6 +164,11 @@ MLA 解压 & \texttt{'bhtj,hvj->bhtv'} & $\tilde{o}$ \shape{B,H,T,d_v} \\
AttnRes 深度打分 & \texttt{'d,nbtd->nbt'} & $s_{l,i}$ \shape{n,B,T} \\ AttnRes 深度打分 & \texttt{'d,nbtd->nbt'} & $s_{l,i}$ \shape{n,B,T} \\
AttnRes 深度加权和 & \texttt{'nbt,nbtd->btd'} & $h_l$ \shape{B,T,D} \\ AttnRes 深度加权和 & \texttt{'nbt,nbtd->btd'} & $h_l$ \shape{B,T,D} \\
AttnRes 批量打分(inter) & \texttt{'qd,nbtd->qnbt'} & logits \shape{S,n,B,T} \\ AttnRes 批量打分(inter) & \texttt{'qd,nbtd->qnbt'} & logits \shape{S,n,B,T} \\
MoE dispatch pad & \texttt{index\_put} & padded \shape{R,C_{\mathrm{moe}},\ell} \\
MoE gate 投影(grouped) & \texttt{bmm(padded, w\_g.T)} & $wg$ \shape{R,C_{\mathrm{moe}},ff} \\
MoE up 投影(grouped) & \texttt{bmm(padded, w\_u.T)} & $wu$ \shape{R,C_{\mathrm{moe}},ff} \\
MoE 输出投影(grouped) & \texttt{bmm(g$\odot$h, w\_o.T)} & out \shape{R,C_{\mathrm{moe}},\ell} \\
MoE scatter-add & \texttt{index\_add(0, tok, ...)} & $u$ \shape{N,\ell} \\
\bottomrule \bottomrule
\end{tabular} \end{tabular}
\end{center} \end{center}
@@ -177,7 +182,9 @@ AttnRes 批量打分(inter) & \texttt{'qd,nbtd->qnbt'} & logits \shape{S,n,B
\item \textbf{分块} = chunk 内下三角解 + chunk 间状态递推,等价于 naive recurrent \item \textbf{分块} = chunk 内下三角解 + chunk 间状态递推,等价于 naive recurrent
\item \textbf{GVA} = $H_V = G \cdot H$,forward repeat\_interleave / backward view+sum \item \textbf{GVA} = $H_V = G \cdot H$,forward repeat\_interleave / backward view+sum
\item \textbf{MLA} = 低秩 latent + 矩阵吸收,KV cache 从 $2Hd$ 降到 $r$ \item \textbf{MLA} = 低秩 latent + 矩阵吸收,KV cache 从 $2Hd$ 降到 $r$
\item \textbf{LatentMoE} = shared 全宽 + routed 半宽 latent + SiTU-GLU 防溢出 \item \textbf{LatentMoE} = shared 全宽 + routed 半宽 latent + SiTU-GLU 防溢出;
K3 sigmoid-TopK 路由 + 稀疏 permute-dispatch(每 token 只算 $k$ 个专家)+
Switch/GShard aux \& z-loss 防塌缩
\item \textbf{K3 Hybrid} = 3 KDA + 1 MLA,KDA 提供位置感知 \item \textbf{K3 Hybrid} = 3 KDA + 1 MLA,KDA 提供位置感知
\item \textbf{AttnRes} = 深度维 softmax 残差,Block 版把源数压到 $O(N/S)$, \item \textbf{AttnRes} = 深度维 softmax 残差,Block 版把源数压到 $O(N/S)$,
两阶段 = inter 批量 + intra online-softmax 合并 两阶段 = inter 批量 + intra online-softmax 合并
+10
View File
@@ -28,6 +28,16 @@ def test_good_zh2en_passes():
assert translation_success(src, hyp, ref, target_lang="en") is True assert translation_success(src, hyp, ref, target_lang="en") is True
def test_english_wiki_garbage_does_not_pass_zh2en():
from kda.training.success import _chrf
src = "今天天气很好。"
hyp = "The first one's the time."
ref = "The weather is very nice today."
assert _chrf(hyp, ref) < 40.0
assert translation_success(src, hyp, ref, target_lang="en") is False
def test_container_help_exits_2(): def test_container_help_exits_2():
import importlib.util import importlib.util
from pathlib import Path from pathlib import Path
+1
View File
@@ -78,6 +78,7 @@ def test_preset_0_5b_schedule():
assert cfg.hidden_size == 768 assert cfg.hidden_size == 768
assert cfg.num_heads * cfg.head_dim == cfg.hidden_size assert cfg.num_heads * cfg.head_dim == cfg.hidden_size
assert cfg.num_hidden_layers == 24 assert cfg.num_hidden_layers == 24
assert cfg.vocab_size == 64000
assert cfg.tie_word_embeddings assert cfg.tie_word_embeddings
assert cfg.chunk_size == 64 assert cfg.chunk_size == 64
assert cfg.gradient_checkpointing is True assert cfg.gradient_checkpointing is True
+13
View File
@@ -0,0 +1,13 @@
from train_sft import _is_better_eval, _sibling
def test_sibling_last_best():
assert _sibling("ckpts/k3_sft.pt", "_last") == "ckpts/k3_sft_last.pt"
assert _sibling("ckpts/k3_sft.pt", "_best") == "ckpts/k3_sft_best.pt"
def test_best_prefers_success_then_chrf():
assert _is_better_eval(1.0, 70.0, 0.95, 90.0)
assert not _is_better_eval(0.95, 99.0, 1.0, 70.0)
assert _is_better_eval(1.0, 91.0, 1.0, 81.0)
assert not _is_better_eval(1.0, 70.0, 1.0, 81.0)
+39
View File
@@ -0,0 +1,39 @@
import os
from kda.training.swanlab_env import prepare_swanlab_env, swanlab_run_id
def test_prepare_drops_string_project(monkeypatch):
monkeypatch.setenv("SWANLAB_PROJECT", "kda")
monkeypatch.delenv("SWANLAB_PROJ_NAME", raising=False)
assert prepare_swanlab_env() == "kda"
assert "SWANLAB_PROJECT" not in os.environ
assert os.environ["SWANLAB_PROJ_NAME"] == "kda"
def test_prepare_strips_quotes(monkeypatch):
monkeypatch.setenv("SWANLAB_PROJECT", '"kda"')
monkeypatch.delenv("SWANLAB_PROJ_NAME", raising=False)
assert prepare_swanlab_env() == "kda"
assert "SWANLAB_PROJECT" not in os.environ
def test_prepare_prefers_proj_name(monkeypatch):
monkeypatch.setenv("SWANLAB_PROJECT", "ignored")
monkeypatch.setenv("SWANLAB_PROJ_NAME", "mine")
assert prepare_swanlab_env() == "mine"
assert os.environ["SWANLAB_PROJ_NAME"] == "mine"
def test_prepare_rejects_json_blob(monkeypatch):
monkeypatch.setenv("SWANLAB_PROJECT", '{"name": "x"}')
monkeypatch.delenv("SWANLAB_PROJ_NAME", raising=False)
assert prepare_swanlab_env() == "kda"
def test_swanlab_run_id():
class _Run:
id = "ilgne5ro"
assert swanlab_run_id(_Run()) == "ilgne5ro"
assert swanlab_run_id(object()) is None
+81
View File
@@ -0,0 +1,81 @@
"""Resume must not clobber attnres or skip a chunk at budget exit."""
from argparse import Namespace
from dataclasses import asdict
import pytest
import torch
from kda.models.causal_lm import CausalLM
from kda.models.k3_config import K3Config
from kda.training.toy import load_ckpt
from train_k3 import _apply_cli_overrides
def _tiny_block():
return K3Config(
hidden_size=32,
num_hidden_layers=4,
num_heads=4,
head_dim=8,
chunk_size=4,
vocab_size=64,
moe_latent_size=16,
moe_d_ff=16,
n_routed=4,
top_k=2,
n_shared=1,
kv_lora_rank=8,
q_lora_rank=16,
qk_nope_head_dim=8,
v_head_dim=8,
attnres="block",
attnres_block_size=2,
)
def _cli(**kwargs):
base = dict(
attnres=None,
attnres_block_size=None,
grad_checkpoint=None,
moe_aux_coef=None,
moe_z_coef=None,
)
base.update(kwargs)
return Namespace(**base)
def test_resume_without_attnres_flag_keeps_block():
cfg = _tiny_block()
_apply_cli_overrides(cfg, _cli())
assert cfg.attnres == "block"
assert cfg.attnres_block_size == 2
def test_explicit_attnres_overrides_resume():
cfg = _tiny_block()
_apply_cli_overrides(cfg, _cli(attnres="full", attnres_block_size=1))
assert cfg.attnres == "full"
assert cfg.attnres_block_size == 1
def test_block_state_with_off_config_cannot_load(tmp_path):
cfg = _tiny_block()
model = CausalLM(cfg)
payload = {"config": asdict(cfg), "model_state": model.state_dict()}
payload["config"]["attnres"] = "off"
payload["config"]["attnres_block_size"] = None
path = str(tmp_path / "polluted.pt")
torch.save(payload, path)
with pytest.raises(RuntimeError, match="Unexpected key"):
load_ckpt(path)
def test_budget_break_does_not_skip_yielded_chunk():
next_chunk = 10
for chunk_index in (10, 11, 12):
tokens = 100
if tokens >= 100:
break
next_chunk = chunk_index + 1
assert next_chunk == 10
+48 -44
View File
@@ -22,6 +22,7 @@ from kda.models.causal_lm import CausalLM
from kda.models.k3_config import K3Config from kda.models.k3_config import K3Config
from kda.training.data import iter_indexed, load_pretrain_chunks, load_tokenizer from kda.training.data import iter_indexed, load_pretrain_chunks, load_tokenizer
from kda.training.schedule import lr_scale, tokens_per_micro, total_opt_steps from kda.training.schedule import lr_scale, tokens_per_micro, total_opt_steps
from kda.training.swanlab_env import prepare_swanlab_env, swanlab_run_id
from kda.training.toy import load_ckpt from kda.training.toy import load_ckpt
_TOY_TRAIN = { _TOY_TRAIN = {
@@ -40,7 +41,7 @@ _TOY_TRAIN = {
"gen_every": 200, "gen_every": 200,
} }
_B500M_TRAIN = { _B500M_TRAIN = {
"tokenizer": "Qwen/Qwen3-8B", "tokenizer": "01-ai/Yi-6B",
"out": "ckpts/k3_0.5b.pt", "out": "ckpts/k3_0.5b.pt",
"limit": 20000, "limit": 20000,
"batch": 2, "batch": 2,
@@ -66,25 +67,14 @@ def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None:
group["lr"] = lr group["lr"] = lr
def _swanlab_run_id(run) -> str | None: def _init_swanlab(
for attr in ("id", "run_id"): cfg: K3Config, args: argparse.Namespace, resume_id: str | None = None
val = getattr(run, attr, None) ):
if isinstance(val, str) and val:
return val
public = getattr(run, "public", None)
if public is not None:
for attr in ("cloud_run_id", "run_id", "id"):
val = getattr(public, attr, None)
if isinstance(val, str) and val:
return val
return None
def _init_swanlab(cfg: K3Config, args: argparse.Namespace, resume_id: str | None = None):
"""Cloud monitor if SWANLAB_API_KEY is set; otherwise no-op.""" """Cloud monitor if SWANLAB_API_KEY is set; otherwise no-op."""
key = os.environ.get("SWANLAB_API_KEY") key = os.environ.get("SWANLAB_API_KEY")
if not key: if not key:
return None return None
project = prepare_swanlab_env()
try: try:
import swanlab import swanlab
except ImportError: except ImportError:
@@ -92,8 +82,6 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace, resume_id: str | None
return None return None
try: try:
swanlab.login(api_key=key, save=False) swanlab.login(api_key=key, save=False)
# swanlab 0.9 Settings.project is nested; a string SWANLAB_PROJECT env crashes init.
project = os.environ.pop("SWANLAB_PROJECT", None) or "kda"
run_id = resume_id or os.environ.get("SWANLAB_RUN_ID") run_id = resume_id or os.environ.get("SWANLAB_RUN_ID")
init_kw = dict( init_kw = dict(
project=project, project=project,
@@ -125,7 +113,7 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace, resume_id: str | None
init_kw["resume"] = True init_kw["resume"] = True
print(f"swanlab resume id={run_id}") print(f"swanlab resume id={run_id}")
run = swanlab.init(**init_kw) run = swanlab.init(**init_kw)
got = _swanlab_run_id(run) got = swanlab_run_id(run)
if got: if got:
args.swanlab_id = got args.swanlab_id = got
print(f"swanlab run id {got}") print(f"swanlab run id {got}")
@@ -196,6 +184,20 @@ def _apply_moe_coefs(model, cfg: K3Config) -> None:
module.z_loss_coef = cfg.moe_z_loss_coef module.z_loss_coef = cfg.moe_z_loss_coef
def _apply_cli_overrides(cfg: K3Config, args: argparse.Namespace) -> None:
"""Copy only flags the user actually passed. CLI defaults must not clobber a resume."""
if args.attnres is not None:
cfg.attnres = args.attnres
if args.attnres_block_size is not None:
cfg.attnres_block_size = args.attnres_block_size
if args.grad_checkpoint is not None:
cfg.gradient_checkpointing = args.grad_checkpoint
if args.moe_aux_coef is not None:
cfg.moe_aux_loss_coef = args.moe_aux_coef
if args.moe_z_coef is not None:
cfg.moe_z_loss_coef = args.moe_z_coef
def _moe_log(model) -> dict: def _moe_log(model) -> dict:
frac = moe_route_frac(model) frac = moe_route_frac(model)
if frac is None: if frac is None:
@@ -279,9 +281,10 @@ def main() -> None:
p.add_argument("--device", default="auto") p.add_argument("--device", default="auto")
p.add_argument( p.add_argument(
"--attnres", "--attnres",
default="off", default=None,
choices=["off", "full", "block"], choices=["off", "full", "block"],
help="depth mixer: off=standard residual, block=K3 AttnRes, full=per-layer AttnRes", help="depth mixer: off=standard residual (preset default), block=K3 AttnRes, "
"full=per-layer AttnRes. Omit on --resume to keep the checkpoint value",
) )
p.add_argument( p.add_argument(
"--attnres-block-size", "--attnres-block-size",
@@ -341,14 +344,7 @@ def main() -> None:
tok = load_tokenizer(args.tokenizer) tok = load_tokenizer(args.tokenizer)
cfg = K3Config.preset(args.preset) cfg = K3Config.preset(args.preset)
cfg.vocab_size = tok.vocab_size cfg.vocab_size = tok.vocab_size
cfg.attnres = args.attnres _apply_cli_overrides(cfg, args)
cfg.attnres_block_size = args.attnres_block_size
if args.grad_checkpoint is not None:
cfg.gradient_checkpointing = args.grad_checkpoint
if args.moe_aux_coef is not None:
cfg.moe_aux_loss_coef = args.moe_aux_coef
if args.moe_z_coef is not None:
cfg.moe_z_loss_coef = args.moe_z_coef
langs = [part.strip() for part in args.langs.split(",") if part.strip()] langs = [part.strip() for part in args.langs.split(",") if part.strip()]
tpm = tokens_per_micro(args.batch, args.seq_len) tpm = tokens_per_micro(args.batch, args.seq_len)
@@ -376,19 +372,16 @@ def main() -> None:
f"{type(loaded_cfg).__name__}" f"{type(loaded_cfg).__name__}"
) )
cfg = loaded_cfg cfg = loaded_cfg
cfg.attnres = args.attnres _apply_cli_overrides(cfg, args)
cfg.attnres_block_size = args.attnres_block_size
if args.grad_checkpoint is not None:
cfg.gradient_checkpointing = args.grad_checkpoint
if args.moe_aux_coef is not None:
cfg.moe_aux_loss_coef = args.moe_aux_coef
if args.moe_z_coef is not None:
cfg.moe_z_loss_coef = args.moe_z_coef
model.gradient_checkpointing = cfg.gradient_checkpointing model.gradient_checkpointing = cfg.gradient_checkpointing
model.to(device) model.to(device)
payload = torch.load(args.resume, map_location="cpu", weights_only=False) payload = torch.load(args.resume, map_location="cpu", weights_only=False)
if payload.get("tokenizer") and payload["tokenizer"] != args.tokenizer: if payload.get("tokenizer") and payload["tokenizer"] != args.tokenizer:
print(f"warning: ckpt tokenizer {payload['tokenizer']} != {args.tokenizer}") raise SystemExit(
f"tokenizer mismatch: ckpt {payload['tokenizer']!r} vs "
f"CLI {args.tokenizer!r}; embeddings are not interchangeable "
f"(do not resume a Qwen ckpt with Yi)"
)
micro_step = int(payload.get("micro_step", 0)) micro_step = int(payload.get("micro_step", 0))
opt_step = int(payload.get("opt_step", 0)) opt_step = int(payload.get("opt_step", 0))
tokens = int(payload.get("tokens", 0)) tokens = int(payload.get("tokens", 0))
@@ -481,6 +474,10 @@ def main() -> None:
model.train() model.train()
t0 = time.perf_counter() t0 = time.perf_counter()
tokens_at_t0 = tokens tokens_at_t0 = tokens
# Index of the next untrained chunk. Mid-loop saves use last_trained+1.
# The final save must NOT +1 again: the loop may break on a yielded chunk
# that was never trained (budget check is at the top).
next_chunk = chunk_index
for chunk_index, x, y in iter_indexed(train_chunks, start=chunk_index): for chunk_index, x, y in iter_indexed(train_chunks, start=chunk_index):
if args.max_tokens is not None: if args.max_tokens is not None:
if tokens >= args.max_tokens: if tokens >= args.max_tokens:
@@ -509,6 +506,7 @@ def main() -> None:
lr_now = optim.param_groups[0]["lr"] lr_now = optim.param_groups[0]["lr"]
if raw_loss < best_train: if raw_loss < best_train:
best_train = raw_loss best_train = raw_loss
next_chunk = chunk_index + 1
metrics = { metrics = {
"train/loss": raw_loss, "train/loss": raw_loss,
@@ -523,9 +521,9 @@ def main() -> None:
if elapsed > 0: if elapsed > 0:
metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed
ended = ( ended = (args.max_tokens is not None and tokens >= args.max_tokens) or (
args.max_tokens is not None and tokens >= args.max_tokens args.max_tokens is None and micro_step >= args.steps
) or (args.max_tokens is None and micro_step >= args.steps) )
log_now = micro_step % args.log_every == 0 or micro_step == 1 or ended log_now = micro_step % args.log_every == 0 or micro_step == 1 or ended
eval_now = micro_step % args.eval_every == 0 or micro_step == 1 or ended eval_now = micro_step % args.eval_every == 0 or micro_step == 1 or ended
ckpt_now = micro_step % args.ckpt_every == 0 or ended ckpt_now = micro_step % args.ckpt_every == 0 or ended
@@ -554,7 +552,7 @@ def main() -> None:
micro_step=micro_step, micro_step=micro_step,
opt_step=opt_step, opt_step=opt_step,
tokens=tokens, tokens=tokens,
chunk_index=chunk_index + 1, chunk_index=next_chunk,
best_heldout=best_heldout, best_heldout=best_heldout,
) )
_save(_sibling(args.out, "_best"), payload) _save(_sibling(args.out, "_best"), payload)
@@ -582,7 +580,7 @@ def main() -> None:
micro_step=micro_step, micro_step=micro_step,
opt_step=opt_step, opt_step=opt_step,
tokens=tokens, tokens=tokens,
chunk_index=chunk_index + 1, chunk_index=next_chunk,
best_heldout=best_heldout, best_heldout=best_heldout,
) )
_save(_sibling(args.out, "_last"), payload) _save(_sibling(args.out, "_last"), payload)
@@ -598,7 +596,7 @@ def main() -> None:
micro_step=micro_step, micro_step=micro_step,
opt_step=opt_step, opt_step=opt_step,
tokens=tokens, tokens=tokens,
chunk_index=chunk_index + 1, chunk_index=next_chunk,
best_heldout=best_heldout, best_heldout=best_heldout,
) )
_save(args.out, payload) _save(args.out, payload)
@@ -608,6 +606,12 @@ def main() -> None:
) )
if tracker is not None: if tracker is not None:
tracker.finish() tracker.finish()
if device == "cuda":
try:
torch.cuda.synchronize()
torch.cuda.empty_cache()
except Exception:
pass
if __name__ == "__main__": if __name__ == "__main__":
+200 -32
View File
@@ -4,6 +4,7 @@
uv run python train_sft.py --ckpt ckpts/k3_wiki.pt --data data/sft/toy.jsonl uv run python train_sft.py --ckpt ckpts/k3_wiki.pt --data data/sft/toy.jsonl
uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data opus-100 \\ uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data opus-100 \\
--limit 100000 --seq-len 512 --batch 4 --lr 5e-5 --epochs 2 --limit 100000 --seq-len 512 --batch 4 --lr 5e-5 --epochs 2
uv run python train_sft.py --resume --out ckpts/k3_sft.pt
""" """
from __future__ import annotations from __future__ import annotations
@@ -23,6 +24,7 @@ from kda.training.data import (
) )
from kda.training.eval_mt import evaluate_pairs from kda.training.eval_mt import evaluate_pairs
from kda.training.schedule import lr_scale, total_opt_steps from kda.training.schedule import lr_scale, total_opt_steps
from kda.training.swanlab_env import prepare_swanlab_env, swanlab_run_id
from kda.training.toy import load_ckpt from kda.training.toy import load_ckpt
@@ -31,20 +33,43 @@ def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None:
group["lr"] = lr group["lr"] = lr
def _init_swanlab(args: argparse.Namespace): def _sibling(path: str, suffix: str) -> str:
root, ext = os.path.splitext(path)
return f"{root}{suffix}{ext}"
def _save(path: str, payload: dict) -> None:
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
tmp = path + ".tmp"
torch.save(payload, tmp)
os.replace(tmp, path)
def _is_better_eval(
success: float, chrf: float, best_success: float, best_chrf: float
) -> bool:
if success > best_success:
return True
if success == best_success and chrf > best_chrf:
return True
return False
def _init_swanlab(args: argparse.Namespace, resume_id: str | None = None):
key = os.environ.get("SWANLAB_API_KEY") key = os.environ.get("SWANLAB_API_KEY")
if not key: if not key:
return None return None
project = prepare_swanlab_env()
try: try:
import swanlab import swanlab
except ImportError: except ImportError:
print("SWANLAB_API_KEY set but swanlab is not installed")
return None return None
try: try:
swanlab.login(api_key=key, save=False) swanlab.login(api_key=key, save=False)
project = os.environ.pop("SWANLAB_PROJECT", None) or "kda" init_kw = dict(
return swanlab.init(
project=project, project=project,
name=f"sft-{os.path.basename(args.ckpt)}", name=f"sft-{os.path.basename(args.ckpt or args.out)}",
config={ config={
"ckpt": args.ckpt, "ckpt": args.ckpt,
"data": args.data, "data": args.data,
@@ -52,8 +77,20 @@ def _init_swanlab(args: argparse.Namespace):
"batch": args.batch, "batch": args.batch,
"seq_len": args.seq_len, "seq_len": args.seq_len,
"epochs": args.epochs, "epochs": args.epochs,
"grad_acc": args.grad_acc,
"limit": args.limit,
}, },
) )
if resume_id:
init_kw["id"] = resume_id
init_kw["resume"] = True
print(f"swanlab resume id={resume_id}")
run = swanlab.init(**init_kw)
got = swanlab_run_id(run)
if got:
args.swanlab_id = got
print(f"swanlab run id {got}")
return run
except Exception as exc: except Exception as exc:
print(f"swanlab init failed ({exc}); continuing without cloud monitor") print(f"swanlab init failed ({exc}); continuing without cloud monitor")
return None return None
@@ -69,9 +106,50 @@ def _read_lines(path: str) -> list[str]:
] ]
def _payload(
cfg,
model,
optim: torch.optim.Optimizer,
args: argparse.Namespace,
*,
tok_src: str,
step: int,
opt_step: int,
row_index: int,
best_success: float,
best_chrf: float,
best_train: float,
):
return {
"config": asdict(cfg),
"model_state": model.state_dict(),
"optimizer_state": optim.state_dict(),
"tokenizer": tok_src,
"sft_data": args.data,
"pretrained_ckpt": args.ckpt,
"sft_step": step,
"opt_step": opt_step,
"row_index": row_index,
"best_success": best_success,
"best_chrf": best_chrf,
"best_train": best_train,
"batch": args.batch,
"seq_len": args.seq_len,
"grad_acc": args.grad_acc,
"swanlab_id": getattr(args, "swanlab_id", None),
}
def main() -> None: def main() -> None:
p = argparse.ArgumentParser(description=__doc__) p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--ckpt", required=True) p.add_argument("--ckpt", default=None, help="pretrained (or SFT) checkpoint to start from")
p.add_argument(
"--resume",
nargs="?",
const="__last__",
default=None,
help="resume SFT; default path is <out>_last",
)
p.add_argument( p.add_argument(
"--data", "--data",
default="opus-100", default="opus-100",
@@ -98,12 +176,26 @@ def main() -> None:
p.add_argument("--max-steps", type=int, default=None) p.add_argument("--max-steps", type=int, default=None)
p.add_argument("--grad-acc", type=int, default=1) p.add_argument("--grad-acc", type=int, default=1)
p.add_argument("--eval-every", type=int, default=50) p.add_argument("--eval-every", type=int, default=50)
p.add_argument(
"--ckpt-every",
type=int,
default=500,
help="write _last this many steps (0 = only interrupt + end)",
)
p.add_argument("--src", default=None, help="frozen eval src (not used as train)") p.add_argument("--src", default=None, help="frozen eval src (not used as train)")
p.add_argument("--ref", default=None) p.add_argument("--ref", default=None)
p.add_argument("--target-lang", default="en", choices=["en", "zh"]) p.add_argument("--target-lang", default="en", choices=["en", "zh"])
p.add_argument("--device", default="auto") p.add_argument("--device", default="auto")
args = p.parse_args() args = p.parse_args()
resume_path = args.resume
if resume_path == "__last__":
resume_path = _sibling(args.out, "_last")
if resume_path is None and not args.ckpt:
raise SystemExit("need --ckpt or --resume")
if resume_path is not None and not os.path.isfile(resume_path):
raise SystemExit(f"resume checkpoint not found: {resume_path}")
device = args.device device = args.device
if device == "auto": if device == "auto":
device = "cuda" if torch.cuda.is_available() else "cpu" device = "cuda" if torch.cuda.is_available() else "cpu"
@@ -111,10 +203,13 @@ def main() -> None:
if device == "cuda" and not use_bf16: if device == "cuda" and not use_bf16:
raise SystemExit("KDA training needs bf16") raise SystemExit("KDA training needs bf16")
model, cfg = load_ckpt(args.ckpt) start_path = resume_path or args.ckpt
model, cfg = load_ckpt(start_path)
model.to(device) model.to(device)
payload = torch.load(args.ckpt, map_location="cpu", weights_only=False) loaded = torch.load(start_path, map_location="cpu", weights_only=False)
tok_src = args.tokenizer or payload.get("tokenizer") if args.ckpt is None:
args.ckpt = loaded.get("pretrained_ckpt")
tok_src = args.tokenizer or loaded.get("tokenizer")
if not tok_src: if not tok_src:
raise SystemExit("need --tokenizer or a tokenizer field in the checkpoint") raise SystemExit("need --tokenizer or a tokenizer field in the checkpoint")
tok = load_tokenizer(tok_src) tok = load_tokenizer(tok_src)
@@ -125,7 +220,6 @@ def main() -> None:
) )
if not rows: if not rows:
raise SystemExit(f"no SFT rows from {args.data}") raise SystemExit(f"no SFT rows from {args.data}")
print(f"SFT {len(rows)} rows from {args.data}; model {cfg.__class__.__name__}")
steps_per_epoch = max((len(rows) + args.batch - 1) // args.batch, 1) steps_per_epoch = max((len(rows) + args.batch - 1) // args.batch, 1)
max_micro = args.max_steps max_micro = args.max_steps
@@ -138,14 +232,61 @@ def main() -> None:
seq_len=args.seq_len, seq_len=args.seq_len,
grad_acc=args.grad_acc, grad_acc=args.grad_acc,
) )
print(
f"SFT {len(rows)} rows from {args.data}; model {cfg.__class__.__name__} "
f"{max_micro} steps ({args.epochs} epoch, batch {args.batch}) "
f"eval/{args.eval_every} ckpt/{args.ckpt_every}"
)
optim = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01) optim = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01)
tracker = _init_swanlab(args)
model.train()
step = 0 step = 0
opt_step = 0 opt_step = 0
best = float("inf") row_start = 0
best_train = float("inf")
best_success = -1.0
best_chrf = -1.0
if resume_path is not None:
if loaded.get("optimizer_state"):
optim.load_state_dict(loaded["optimizer_state"])
step = int(loaded.get("sft_step", 0))
opt_step = int(loaded.get("opt_step", 0))
row_start = int(loaded.get("row_index", 0))
best_train = float(loaded.get("best_train", best_train))
best_success = float(loaded.get("best_success", best_success))
best_chrf = float(loaded.get("best_chrf", best_chrf))
print(
f"resume {resume_path} step {step} opt {opt_step} "
f"row {row_start} best success {best_success:.2f} chrf {best_chrf:.2f}"
)
for _, x, y in iter_sft_batches(rows, tok, args.batch, args.seq_len): tracker = _init_swanlab(
args,
resume_id=loaded.get("swanlab_id") if resume_path is not None else None,
)
last_path = _sibling(args.out, "_last")
best_path = _sibling(args.out, "_best")
model.train()
next_row = row_start
def dump() -> dict:
return _payload(
cfg,
model,
optim,
args,
tok_src=tok_src,
step=step,
opt_step=opt_step,
row_index=next_row,
best_success=best_success,
best_chrf=best_chrf,
best_train=best_train,
)
try:
for row_index, x, y in iter_sft_batches(
rows, tok, args.batch, args.seq_len, start=row_start
):
if step >= max_micro: if step >= max_micro:
break break
x, y = x.to(device), y.to(device) x, y = x.to(device), y.to(device)
@@ -161,8 +302,9 @@ def main() -> None:
optim.zero_grad(set_to_none=True) optim.zero_grad(set_to_none=True)
opt_step += 1 opt_step += 1
raw = float(task.detach()) raw = float(task.detach())
if raw < best: if raw < best_train:
best = raw best_train = raw
next_row = row_index + args.batch
if tracker is not None: if tracker is not None:
tracker.log( tracker.log(
{ {
@@ -173,9 +315,14 @@ def main() -> None:
}, },
step=step, step=step,
) )
if step % args.eval_every == 0 or step == max_micro - 1: eval_now = step % args.eval_every == 0 or step == max_micro - 1
ckpt_now = args.ckpt_every > 0 and step > 0 and (
step % args.ckpt_every == 0 or step == max_micro - 1
)
if eval_now:
print( print(
f"step {step:4d} sft loss {raw:.4f} lr {optim.param_groups[0]['lr']:.2e}" f"step {step:4d} sft loss {raw:.4f} "
f"lr {optim.param_groups[0]['lr']:.2e}"
) )
if args.src and args.ref: if args.src and args.ref:
model.eval() model.eval()
@@ -192,6 +339,8 @@ def main() -> None:
) )
printable = {k: v for k, v in out.items() if k != "hyps"} printable = {k: v for k, v in out.items() if k != "hyps"}
print(printable) print(printable)
for i, hyp in enumerate((out.get("hyps") or [])[:2]):
print(f" hyp[{i}] {hyp}")
if tracker is not None: if tracker is not None:
tracker.log( tracker.log(
{ {
@@ -201,22 +350,41 @@ def main() -> None:
}, },
step=step, step=step,
) )
model.train() if step > 0 and _is_better_eval(
step += 1 printable["success_rate"],
printable["chrf"],
os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True) best_success,
torch.save( best_chrf,
{ ):
"config": asdict(cfg), best_success = float(printable["success_rate"])
"model_state": model.state_dict(), best_chrf = float(printable["chrf"])
"optimizer_state": optim.state_dict(), _save(best_path, dump())
"tokenizer": tok_src, print(
"sft_data": args.data, f" best success {best_success:.2f} "
"pretrained_ckpt": args.ckpt, f"chrf {best_chrf:.2f} -> {best_path}"
},
args.out,
) )
print(f"best sft loss {best:.4f}; checkpoint -> {args.out}") model.train()
if device == "cuda":
torch.cuda.empty_cache()
if ckpt_now:
_save(last_path, dump())
print(f" last -> {last_path}")
step += 1
payload = dump()
_save(last_path, payload)
_save(args.out, payload)
print(
f"best train {best_train:.4f} best success {best_success:.2f} "
f"chrf {best_chrf:.2f}; checkpoint -> {args.out}"
)
except KeyboardInterrupt:
print("interrupt; writing last checkpoint")
_save(last_path, dump())
print(f" last -> {last_path}")
if os.path.isfile(best_path):
print(f" best remains {best_path}")
raise SystemExit(130) from None
finally:
if tracker is not None: if tracker is not None:
tracker.finish() tracker.finish()