Compare commits
16
Commits
584f7e9e73
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a2c4217dae | ||
|
|
ea7167b3f7 | ||
|
|
94d0f2ff6a | ||
|
|
24c9d56b72 | ||
|
|
5cc0555563 | ||
|
|
5a7d949b01 | ||
|
|
9652a9a7eb | ||
|
|
071dfaf42c | ||
|
|
9a4862a866 | ||
|
|
47c72e5bb8 | ||
|
|
e7185cbf49 | ||
|
|
53d0f4b17a | ||
|
|
8442f92c58 | ||
|
|
49aede9cb2 | ||
|
|
7a12f61de1 | ||
|
|
d1da0816f2 |
@@ -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。
|
||||||
|
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ import torch
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from einops import rearrange
|
from einops import rearrange
|
||||||
from torch import Tensor, nn
|
from torch import Tensor, nn
|
||||||
|
from torch.utils.checkpoint import checkpoint as activation_checkpoint
|
||||||
|
|
||||||
|
|
||||||
ATTNRES_MODES = ("off", "full", "block")
|
ATTNRES_MODES = ("off", "full", "block")
|
||||||
@@ -247,6 +248,7 @@ class BlockAttnResStack(nn.Module):
|
|||||||
if is_final_aggregate
|
if is_final_aggregate
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
self.gradient_checkpointing = False
|
||||||
|
|
||||||
def forward_naive(self, x: Tensor) -> Tensor:
|
def forward_naive(self, x: Tensor) -> Tensor:
|
||||||
blocks = [x] # b_0=embedding/input representation
|
blocks = [x] # b_0=embedding/input representation
|
||||||
@@ -294,8 +296,19 @@ class BlockAttnResStack(nn.Module):
|
|||||||
blocks = [x]
|
blocks = [x]
|
||||||
depth = len(self.layers)
|
depth = len(self.layers)
|
||||||
start = 0
|
start = 0
|
||||||
|
use_ckpt = (
|
||||||
|
self.gradient_checkpointing and self.training and torch.is_grad_enabled()
|
||||||
|
)
|
||||||
while start < depth:
|
while start < depth:
|
||||||
end = min(start + self.block_size, depth)
|
end = min(start + self.block_size, depth)
|
||||||
|
if use_ckpt:
|
||||||
|
def _run(*srcs, _start=start, _end=end):
|
||||||
|
return self._run_block_two_phase(list(srcs), _start, _end)
|
||||||
|
|
||||||
|
blocks.append(
|
||||||
|
activation_checkpoint(_run, *blocks, use_reentrant=False)
|
||||||
|
)
|
||||||
|
else:
|
||||||
blocks.append(self._run_block_two_phase(blocks, start, end))
|
blocks.append(self._run_block_two_phase(blocks, start, end))
|
||||||
start = end
|
start = end
|
||||||
|
|
||||||
|
|||||||
+123
-16
@@ -9,12 +9,14 @@ SiTU-GLU: gate = β1·tanh(W_g x/β1)⊙σ(W_g x); up = β2·tanh(W_u x/β2)
|
|||||||
||SiTU-GLU||_∞ ≤ β1·β2 (=100), 原点附近≈SwiGLU, 远端软饱和防低精度溢出.
|
||SiTU-GLU||_∞ ≤ β1·β2 (=100), 原点附近≈SwiGLU, 远端软饱和防低精度溢出.
|
||||||
E: R^in → R^in (内部中间维 d_ff).
|
E: R^in → R^in (内部中间维 d_ff).
|
||||||
|
|
||||||
Router: Top-k logits 基于全宽 x (笔记 Topk(x)); 归一化权重取 softmax(topk).
|
Router (K3 eq.13): s=σ(W_r x), Top-k(s+b), p_i = s_i / Σ_{j∈T} s_j.
|
||||||
|
Routed 执行: permute-dispatch, pad 到 [R, C, ℓ], 三次 bmm(SiTU 参数结构不变).
|
||||||
|
训练: Switch/GShard aux + router z-loss, 由 train loop 加到 CE 上.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
from .rmsnorm import RMSNorm
|
from .rmsnorm import RMSNorm
|
||||||
@@ -23,7 +25,9 @@ from .rmsnorm import RMSNorm
|
|||||||
class SiTU(nn.Module):
|
class SiTU(nn.Module):
|
||||||
"""SiTU-GLU expert: gate 支软上限 β1, up 支软上限 β2, 输出回到输入维."""
|
"""SiTU-GLU expert: gate 支软上限 β1, up 支软上限 β2, 输出回到输入维."""
|
||||||
|
|
||||||
def __init__(self, dim_in: int, dim_ff: int, beta1: float = 4.0, beta2: float = 25.0):
|
def __init__(
|
||||||
|
self, dim_in: int, dim_ff: int, beta1: float = 4.0, beta2: float = 25.0
|
||||||
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.beta1, self.beta2 = beta1, beta2
|
self.beta1, self.beta2 = beta1, beta2
|
||||||
self.w_g = nn.Linear(dim_in, dim_ff, bias=False)
|
self.w_g = nn.Linear(dim_in, dim_ff, bias=False)
|
||||||
@@ -48,14 +52,19 @@ class LatentMoE(nn.Module):
|
|||||||
d_ff: int,
|
d_ff: int,
|
||||||
beta1: float = 4.0,
|
beta1: float = 4.0,
|
||||||
beta2: float = 25.0,
|
beta2: float = 25.0,
|
||||||
|
aux_loss_coef: float = 1e-2,
|
||||||
|
z_loss_coef: float = 1e-3,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.latent_size = latent_size
|
self.latent_size = latent_size
|
||||||
self.n_routed = n_routed
|
self.n_routed = n_routed
|
||||||
self.top_k = top_k
|
self.top_k = top_k
|
||||||
|
self.aux_loss_coef = aux_loss_coef
|
||||||
|
self.z_loss_coef = z_loss_coef
|
||||||
|
|
||||||
self.down = nn.Linear(hidden_size, latent_size, bias=False) # W↓
|
self.down = nn.Linear(hidden_size, latent_size, bias=False) # W↓
|
||||||
self.router = nn.Linear(hidden_size, n_routed, bias=False) # Top-k logits
|
self.router = nn.Linear(hidden_size, n_routed, bias=False) # logits → sigmoid
|
||||||
|
self.register_buffer("expert_bias", torch.zeros(n_routed), persistent=False)
|
||||||
self.shared = nn.ModuleList(
|
self.shared = nn.ModuleList(
|
||||||
[SiTU(hidden_size, d_ff, beta1, beta2) for _ in range(n_shared)]
|
[SiTU(hidden_size, d_ff, beta1, beta2) for _ in range(n_shared)]
|
||||||
)
|
)
|
||||||
@@ -65,6 +74,9 @@ class LatentMoE(nn.Module):
|
|||||||
self.norm = RMSNorm(latent_size)
|
self.norm = RMSNorm(latent_size)
|
||||||
self.up = nn.Linear(latent_size, hidden_size, bias=False) # W↑
|
self.up = nn.Linear(latent_size, hidden_size, bias=False) # W↑
|
||||||
self.last_route_ids: torch.Tensor | None = None
|
self.last_route_ids: torch.Tensor | None = None
|
||||||
|
self.last_capacity: int = 0
|
||||||
|
self.last_aux_loss: torch.Tensor | None = None
|
||||||
|
self.last_z_loss: torch.Tensor | None = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_config(cls, config) -> LatentMoE:
|
def from_config(cls, config) -> LatentMoE:
|
||||||
@@ -77,27 +89,104 @@ class LatentMoE(nn.Module):
|
|||||||
config.moe_d_ff,
|
config.moe_d_ff,
|
||||||
config.situ_beta1,
|
config.situ_beta1,
|
||||||
config.situ_beta2,
|
config.situ_beta2,
|
||||||
|
float(getattr(config, "moe_aux_loss_coef", 1e-2)),
|
||||||
|
float(getattr(config, "moe_z_loss_coef", 1e-3)),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _routed_u(
|
||||||
|
self, z: torch.Tensor, ids: torch.Tensor, probs: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Permute-dispatch + pad to [R, C, ℓ] + 3 bmm + scatter-add.
|
||||||
|
|
||||||
|
``C = max(counts)``: FLOPs are ``R·C``, not ``sum(counts)``. Padding
|
||||||
|
slots are not gathered, so they contribute zero gradient. Empty
|
||||||
|
experts stay in the stacked weights (padded grouped GEMM).
|
||||||
|
"""
|
||||||
|
B, T, ell = z.shape
|
||||||
|
N = B * T
|
||||||
|
R, k = self.n_routed, self.top_k
|
||||||
|
device = z.device
|
||||||
|
|
||||||
|
tok = torch.arange(N, device=device).unsqueeze(1).expand(N, k).reshape(-1)
|
||||||
|
eid = ids.reshape(-1)
|
||||||
|
pw = probs.reshape(-1)
|
||||||
|
|
||||||
|
order = eid.argsort(stable=True)
|
||||||
|
tok, eid, pw = tok[order], eid[order], pw[order]
|
||||||
|
|
||||||
|
counts = torch.bincount(eid, minlength=R)
|
||||||
|
offsets = counts.cumsum(0) - counts
|
||||||
|
local_pos = torch.arange(N * k, device=device) - offsets[eid]
|
||||||
|
C = int(counts.max().item()) if eid.numel() else 0
|
||||||
|
self.last_capacity = C
|
||||||
|
|
||||||
|
u_flat = z.new_zeros(N, ell)
|
||||||
|
if C == 0 or eid.numel() == 0:
|
||||||
|
return u_flat.view(B, T, ell)
|
||||||
|
|
||||||
|
dtype = z.dtype
|
||||||
|
gathered = z.reshape(N, ell)[tok]
|
||||||
|
padded = torch.index_put(
|
||||||
|
gathered.new_zeros(R, C, ell), (eid, local_pos), gathered
|
||||||
|
)
|
||||||
|
|
||||||
|
w_g = torch.stack([e.w_g.weight for e in self.experts]).to(dtype)
|
||||||
|
w_u = torch.stack([e.w_u.weight for e in self.experts]).to(dtype)
|
||||||
|
w_o = torch.stack([e.w_o.weight for e in self.experts]).to(dtype)
|
||||||
|
beta1 = padded.new_tensor(self.experts[0].beta1)
|
||||||
|
beta2 = padded.new_tensor(self.experts[0].beta2)
|
||||||
|
|
||||||
|
wg = torch.bmm(padded, w_g.transpose(-1, -2)) # [R, C, ff]
|
||||||
|
g = beta1 * torch.tanh(wg / beta1) * torch.sigmoid(wg)
|
||||||
|
wu = torch.bmm(padded, w_u.transpose(-1, -2))
|
||||||
|
hidden = beta2 * torch.tanh(wu / beta2)
|
||||||
|
out = torch.bmm(g * hidden, w_o.transpose(-1, -2)) # [R, C, ℓ]
|
||||||
|
|
||||||
|
weighted = pw.to(dtype).unsqueeze(-1) * out[eid, local_pos]
|
||||||
|
u_flat = u_flat.index_add(0, tok, weighted)
|
||||||
|
return u_flat.view(B, T, ell)
|
||||||
|
|
||||||
|
def _route(self, logits: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""K3: s=σ(l), T=TopK(s+b), p_i = s_i / Σ_{j∈T} s_j. Bias does not enter p."""
|
||||||
|
scores = torch.sigmoid(logits)
|
||||||
|
ids = torch.topk(
|
||||||
|
scores + self.expert_bias.to(dtype=scores.dtype), self.top_k, dim=-1
|
||||||
|
).indices
|
||||||
|
selected = scores.gather(-1, ids)
|
||||||
|
probs = selected / selected.sum(dim=-1, keepdim=True).clamp_min(1e-9)
|
||||||
|
return ids, probs
|
||||||
|
|
||||||
|
def _balancing_losses(
|
||||||
|
self, logits: torch.Tensor, ids: torch.Tensor
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Switch/GShard aux on sigmoid scores; z-loss on raw logits."""
|
||||||
|
routed = self.n_routed
|
||||||
|
flat_logits = logits.reshape(-1, routed).float()
|
||||||
|
scores = torch.sigmoid(flat_logits)
|
||||||
|
counts = torch.bincount(ids.reshape(-1), minlength=routed).to(
|
||||||
|
dtype=scores.dtype
|
||||||
|
)
|
||||||
|
frac = counts / counts.sum().clamp_min(1.0)
|
||||||
|
prob_mean = scores.mean(dim=0)
|
||||||
|
aux = routed * (frac * prob_mean).sum()
|
||||||
|
z_loss = torch.logsumexp(flat_logits, dim=-1).square().mean()
|
||||||
|
return self.aux_loss_coef * aux, self.z_loss_coef * z_loss
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor):
|
def forward(self, x: torch.Tensor):
|
||||||
B, T, _ = x.shape
|
B, T, _ = x.shape
|
||||||
z = self.down(x) # [B, T, ℓ]
|
z = self.down(x) # [B, T, ℓ]
|
||||||
|
|
||||||
logits = self.router(x) # [B, T, n_routed]
|
logits = self.router(x) # [B, T, n_routed]
|
||||||
topk = torch.topk(logits, self.top_k, dim=-1)
|
ids, probs = self._route(logits)
|
||||||
ids = topk.indices # [B, T, k]
|
|
||||||
self.last_route_ids = ids.detach()
|
self.last_route_ids = ids.detach()
|
||||||
probs = F.softmax(topk.values, dim=-1) # [B, T, k]
|
if self.training and (self.aux_loss_coef != 0.0 or self.z_loss_coef != 0.0):
|
||||||
|
self.last_aux_loss, self.last_z_loss = self._balancing_losses(logits, ids)
|
||||||
# 向量化 routed: 预计算全部专家输出, 按 token 的 Top-k id 取
|
else:
|
||||||
all_out = torch.stack([e(z) for e in self.experts]) # [R, B, T, ℓ]
|
zero = logits.new_zeros(())
|
||||||
all_out = all_out.permute(1, 2, 0, 3).reshape(B * T, self.n_routed, self.latent_size)
|
self.last_aux_loss = zero
|
||||||
u = torch.zeros(B, T, self.latent_size, device=x.device, dtype=x.dtype)
|
self.last_z_loss = zero
|
||||||
for i in range(self.top_k):
|
|
||||||
idx = ids[:, :, i].reshape(B * T) # [B*T]
|
|
||||||
sel = all_out[torch.arange(B * T, device=x.device), idx] # [B*T, ℓ]
|
|
||||||
u += probs[:, :, i : i + 1] * sel.reshape(B, T, self.latent_size)
|
|
||||||
|
|
||||||
|
u = self._routed_u(z, ids, probs)
|
||||||
shared_out = torch.stack([e(x) for e in self.shared]).sum(0) # [B, T, d]
|
shared_out = torch.stack([e(x) for e in self.shared]).sum(0) # [B, T, d]
|
||||||
return shared_out + self.up(self.norm(u))
|
return shared_out + self.up(self.norm(u))
|
||||||
|
|
||||||
@@ -116,3 +205,21 @@ def moe_route_frac(model: nn.Module) -> torch.Tensor | None:
|
|||||||
return None
|
return None
|
||||||
stacked = torch.stack(hists).sum(0)
|
stacked = torch.stack(hists).sum(0)
|
||||||
return stacked / stacked.sum().clamp_min(1.0)
|
return stacked / stacked.sum().clamp_min(1.0)
|
||||||
|
|
||||||
|
|
||||||
|
def moe_router_losses(model: nn.Module) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Sum Switch aux and z-loss over every LatentMoE layer (0 if none)."""
|
||||||
|
auxes: list[torch.Tensor] = []
|
||||||
|
z_losses: list[torch.Tensor] = []
|
||||||
|
for module in model.modules():
|
||||||
|
if not isinstance(module, LatentMoE):
|
||||||
|
continue
|
||||||
|
if module.last_aux_loss is None or module.last_z_loss is None:
|
||||||
|
continue
|
||||||
|
auxes.append(module.last_aux_loss)
|
||||||
|
z_losses.append(module.last_z_loss)
|
||||||
|
if not auxes:
|
||||||
|
param = next(model.parameters(), None)
|
||||||
|
zero = param.new_zeros(()) if param is not None else torch.zeros(())
|
||||||
|
return zero, zero
|
||||||
|
return torch.stack(auxes).sum(), torch.stack(z_losses).sum()
|
||||||
|
|||||||
+7
-13
@@ -78,20 +78,14 @@ class GatedMLA(nn.Module):
|
|||||||
w_uk = w[: H * self.qk_nope_head_dim].view(H, self.qk_nope_head_dim, r)
|
w_uk = w[: H * self.qk_nope_head_dim].view(H, self.qk_nope_head_dim, r)
|
||||||
w_uv = w[H * self.qk_nope_head_dim :].view(H, self.v_head_dim, r)
|
w_uv = w[H * self.qk_nope_head_dim :].view(H, self.v_head_dim, r)
|
||||||
|
|
||||||
# 吸收 W_UK 进 query: score = (q @ W_UK^T) @ c^T
|
# 吸收 W_UK 进 query: score = (q @ W_UK^T) @ c^T, scale=1 matches the unscaled einsum.
|
||||||
q_absorb = torch.einsum("bthd,hdj->bthj", q, w_uk) # [B, T, H, r]
|
q_absorb = torch.einsum("bthd,hdj->bthj", q, w_uk) # [B, T, H, r]
|
||||||
scores = torch.einsum("bthj,bsj->bhts", q_absorb, c) # [B, H, T, T]
|
q_h = q_absorb.transpose(1, 2) # [B, H, T, r]
|
||||||
|
kv = c.unsqueeze(1).expand(B, H, T, r)
|
||||||
mask = torch.triu(
|
latent_out = F.scaled_dot_product_attention(
|
||||||
torch.ones(T, T, dtype=torch.bool, device=x.device), diagonal=1
|
q_h, kv, kv, is_causal=True, scale=1.0
|
||||||
)
|
) # [B, H, T, r]
|
||||||
scores = scores.masked_fill(mask, float("-inf"))
|
o_heads = torch.einsum("bhtr,hvr->bhtv", latent_out, w_uv)
|
||||||
attn = F.softmax(scores, dim=-1) # [B, H, T, T]
|
|
||||||
|
|
||||||
# 先在 latent 加权, 再乘 W_UV^T 还原 v —— 永不解压
|
|
||||||
latent_out = torch.einsum("bhts,bsj->bhtj", attn, c) # [B, H, T, r]
|
|
||||||
o_heads = torch.einsum("bhtj,hvj->bhtv", latent_out, w_uv) # [B, H, T, d_v]
|
|
||||||
|
|
||||||
o_heads = o_heads.transpose(1, 2).reshape(B, T, H * self.v_head_dim)
|
o_heads = o_heads.transpose(1, 2).reshape(B, T, H * self.v_head_dim)
|
||||||
gate = torch.sigmoid(self.gate(x)) # [B, T, H*d_v]
|
gate = torch.sigmoid(self.gate(x)) # [B, T, H*d_v]
|
||||||
return self.o_proj(gate * o_heads) # [B, T, d]
|
return self.o_proj(gate * o_heads) # [B, T, d]
|
||||||
|
|||||||
+27
-8
@@ -26,6 +26,28 @@ from ..layers.block import DecoderBlock
|
|||||||
from ..layers.rmsnorm import RMSNorm
|
from ..layers.rmsnorm import RMSNorm
|
||||||
|
|
||||||
|
|
||||||
|
def _chunked_linear_cross_entropy(
|
||||||
|
hidden: torch.Tensor,
|
||||||
|
weight: torch.Tensor,
|
||||||
|
labels: torch.Tensor,
|
||||||
|
ignore_index: int = -100,
|
||||||
|
chunk_size: int = 256,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""CE without materializing [B, T, vocab]. Match mean reduction over valid labels."""
|
||||||
|
features = hidden[:, :-1].reshape(-1, hidden.size(-1))
|
||||||
|
targets = labels[:, 1:].reshape(-1)
|
||||||
|
total = hidden.new_zeros(())
|
||||||
|
n_valid = hidden.new_zeros((), dtype=torch.long)
|
||||||
|
for start in range(0, features.size(0), chunk_size):
|
||||||
|
sl = slice(start, start + chunk_size)
|
||||||
|
logits = F.linear(features[sl], weight)
|
||||||
|
total = total + F.cross_entropy(
|
||||||
|
logits, targets[sl], ignore_index=ignore_index, reduction="sum"
|
||||||
|
)
|
||||||
|
n_valid = n_valid + (targets[sl] != ignore_index).sum()
|
||||||
|
return total / n_valid.clamp_min(1).to(dtype=total.dtype)
|
||||||
|
|
||||||
|
|
||||||
def _build_mixer(config, blocks: nn.ModuleList):
|
def _build_mixer(config, blocks: nn.ModuleList):
|
||||||
mode = getattr(config, "attnres", "off")
|
mode = getattr(config, "attnres", "off")
|
||||||
if mode == "off":
|
if mode == "off":
|
||||||
@@ -89,17 +111,14 @@ class CausalLM(nn.Module):
|
|||||||
x = activation_checkpoint(block, x, use_reentrant=False)
|
x = activation_checkpoint(block, x, use_reentrant=False)
|
||||||
else:
|
else:
|
||||||
x = block(x)
|
x = block(x)
|
||||||
elif self.gradient_checkpointing and self.training:
|
|
||||||
x = activation_checkpoint(self.mixer, x, use_reentrant=False)
|
|
||||||
else:
|
else:
|
||||||
|
self.mixer.gradient_checkpointing = self.gradient_checkpointing
|
||||||
x = self.mixer(x)
|
x = self.mixer(x)
|
||||||
logits = self.lm_head(self.norm(x))
|
hidden = self.norm(x)
|
||||||
if labels is None:
|
if labels is None:
|
||||||
return logits
|
return self.lm_head(hidden)
|
||||||
return F.cross_entropy(
|
return _chunked_linear_cross_entropy(
|
||||||
logits[:, :-1].reshape(-1, logits.size(-1)),
|
hidden, self.lm_head.weight, labels, ignore_index=ignore_index
|
||||||
labels[:, 1:].reshape(-1),
|
|
||||||
ignore_index=ignore_index,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
@torch.inference_mode()
|
@torch.inference_mode()
|
||||||
|
|||||||
+10
-6
@@ -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
|
||||||
@@ -53,6 +55,8 @@ class K3Config:
|
|||||||
moe_d_ff: int = 96
|
moe_d_ff: int = 96
|
||||||
situ_beta1: float = 4.0
|
situ_beta1: float = 4.0
|
||||||
situ_beta2: float = 25.0
|
situ_beta2: float = 25.0
|
||||||
|
moe_aux_loss_coef: float = 1e-2 # Switch/GShard N Σ f_e P_e
|
||||||
|
moe_z_loss_coef: float = 1e-3 # mean (logsumexp logits)^2
|
||||||
|
|
||||||
kda_backend: str = "reference"
|
kda_backend: str = "reference"
|
||||||
|
|
||||||
@@ -73,12 +77,12 @@ class K3Config:
|
|||||||
if name == "toy":
|
if name == "toy":
|
||||||
return cls()
|
return cls()
|
||||||
if name in {"0.5b", "500m"}:
|
if name in {"0.5b", "500m"}:
|
||||||
# H * head_dim == hidden. Routed 16: LatentMoE still runs every expert.
|
# 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,
|
||||||
|
|||||||
+122
-5
@@ -19,12 +19,15 @@ import torch
|
|||||||
from .prompts import instruction_prompt
|
from .prompts import instruction_prompt
|
||||||
|
|
||||||
WIKI_SHARD_TOTAL = {"zh": 6, "en": 41}
|
WIKI_SHARD_TOTAL = {"zh": 6, "en": 41}
|
||||||
WIKI_BASE = (
|
WIKI_PATH = "datasets/wikimedia/wikipedia/resolve/main/20231101.{lang}"
|
||||||
"https://huggingface.co/datasets/wikimedia/wikipedia/resolve/main/20231101.{lang}"
|
|
||||||
)
|
|
||||||
IGNORE_INDEX = -100
|
IGNORE_INDEX = -100
|
||||||
|
|
||||||
|
|
||||||
|
def _hf_endpoint() -> str:
|
||||||
|
"""Hub origin. OpenBayes/CN: export HF_ENDPOINT=https://hf-mirror.com"""
|
||||||
|
return os.environ.get("HF_ENDPOINT", "https://huggingface.co").rstrip("/")
|
||||||
|
|
||||||
|
|
||||||
class Tokenizer(Protocol):
|
class Tokenizer(Protocol):
|
||||||
vocab_size: int
|
vocab_size: int
|
||||||
|
|
||||||
@@ -54,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))
|
||||||
@@ -72,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)
|
||||||
|
|
||||||
|
|
||||||
@@ -86,12 +95,23 @@ def pretrain_dir() -> Path:
|
|||||||
return Path("data/pretrain")
|
return Path("data/pretrain")
|
||||||
|
|
||||||
|
|
||||||
|
def sft_dir() -> Path:
|
||||||
|
for candidate in (
|
||||||
|
os.environ.get("KDA_SFT_DIR"),
|
||||||
|
"/data/sft",
|
||||||
|
"data/sft",
|
||||||
|
):
|
||||||
|
if candidate and Path(candidate).is_dir():
|
||||||
|
return Path(candidate)
|
||||||
|
return Path("data/sft")
|
||||||
|
|
||||||
|
|
||||||
def _wiki_files(lang: str, n_shards: int) -> list[str]:
|
def _wiki_files(lang: str, n_shards: int) -> list[str]:
|
||||||
if lang not in WIKI_SHARD_TOTAL:
|
if lang not in WIKI_SHARD_TOTAL:
|
||||||
raise ValueError(f"unsupported wiki lang {lang!r}; expected zh or en")
|
raise ValueError(f"unsupported wiki lang {lang!r}; expected zh or en")
|
||||||
total = WIKI_SHARD_TOTAL[lang]
|
total = WIKI_SHARD_TOTAL[lang]
|
||||||
n = min(max(n_shards, 1), total)
|
n = min(max(n_shards, 1), total)
|
||||||
base = WIKI_BASE.format(lang=lang)
|
base = f"{_hf_endpoint()}/{WIKI_PATH.format(lang=lang)}"
|
||||||
return [f"{base}/train-{i:05d}-of-{total:05d}.parquet" for i in range(n)]
|
return [f"{base}/train-{i:05d}-of-{total:05d}.parquet" for i in range(n)]
|
||||||
|
|
||||||
|
|
||||||
@@ -284,6 +304,103 @@ def encode_sft_row(
|
|||||||
return ids, labels
|
return ids, labels
|
||||||
|
|
||||||
|
|
||||||
|
def _eval_blocklist(eval_dir: str | Path | None = None) -> set[str]:
|
||||||
|
"""Frozen eval sentences must not appear in SFT bitext."""
|
||||||
|
blocked: set[str] = set()
|
||||||
|
folders = []
|
||||||
|
if eval_dir is not None:
|
||||||
|
folders.append(Path(eval_dir))
|
||||||
|
folders.extend(
|
||||||
|
[
|
||||||
|
Path(os.environ["KDA_EVAL_DIR"]) if os.environ.get("KDA_EVAL_DIR") else None,
|
||||||
|
Path("/data/eval"),
|
||||||
|
Path("data/eval"),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
for folder in folders:
|
||||||
|
if folder is None or not folder.is_dir():
|
||||||
|
continue
|
||||||
|
for path in folder.glob("*.txt"):
|
||||||
|
for line in path.read_text(encoding="utf-8").splitlines():
|
||||||
|
text = line.strip()
|
||||||
|
if text:
|
||||||
|
blocked.add(text)
|
||||||
|
return blocked
|
||||||
|
|
||||||
|
|
||||||
|
def fetch_opus100_enzh(
|
||||||
|
limit: int,
|
||||||
|
*,
|
||||||
|
both_dirs: bool = True,
|
||||||
|
cache_dir: str | Path | None = None,
|
||||||
|
eval_dir: str | Path | None = None,
|
||||||
|
) -> list[dict]:
|
||||||
|
"""Stream Helsinki-NLP/opus-100 ``en-zh`` train. ``limit`` is source pairs."""
|
||||||
|
if limit < 1:
|
||||||
|
raise ValueError(f"limit must be >= 1, got {limit}")
|
||||||
|
cache = Path(cache_dir) if cache_dir is not None else sft_dir()
|
||||||
|
cache.mkdir(parents=True, exist_ok=True)
|
||||||
|
tag = "both" if both_dirs else "enzh"
|
||||||
|
path = cache / f"opus100-en-zh-{tag}-limit{limit}.jsonl"
|
||||||
|
if path.exists():
|
||||||
|
rows = load_sft_rows(path)
|
||||||
|
if rows:
|
||||||
|
return rows
|
||||||
|
|
||||||
|
from datasets import load_dataset
|
||||||
|
|
||||||
|
ds = load_dataset("Helsinki-NLP/opus-100", "en-zh", split="train", streaming=True)
|
||||||
|
blocked = _eval_blocklist(eval_dir)
|
||||||
|
rows: list[dict] = []
|
||||||
|
n_src = 0
|
||||||
|
for row in ds:
|
||||||
|
trans = row.get("translation") if isinstance(row, dict) else None
|
||||||
|
blob = trans if isinstance(trans, dict) else row
|
||||||
|
en = str(blob.get("en") or "").strip()
|
||||||
|
zh = str(blob.get("zh") or "").strip()
|
||||||
|
if not en or not zh or en == zh:
|
||||||
|
continue
|
||||||
|
if en in blocked or zh in blocked:
|
||||||
|
continue
|
||||||
|
if min(len(en), len(zh)) < 2:
|
||||||
|
continue
|
||||||
|
n_src += 1
|
||||||
|
rows.append({"src": zh, "tgt": en, "target_lang": "en"})
|
||||||
|
if both_dirs:
|
||||||
|
rows.append({"src": en, "tgt": zh, "target_lang": "zh"})
|
||||||
|
if n_src >= limit:
|
||||||
|
break
|
||||||
|
tmp = path.with_suffix(path.suffix + ".tmp")
|
||||||
|
with tmp.open("w", encoding="utf-8") as fh:
|
||||||
|
for row in rows:
|
||||||
|
fh.write(json.dumps(row, ensure_ascii=False) + "\n")
|
||||||
|
tmp.replace(path)
|
||||||
|
return rows
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_sft_rows(
|
||||||
|
source: str,
|
||||||
|
*,
|
||||||
|
limit: int = 100_000,
|
||||||
|
both_dirs: bool = True,
|
||||||
|
cache_dir: str | Path | None = None,
|
||||||
|
eval_dir: str | Path | None = None,
|
||||||
|
) -> list[dict]:
|
||||||
|
"""Local jsonl/tsv, or ``opus-100`` / ``opus`` to pull OPUS-100 en-zh from HF."""
|
||||||
|
path = Path(source)
|
||||||
|
if path.is_file():
|
||||||
|
return load_sft_rows(path)
|
||||||
|
key = source.strip().lower().replace("_", "-")
|
||||||
|
if key in {"opus", "opus-100", "opus100", "helsinki-nlp/opus-100"}:
|
||||||
|
print(f"fetching OPUS-100 en-zh (limit {limit} pairs, both_dirs={both_dirs})")
|
||||||
|
return fetch_opus100_enzh(
|
||||||
|
limit, both_dirs=both_dirs, cache_dir=cache_dir, eval_dir=eval_dir
|
||||||
|
)
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"SFT source {source!r} is not a file; use a jsonl path or 'opus-100'"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def load_sft_rows(path: str | Path) -> list[dict]:
|
def load_sft_rows(path: str | Path) -> list[dict]:
|
||||||
"""jsonl ``{src,tgt,target_lang}`` or TSV ``src\\ttgt\\ttarget_lang``."""
|
"""jsonl ``{src,tgt,target_lang}`` or TSV ``src\\ttgt\\ttarget_lang``."""
|
||||||
p = Path(path)
|
p = Path(path)
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|
||||||
|
|||||||
@@ -34,14 +34,16 @@ def total_opt_steps(
|
|||||||
seq_len: int,
|
seq_len: int,
|
||||||
grad_acc: int,
|
grad_acc: int,
|
||||||
) -> int:
|
) -> int:
|
||||||
"""Optimizer-step horizon used by cosine. At least 1."""
|
"""Optimizer-step horizon used by cosine. At least 1.
|
||||||
|
|
||||||
|
``max_tokens`` is the training budget when set; ``max_micro`` is only used
|
||||||
|
when ``max_tokens`` is None. Otherwise a default ``--steps 2000`` would
|
||||||
|
shrink a 1B-token cosine to 250 opt steps.
|
||||||
|
"""
|
||||||
acc = max(grad_acc, 1)
|
acc = max(grad_acc, 1)
|
||||||
candidates: list[int] = []
|
|
||||||
if max_tokens is not None and max_tokens > 0:
|
if max_tokens is not None and max_tokens > 0:
|
||||||
tpm = max(tokens_per_micro(batch, seq_len), 1)
|
tpm = max(tokens_per_micro(batch, seq_len), 1)
|
||||||
candidates.append(math.ceil(max_tokens / (tpm * acc)))
|
return max(math.ceil(max_tokens / (tpm * acc)), 1)
|
||||||
if max_micro is not None and max_micro > 0:
|
if max_micro is not None and max_micro > 0:
|
||||||
candidates.append(math.ceil(max_micro / acc))
|
return max(math.ceil(max_micro / acc), 1)
|
||||||
if not candidates:
|
|
||||||
return 1
|
return 1
|
||||||
return max(min(candidates), 1)
|
|
||||||
|
|||||||
+27
-11
@@ -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(
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
|
|||||||
@@ -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}
|
||||||
|
|||||||
Binary file not shown.
+145
-40
@@ -1,8 +1,9 @@
|
|||||||
% teach:
|
% teach:
|
||||||
% gap: 读者知道 MoE 的 top-k 路由但不知道 LatentMoE 的 latent 接口和 SiTU-GLU
|
% gap: 读者知道 MoE 的 top-k 路由但不知道 LatentMoE 的 latent 接口、稀疏 dispatch 和负载均衡损失
|
||||||
% takeaway: LatentMoE 通过 latent 接口把 routed 专家限制在 ℓ=d/2 上算, SiTU-GLU 用软上限防溢出
|
% takeaway: LatentMoE = shared 全宽 + routed 半宽 latent;路由用 K3 sigmoid-TopK + L1 归一化;
|
||||||
% jump: 为什么 routed 专家在 latent 空间而 shared 在全宽?省参数
|
% 执行用 permute-dispatch + padded bmm(每个 token 只算 k 个专家);训练加 Switch/GShard aux + z-loss
|
||||||
% omit: load balancing loss
|
% jump: 为什么不用 dense stack 算全部专家?稀疏 dispatch 让算力只随 k 不随 R 涨
|
||||||
|
% omit: none
|
||||||
|
|
||||||
\section{SiTU-GLU 与 Stable LatentMoE}
|
\section{SiTU-GLU 与 Stable LatentMoE}
|
||||||
\splabel{C5}
|
\splabel{C5}
|
||||||
@@ -54,7 +55,7 @@ class SiTU(nn.Module):
|
|||||||
\textbf{Shared 专家} & $n_{\mathrm{shared}}$ 个 SiTU,全宽 $d \to d$,所有 token 都经过 \\
|
\textbf{Shared 专家} & $n_{\mathrm{shared}}$ 个 SiTU,全宽 $d \to d$,所有 token 都经过 \\
|
||||||
\textbf{Routed 专家} & $n_{\mathrm{routed}}$ 个 SiTU,半宽 $\ell \to \ell$($\ell = d/2$) \\
|
\textbf{Routed 专家} & $n_{\mathrm{routed}}$ 个 SiTU,半宽 $\ell \to \ell$($\ell = d/2$) \\
|
||||||
\textbf{Latent 接口} & $W_\downarrow: d \to \ell$, $W_\uparrow: \ell \to d$(压缩/还原) \\
|
\textbf{Latent 接口} & $W_\downarrow: d \to \ell$, $W_\uparrow: \ell \to d$(压缩/还原) \\
|
||||||
\textbf{Router} & $W_r: d \to n_{\mathrm{routed}}$,Top-k 选择 + softmax 归一化 \\
|
\textbf{Router} & $W_r: d \to n_{\mathrm{routed}}$,K3:$s=\sigma(W_r x)$,Top-$k(s+b)$,$p_i=s_i/\sum_{j\in T}s_j$ \\
|
||||||
\bottomrule
|
\bottomrule
|
||||||
\end{tabular}
|
\end{tabular}
|
||||||
\end{center}
|
\end{center}
|
||||||
@@ -67,29 +68,31 @@ class SiTU(nn.Module):
|
|||||||
z = W_\downarrow \cdot x \qquad \shape{B, T, \ell}
|
z = W_\downarrow \cdot x \qquad \shape{B, T, \ell}
|
||||||
\]
|
\]
|
||||||
|
|
||||||
\item \textbf{Routing}:
|
\item \textbf{Routing}(K3 eq.13):
|
||||||
\[
|
\[
|
||||||
\mathrm{logits} = W_r \cdot x \qquad \shape{B, T, n_{\mathrm{routed}}}
|
s = \sigma(W_r \cdot x),\quad
|
||||||
\]
|
T = \mathrm{TopK}(s+b, k),\quad
|
||||||
\[
|
p_i = \frac{s_i}{\sum_{j\in T} s_j}
|
||||||
\mathrm{ids}, \mathrm{probs} = \mathrm{TopK}(\mathrm{logits}, k)
|
|
||||||
\qquad \mathrm{ids}: \shape{B, T, k}, \;\; \mathrm{probs}: \shape{B, T, k}
|
|
||||||
\]
|
\]
|
||||||
|
|
||||||
\item \textbf{Routed 专家}(在 latent 空间 $\ell$ 上):
|
\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}
|
||||||
\]
|
\]
|
||||||
|
|
||||||
|
数学上是"对选中的 $k$ 个专家加权求和",但\textbf{实现上不是}用
|
||||||
|
\texttt{stack([e(z) for e in experts])} 把所有专家都算一遍——
|
||||||
|
而是每个 token 只被送进它选中的 $k$ 个专家(permute-dispatch + padded bmm)。
|
||||||
|
|
||||||
\item \textbf{Shared 专家}(全宽 $d$):
|
\item \textbf{Shared 专家}(全宽 $d$):
|
||||||
\[
|
\[
|
||||||
s = \sum_j E_j^{\mathrm{sh}}(x) \qquad \shape{B, T, d}
|
s_{\mathrm{sh}} = \sum_j E_j^{\mathrm{sh}}(x) \qquad \shape{B, T, d}
|
||||||
\]
|
\]
|
||||||
|
|
||||||
\item \textbf{合并}:
|
\item \textbf{合并}:
|
||||||
\[
|
\[
|
||||||
y = s + W_\uparrow \cdot \mathrm{RMSNorm}(u) \qquad \shape{B, T, d}
|
y = s_{\mathrm{sh}} + W_\uparrow \cdot \mathrm{RMSNorm}(u) \qquad \shape{B, T, d}
|
||||||
\]
|
\]
|
||||||
\end{enumerate}
|
\end{enumerate}
|
||||||
|
|
||||||
@@ -99,38 +102,133 @@ Routed 专家只在 $\ell = d/2$ 的 latent 空间操作,
|
|||||||
Shared 专家保持全宽 $d$,提供基础表达能力。
|
Shared 专家保持全宽 $d$,提供基础表达能力。
|
||||||
\end{importantbox}
|
\end{importantbox}
|
||||||
|
|
||||||
\subsection{代码对照}
|
\subsection{代码对照:路由}
|
||||||
|
|
||||||
\begin{codemathtop}{layers/latent\_moe.py — LatentMoE.forward}
|
\begin{codemathtop}{layers/latent\_moe.py — \_route(K3 eq.13)}
|
||||||
\begin{lstlisting}
|
\begin{lstlisting}
|
||||||
def forward(self, x): # [B, T, d]
|
def _route(self, logits): # logits: [B, T, n_routed]
|
||||||
z = self.down(x) # [B, T, ell]
|
scores = sigmoid(logits) # s = sigma(W_r x)
|
||||||
|
ids = topk(scores + self.expert_bias, k).indices # T = TopK(s+b)
|
||||||
logits = self.router(x) # [B, T, n_routed]
|
selected = scores.gather(-1, ids)
|
||||||
topk = torch.topk(logits, self.top_k, dim=-1)
|
probs = selected / selected.sum(-1).clamp_min(1e-9) # p_i = s_i / sum_{j in T} s_j
|
||||||
ids = topk.indices # [B, T, k]
|
return ids, probs
|
||||||
probs = F.softmax(topk.values, dim=-1) # [B, T, k]
|
|
||||||
|
|
||||||
# All expert outputs (vectorized)
|
|
||||||
all_out = stack([e(z) for e in self.experts]) # [R, B, T, ell]
|
|
||||||
# Gather top-k and weighted sum
|
|
||||||
u = zeros(B, T, ell)
|
|
||||||
for i in range(self.top_k):
|
|
||||||
idx = ids[:,:,i].reshape(B*T)
|
|
||||||
sel = all_out[arange, idx]
|
|
||||||
u += probs[:,:,i:i+1] * sel.reshape(B, T, ell)
|
|
||||||
|
|
||||||
shared_out = stack([e(x) for e in self.shared]).sum(0) # [B, T, d]
|
|
||||||
return shared_out + self.up(self.norm(u)) # [B, T, d]
|
|
||||||
\end{lstlisting}
|
\end{lstlisting}
|
||||||
\end{codemathtop}
|
\end{codemathtop}
|
||||||
|
|
||||||
|
注意三点:
|
||||||
|
|
||||||
|
\begin{enumerate}[leftmargin=2em]
|
||||||
|
\item \textbf{sigmoid 代替 softmax}:K3 的 router 对每个专家独立打分
|
||||||
|
$s_i = \sigma(w_i \cdot x)$,不再是 softmax 归一化。这样"专家之间"不互相竞争
|
||||||
|
归一化预算,便于用 bias 做负载调节。
|
||||||
|
\item \textbf{bias 只进选择、不进权重}:$\texttt{expert\_bias}$ 是
|
||||||
|
\texttt{register\_buffer(..., persistent=False)} 的非持久 buffer(初始化为 $0$,
|
||||||
|
不进 \texttt{state\_dict}),只参与 $\mathrm{TopK}(s+b)$ 的选择,
|
||||||
|
归一化 $p_i$ 仍用原始 $s_i$。
|
||||||
|
\item \textbf{L1 归一化}:$p_i = s_i / \sum_{j \in T} s_j$(在选中的 $k$ 个上做),
|
||||||
|
权重之和为 1,等价于"选中的 sigmoid 分数重新归一化"。
|
||||||
|
\end{enumerate}
|
||||||
|
|
||||||
|
\subsection{稀疏执行:permute-dispatch + padded bmm}\label{sec:sparse-dispatch}
|
||||||
|
|
||||||
|
\begin{codemathtop}{layers/latent\_moe.py — \_routed\_u}
|
||||||
|
\begin{lstlisting}
|
||||||
|
def _routed_u(self, z, ids, probs): # z: [B,T,ell] ids/probs: [B,T,k]
|
||||||
|
N = B * T; R, k = self.n_routed, self.top_k
|
||||||
|
tok = arange(N).unsqueeze(1).expand(N, k).reshape(-1) # 每个 token 重复 k 次
|
||||||
|
eid, pw = ids.reshape(-1), probs.reshape(-1)
|
||||||
|
|
||||||
|
order = eid.argsort(stable=True) # 按专家 id 排序 -> 同专家连续
|
||||||
|
tok, eid, pw = tok[order], eid[order], pw[order]
|
||||||
|
|
||||||
|
counts = bincount(eid, minlength=R) # 每个专家的 token 数
|
||||||
|
offsets = counts.cumsum(0) - counts
|
||||||
|
local_pos = arange(N*k) - offsets[eid] # 专家内局部位置
|
||||||
|
C = int(counts.max().item()) # capacity = 最大负载
|
||||||
|
|
||||||
|
gathered = z.reshape(N, ell)[tok]
|
||||||
|
padded = index_put(zeros(R, C, ell), (eid, local_pos), gathered) # [R, C, ell]
|
||||||
|
|
||||||
|
w_g = stack([e.w_g.weight for e in self.experts]) # [R, ff, ell]
|
||||||
|
w_u = stack([e.w_u.weight for e in self.experts])
|
||||||
|
w_o = stack([e.w_o.weight for e in self.experts]) # [R, ell, ff]
|
||||||
|
|
||||||
|
wg = bmm(padded, w_g.transpose(-1,-2)) # [R, C, ff] grouped GEMM
|
||||||
|
g = beta1 * tanh(wg / beta1) * sigmoid(wg)
|
||||||
|
wu = bmm(padded, w_u.transpose(-1,-2))
|
||||||
|
h = beta2 * tanh(wu / beta2)
|
||||||
|
out = bmm(g * h, w_o.transpose(-1,-2)) # [R, C, ell]
|
||||||
|
|
||||||
|
weighted = pw.unsqueeze(-1) * out[eid, local_pos]
|
||||||
|
return index_add(zeros(N, ell), 0, tok, weighted).view(B, T, ell) # scatter-add
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
三步走:\textbf{(1) permute-dispatch}(按专家排序 + pad 到 $[R, C, \ell]$)$\to$
|
||||||
|
\textbf{(2) padded bmm}(专家参数堆成 batch 维,三次 batched GEMM 一次算完 $R$ 个专家)$\to$
|
||||||
|
\textbf{(3) scatter-add}(\texttt{index\_add} 把加权输出按 \texttt{tok} 累加回 $u$)。
|
||||||
|
|
||||||
|
\begin{importantbox}{为什么不用 dense stack?}
|
||||||
|
朴素写法 \texttt{stack([e(z) for e in experts])} 会让每个专家都算全部 $B\cdot T$ 个 token,
|
||||||
|
FLOPs 是 $R \cdot N$,退化成 dense,失去 MoE 的加速。稀疏 dispatch 把每个 token 只送进
|
||||||
|
它选中的 $k$ 个专家,FLOPs 是 $R \cdot C$($C = \max_e \text{count}_e \approx k \cdot N / R$),
|
||||||
|
当 $k \ll R$ 时远小于 $R \cdot N$。padding 槽位不被 \texttt{index\_add} 收集,
|
||||||
|
贡献零梯度;空专家仍留在堆叠权重里(padded grouped GEMM)。
|
||||||
|
\end{importantbox}
|
||||||
|
|
||||||
\begin{warningbox}{为什么 router 用 $x$(全宽)而不是 $z$(latent)?}
|
\begin{warningbox}{为什么 router 用 $x$(全宽)而不是 $z$(latent)?}
|
||||||
路由需要看到 token 的完整表示才能做好选择。
|
路由需要看到 token 的完整表示才能做好选择。
|
||||||
如果用 $z$ 路由,压缩过程可能丢失路由需要的信息。
|
如果用 $z$ 路由,压缩过程可能丢失路由需要的信息。
|
||||||
K3 论文里也是用全宽 $x$ 做 Top-k,然后在 $\ell$ 空间计算。
|
K3 论文里也是用全宽 $x$ 做 Top-k,然后在 $\ell$ 空间计算。
|
||||||
\end{warningbox}
|
\end{warningbox}
|
||||||
|
|
||||||
|
\subsection{负载均衡损失:Switch/GShard aux + z-loss}\label{sec:load-balancing}
|
||||||
|
|
||||||
|
Top-k 路由容易"塌缩"到少数专家(router 学出永远选某几个专家),导致负载不均衡、
|
||||||
|
专家利用率低。训练时加两个损失,由 train loop 加到 CE 上:
|
||||||
|
|
||||||
|
\[
|
||||||
|
f_e = \frac{\#\{\text{路由到 } e\}}{N \cdot k},\qquad
|
||||||
|
P_e = \frac{1}{N}\sum_{t} \sigma(W_r x_t)_e
|
||||||
|
\]
|
||||||
|
\[
|
||||||
|
\mathcal{L}_{\mathrm{aux}} = n_{\mathrm{routed}} \sum_e f_e \cdot P_e,\qquad
|
||||||
|
\mathcal{L}_{z} = \frac{1}{N}\sum_t \left(\log\!\sum_j e^{l_{tj}}\right)^2
|
||||||
|
\]
|
||||||
|
|
||||||
|
\begin{center}
|
||||||
|
\begin{tabular}{lp{11.5cm}}
|
||||||
|
\toprule
|
||||||
|
项 & 作用 \\
|
||||||
|
\midrule
|
||||||
|
$\mathcal{L}_{\mathrm{aux}}$ & Switch/GShard 风格:$f_e$ 是专家 $e$ 被路由到的 token 占比,$P_e$ 是它的平均 sigmoid 分数。塌缩时 $f$ 集中到单个专家、损失变大,逼着路由均匀化 \\
|
||||||
|
$\mathcal{L}_{z}$ & z-loss:对 raw logits 的 logsumexp 求平方,压住 logits 幅度、防 router 分数爆炸 \\
|
||||||
|
\bottomrule
|
||||||
|
\end{tabular}
|
||||||
|
\end{center}
|
||||||
|
|
||||||
|
\begin{codemathtop}{layers/latent\_moe.py — \_balancing\_losses}
|
||||||
|
\begin{lstlisting}
|
||||||
|
def _balancing_losses(self, logits, ids):
|
||||||
|
flat = logits.reshape(-1, self.n_routed).float()
|
||||||
|
scores = sigmoid(flat)
|
||||||
|
counts = bincount(ids.reshape(-1), minlength=self.n_routed).float()
|
||||||
|
frac = counts / counts.sum().clamp_min(1.0) # f_e
|
||||||
|
prob_mean = scores.mean(dim=0) # 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
|
||||||
|
return self.aux_loss_coef * aux, self.z_loss_coef * z_loss
|
||||||
|
\end{lstlisting}
|
||||||
|
\end{codemathtop}
|
||||||
|
|
||||||
|
两个损失只在 \texttt{self.training} 且系数非零时计算;系数默认
|
||||||
|
$\alpha_{\mathrm{aux}} = 10^{-2}$、$\alpha_z = 10^{-3}$(即 \texttt{K3Config} 的
|
||||||
|
\texttt{moe\_aux\_loss\_coef} / \texttt{moe\_z\_loss\_coef},可用 \texttt{--moe-aux-coef}
|
||||||
|
/ \texttt{--moe-z-coef} 覆盖)。train loop 里 \texttt{moe\_router\_losses(model)}
|
||||||
|
把所有 LatentMoE 层的损失求和,\texttt{loss = task + aux + z\_loss} 一起反传。
|
||||||
|
两个损失只依赖 router 输出 \texttt{logits} 与不可微的索引 \texttt{ids},
|
||||||
|
所以梯度只流回 router 的 $W_r$,不碰专家权重。
|
||||||
|
|
||||||
\subsection{形状总览}
|
\subsection{形状总览}
|
||||||
|
|
||||||
\begin{center}
|
\begin{center}
|
||||||
@@ -140,12 +238,17 @@ K3 论文里也是用全宽 $x$ 做 Top-k,然后在 $\ell$ 空间计算。
|
|||||||
\midrule
|
\midrule
|
||||||
$x$ & \shape{B, T, d} & 输入 \\
|
$x$ & \shape{B, T, d} & 输入 \\
|
||||||
$z$ & \shape{B, T, \ell} & latent($\ell = d/2$)\\
|
$z$ & \shape{B, T, \ell} & latent($\ell = d/2$)\\
|
||||||
logits & \shape{B, T, n_r} & router 输出 \\
|
logits & \shape{B, T, n_r} & router logits \\
|
||||||
|
$s$ & \shape{B, T, n_r} & $\sigma(\mathrm{logits})$ \\
|
||||||
|
$b$ & \shape{n_r} & expert bias(非持久,只进 TopK 不进 $p$)\\
|
||||||
ids & \shape{B, T, k} & Top-k 专家索引 \\
|
ids & \shape{B, T, k} & Top-k 专家索引 \\
|
||||||
probs & \shape{B, T, k} & Top-k softmax 权重 \\
|
probs & \shape{B, T, k} & K3 sigmoid-L1 权重 \\
|
||||||
\texttt{all\_out} & \shape{n_r, B, T, \ell} & 所有 routed 专家输出 \\
|
padded & \shape{n_r, C, \ell} & dispatch 后 pad 到容量 $C$ 的张量 \\
|
||||||
|
$C$ & 标量 & 最大专家负载(pad 宽度)\\
|
||||||
$u$ & \shape{B, T, \ell} & 加权求和后的 routed 输出 \\
|
$u$ & \shape{B, T, \ell} & 加权求和后的 routed 输出 \\
|
||||||
\texttt{shared\_out} & \shape{B, T, d} & shared 专家求和 \\
|
\texttt{shared\_out} & \shape{B, T, d} & shared 专家求和 \\
|
||||||
|
$f_e$, $P_e$ & 标量 & aux loss 的负载占比 / 平均分数 \\
|
||||||
|
$\mathcal{L}_{\mathrm{aux}}$, $\mathcal{L}_z$ & 标量 & 负载均衡 / z-loss \\
|
||||||
$y$ & \shape{B, T, d} & 最终输出 \\
|
$y$ & \shape{B, T, d} & 最终输出 \\
|
||||||
\bottomrule
|
\bottomrule
|
||||||
\end{tabular}
|
\end{tabular}
|
||||||
@@ -154,5 +257,7 @@ $y$ & \shape{B, T, d} & 最终输出 \\
|
|||||||
\subsection{本章小结}
|
\subsection{本章小结}
|
||||||
|
|
||||||
LatentMoE 把 routed 专家限制在 $\ell = d/2$ 的 latent 空间,省参数。
|
LatentMoE 把 routed 专家限制在 $\ell = d/2$ 的 latent 空间,省参数。
|
||||||
SiTU-GLU 给 gate 和 up 加 $\tanh$ 软上限($\beta_1=4, \beta_2=25$),
|
SiTU-GLU 给 gate 和 up 加 $\tanh$ 软上限($\beta_1=4, \beta_2=25$),防止低精度溢出。
|
||||||
防止低精度溢出。Shared 专家全宽,提供基础能力;routed 专家通过 Top-k 路由提供专业化能力。
|
路由走 K3 eq.13(sigmoid-TopK + L1 归一化),执行用稀疏 permute-dispatch + padded bmm
|
||||||
|
(每个 token 只算 $k$ 个专家,FLOPs $R \cdot C$ 而非 $R \cdot N$),训练加
|
||||||
|
Switch/GShard aux loss + router z-loss 防路由塌缩。
|
||||||
|
|||||||
@@ -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。
|
||||||
|
|||||||
@@ -30,6 +30,9 @@ $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_{\mathrm{moe}}$ & MoE 专家容量(pad 宽度) & 动态 \\
|
||||||
|
$\alpha_{\mathrm{aux}}$ & Switch/GShard aux 系数 & $10^{-2}$ \\
|
||||||
|
$\alpha_z$ & router z-loss 系数 & $10^{-3}$ \\
|
||||||
$N$ & AttnRes 原子层数 ($= 2L$) & 8 \\
|
$N$ & AttnRes 原子层数 ($= 2L$) & 8 \\
|
||||||
$S$ & AttnRes 块大小(原子层) & 2--24 \\
|
$S$ & AttnRes 块大小(原子层) & 2--24 \\
|
||||||
\bottomrule
|
\bottomrule
|
||||||
@@ -105,12 +108,19 @@ gate & \shape{B, T, H \cdot d_v} & $\sigma(W_g x)$ \\
|
|||||||
\midrule
|
\midrule
|
||||||
$x$ & \shape{B, T, D} & 输入 \\
|
$x$ & \shape{B, T, D} & 输入 \\
|
||||||
$z$ & \shape{B, T, \ell} & latent ($\ell = D/2$) \\
|
$z$ & \shape{B, T, \ell} & latent ($\ell = D/2$) \\
|
||||||
logits & \shape{B, T, n_r} & router logits \\
|
logits & \shape{B, T, n_r} & router logits $W_r x$ \\
|
||||||
|
$s$ & \shape{B, T, n_r} & sigmoid 分数 $\sigma(\mathrm{logits})$ \\
|
||||||
|
$b$ & \shape{n_r} & expert bias(非持久,只进 TopK) \\
|
||||||
ids & \shape{B, T, k} & Top-$k$ 专家索引 \\
|
ids & \shape{B, T, k} & Top-$k$ 专家索引 \\
|
||||||
probs & \shape{B, T, k} & softmax 权重 \\
|
$p_i$ & \shape{B, T, k} & sigmoid-L1 权重 $s_i/\sum_{j\in T}s_j$ \\
|
||||||
|
padded & \shape{n_r, C_{\mathrm{moe}}, \ell} & dispatch 后 pad 到容量 $C_{\mathrm{moe}}$ \\
|
||||||
|
$C_{\mathrm{moe}}$ & 标量 & 最大专家负载(pad 宽度) \\
|
||||||
$u$ & \shape{B, T, \ell} & routed 加权输出 \\
|
$u$ & \shape{B, T, \ell} & routed 加权输出 \\
|
||||||
$s$ & \shape{B, T, D} & shared 专家求和 \\
|
$s_{\mathrm{sh}}$ & \shape{B, T, D} & shared 专家求和 \\
|
||||||
$y$ & \shape{B, T, D} & $s + W_\uparrow \mathrm{RMSNorm}(u)$ \\
|
$y$ & \shape{B, T, D} & $s_{\mathrm{sh}} + W_\uparrow \mathrm{RMSNorm}(u)$ \\
|
||||||
|
$f_e, P_e$ & 标量 & aux loss 负载占比 / 平均分数 \\
|
||||||
|
$\mathcal{L}_{\mathrm{aux}}, \mathcal{L}_z$ & 标量 & 负载均衡 / z-loss \\
|
||||||
|
$\alpha_{\mathrm{aux}}, \alpha_z$ & 标量 & 对应系数($10^{-2}$ / $10^{-3}$) \\
|
||||||
\bottomrule
|
\bottomrule
|
||||||
\end{tabular}
|
\end{tabular}
|
||||||
\end{center}
|
\end{center}
|
||||||
@@ -154,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}
|
||||||
@@ -167,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 合并
|
||||||
|
|||||||
@@ -59,6 +59,46 @@ def test_gradient_checkpointing_matches_eager_grad():
|
|||||||
torch.testing.assert_close(p1.grad, p2.grad, atol=1e-4, rtol=1e-4)
|
torch.testing.assert_close(p1.grad, p2.grad, atol=1e-4, rtol=1e-4)
|
||||||
|
|
||||||
|
|
||||||
|
def test_attnres_block_checkpoint_matches_eager_grad():
|
||||||
|
torch.manual_seed(8)
|
||||||
|
cfg = K3Config(
|
||||||
|
hidden_size=32,
|
||||||
|
num_hidden_layers=4,
|
||||||
|
num_heads=4,
|
||||||
|
head_dim=8,
|
||||||
|
chunk_size=4,
|
||||||
|
vocab_size=32,
|
||||||
|
moe_latent_size=16,
|
||||||
|
moe_d_ff=16,
|
||||||
|
n_routed=4,
|
||||||
|
kv_lora_rank=16,
|
||||||
|
q_lora_rank=32,
|
||||||
|
qk_nope_head_dim=8,
|
||||||
|
v_head_dim=8,
|
||||||
|
attnres="block",
|
||||||
|
attnres_block_size=1,
|
||||||
|
moe_aux_loss_coef=0.0,
|
||||||
|
moe_z_loss_coef=0.0,
|
||||||
|
)
|
||||||
|
tokens = torch.randint(0, cfg.vocab_size, (2, 8))
|
||||||
|
m1 = CausalLM(cfg)
|
||||||
|
m2 = CausalLM(cfg)
|
||||||
|
m2.load_state_dict(m1.state_dict())
|
||||||
|
m2.gradient_checkpointing = True
|
||||||
|
m1.train()
|
||||||
|
m2.train()
|
||||||
|
l1 = m1(tokens, labels=tokens)
|
||||||
|
l2 = m2(tokens, labels=tokens)
|
||||||
|
torch.testing.assert_close(l1, l2, atol=1e-5, rtol=1e-5)
|
||||||
|
l1.backward()
|
||||||
|
l2.backward()
|
||||||
|
for p1, p2 in zip(m1.parameters(), m2.parameters()):
|
||||||
|
if p1.grad is None:
|
||||||
|
assert p2.grad is None
|
||||||
|
continue
|
||||||
|
torch.testing.assert_close(p1.grad, p2.grad, atol=1e-4, rtol=1e-4)
|
||||||
|
|
||||||
|
|
||||||
def test_0_5b_preset_enables_checkpointing():
|
def test_0_5b_preset_enables_checkpointing():
|
||||||
assert K3Config.preset("0.5b").gradient_checkpointing is True
|
assert K3Config.preset("0.5b").gradient_checkpointing is True
|
||||||
assert K3Config.preset("toy").gradient_checkpointing is False
|
assert K3Config.preset("toy").gradient_checkpointing is False
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ import torch.nn.functional as F
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from kda.layers.kda_attn import KDAAttention
|
from kda.layers.kda_attn import KDAAttention
|
||||||
from kda.layers.latent_moe import LatentMoE
|
from kda.layers.latent_moe import LatentMoE, moe_router_losses
|
||||||
from kda.layers.mla import GatedMLA
|
from kda.layers.mla import GatedMLA
|
||||||
from kda.models.causal_lm import CausalLM
|
from kda.models.causal_lm import CausalLM
|
||||||
from kda.models.k3_config import K3Config
|
from kda.models.k3_config import K3Config
|
||||||
@@ -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
|
||||||
@@ -88,30 +89,189 @@ def test_preset_0_5b_schedule():
|
|||||||
assert cfg.layer_specs()[3] == ("mla", "moe")
|
assert cfg.layer_specs()[3] == ("mla", "moe")
|
||||||
|
|
||||||
|
|
||||||
|
def _dense_moe_forward(moe: LatentMoE, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Dense path: run every routed expert, then gather K3 sigmoid-norm top-k."""
|
||||||
|
z = moe.down(x)
|
||||||
|
ids, probs = moe._route(moe.router(x))
|
||||||
|
all_out = torch.stack([expert(z) for expert in moe.experts])
|
||||||
|
B, T, _ = x.shape
|
||||||
|
all_out = all_out.permute(1, 2, 0, 3).reshape(B * T, moe.n_routed, moe.latent_size)
|
||||||
|
u = z.new_zeros(B, T, moe.latent_size)
|
||||||
|
for i in range(moe.top_k):
|
||||||
|
idx = ids[:, :, i].reshape(B * T)
|
||||||
|
sel = all_out[torch.arange(B * T, device=x.device), idx]
|
||||||
|
u = u + probs[:, :, i : i + 1] * sel.reshape(B, T, moe.latent_size)
|
||||||
|
shared = torch.stack([expert(x) for expert in moe.shared]).sum(0)
|
||||||
|
return shared + moe.up(moe.norm(u))
|
||||||
|
|
||||||
|
|
||||||
def test_moe_router_activates_topk_only():
|
def test_moe_router_activates_topk_only():
|
||||||
from kda.layers.latent_moe import LatentMoE
|
|
||||||
torch.manual_seed(3)
|
torch.manual_seed(3)
|
||||||
moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24)
|
moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24)
|
||||||
x = torch.randn(2, 6, 32)
|
x = torch.randn(2, 6, 32)
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
y = moe(x)
|
y = moe(x)
|
||||||
logits = moe.router(x)
|
scores = torch.sigmoid(moe.router(x))
|
||||||
topk = torch.topk(logits, moe.top_k, dim=-1)
|
ids, probs = moe._route(moe.router(x))
|
||||||
z = moe.down(x)
|
z = moe.down(x)
|
||||||
# 手算: 只有 top-k 专家输出被加权, 再经 shared + up(norm(u))
|
|
||||||
expected_u = torch.zeros(2, 6, moe.latent_size)
|
expected_u = torch.zeros(2, 6, moe.latent_size)
|
||||||
all_out = torch.stack([e(z) for e in moe.experts]) # [R,B,T,ℓ]
|
all_out = torch.stack([e(z) for e in moe.experts]) # [R,B,T,ℓ]
|
||||||
probs = F.softmax(topk.values, dim=-1)
|
|
||||||
for i in range(moe.top_k):
|
for i in range(moe.top_k):
|
||||||
idx = topk.indices[:, :, i]
|
idx = ids[:, :, i]
|
||||||
for b in range(2):
|
for b in range(2):
|
||||||
for t in range(6):
|
for t in range(6):
|
||||||
expected_u[b, t] += probs[b, t, i] * all_out[idx[b, t], b, t]
|
expected_u[b, t] += probs[b, t, i] * all_out[idx[b, t], b, t]
|
||||||
shared = torch.stack([e(x) for e in moe.shared]).sum(0)
|
shared = torch.stack([e(x) for e in moe.shared]).sum(0)
|
||||||
expected_y = shared + moe.up(moe.norm(expected_u))
|
expected_y = shared + moe.up(moe.norm(expected_u))
|
||||||
|
selected = scores.gather(-1, ids)
|
||||||
|
torch.testing.assert_close(
|
||||||
|
probs, selected / selected.sum(-1, keepdim=True).clamp_min(1e-9)
|
||||||
|
)
|
||||||
torch.testing.assert_close(y, expected_y, atol=1e-5, rtol=1e-5)
|
torch.testing.assert_close(y, expected_y, atol=1e-5, rtol=1e-5)
|
||||||
assert moe.last_route_ids is not None
|
assert moe.last_route_ids is not None
|
||||||
assert moe.last_route_ids.shape[-1] == moe.top_k
|
assert moe.last_route_ids.shape[-1] == moe.top_k
|
||||||
|
counts = torch.bincount(moe.last_route_ids.reshape(-1), minlength=moe.n_routed)
|
||||||
|
assert moe.last_capacity == int(counts.max())
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_sparse_matches_dense_fwd_bwd():
|
||||||
|
torch.manual_seed(3)
|
||||||
|
moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24)
|
||||||
|
x = torch.randn(2, 6, 32)
|
||||||
|
y_sparse = moe(x)
|
||||||
|
y_dense = _dense_moe_forward(moe, x)
|
||||||
|
torch.testing.assert_close(y_sparse, y_dense, atol=1e-5, rtol=1e-5)
|
||||||
|
|
||||||
|
moe.zero_grad(set_to_none=True)
|
||||||
|
xs = x.detach().requires_grad_(True)
|
||||||
|
moe(xs).square().mean().backward()
|
||||||
|
grads_s = {
|
||||||
|
name: param.grad.detach().clone()
|
||||||
|
for name, param in moe.named_parameters()
|
||||||
|
if param.grad is not None
|
||||||
|
}
|
||||||
|
dx_s = xs.grad.detach().clone()
|
||||||
|
|
||||||
|
moe.zero_grad(set_to_none=True)
|
||||||
|
xd = x.detach().requires_grad_(True)
|
||||||
|
_dense_moe_forward(moe, xd).square().mean().backward()
|
||||||
|
grads_d = {
|
||||||
|
name: param.grad.detach().clone()
|
||||||
|
for name, param in moe.named_parameters()
|
||||||
|
if param.grad is not None
|
||||||
|
}
|
||||||
|
dx_d = xd.grad.detach().clone()
|
||||||
|
|
||||||
|
torch.testing.assert_close(dx_s, dx_d, atol=1e-5, rtol=1e-5)
|
||||||
|
assert grads_s.keys() == grads_d.keys()
|
||||||
|
for name in grads_s:
|
||||||
|
torch.testing.assert_close(grads_s[name], grads_d[name], atol=1e-5, rtol=1e-5)
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_unselected_experts_have_zero_grad():
|
||||||
|
torch.manual_seed(0)
|
||||||
|
moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24)
|
||||||
|
with torch.no_grad():
|
||||||
|
moe.router.weight.zero_()
|
||||||
|
moe.router.weight[0] = 1.0
|
||||||
|
moe.router.weight[1] = 0.5
|
||||||
|
x = torch.ones(2, 4, 32)
|
||||||
|
moe(x).square().mean().backward()
|
||||||
|
selected = set(moe.last_route_ids.reshape(-1).tolist())
|
||||||
|
assert selected == {0, 1}
|
||||||
|
for idx, expert in enumerate(moe.experts):
|
||||||
|
for param in (expert.w_g.weight, expert.w_u.weight, expert.w_o.weight):
|
||||||
|
assert param.grad is not None
|
||||||
|
if idx in selected:
|
||||||
|
assert param.grad.abs().sum() > 0
|
||||||
|
else:
|
||||||
|
assert torch.equal(param.grad, torch.zeros_like(param.grad))
|
||||||
|
|
||||||
|
|
||||||
|
def test_routed_u_bf16_index_add_matches_z_dtype():
|
||||||
|
"""Python float * bf16 promotes to fp32; index_add must still land in z.dtype."""
|
||||||
|
torch.manual_seed(0)
|
||||||
|
moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24)
|
||||||
|
z = torch.randn(2, 6, 16, dtype=torch.bfloat16)
|
||||||
|
logits = torch.randn(2, 6, 8, dtype=torch.bfloat16)
|
||||||
|
ids, probs = moe._route(logits)
|
||||||
|
u = moe._routed_u(z, ids, probs)
|
||||||
|
assert u.dtype == torch.bfloat16
|
||||||
|
u.float().square().mean().backward()
|
||||||
|
assert any(p.grad is not None and p.grad.abs().sum() > 0 for p in moe.experts[0].parameters())
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_aux_loss_penalizes_collapse():
|
||||||
|
torch.manual_seed(0)
|
||||||
|
moe = LatentMoE(
|
||||||
|
hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24,
|
||||||
|
aux_loss_coef=1.0, z_loss_coef=1.0,
|
||||||
|
)
|
||||||
|
moe(torch.randn(4, 16, 32))
|
||||||
|
spread = float(moe.last_aux_loss.detach())
|
||||||
|
with torch.no_grad():
|
||||||
|
moe.router.weight.zero_()
|
||||||
|
moe.router.weight[0] = 1.0
|
||||||
|
moe.router.weight[1] = 0.5
|
||||||
|
moe(torch.ones(4, 16, 32))
|
||||||
|
collapsed = float(moe.last_aux_loss.detach())
|
||||||
|
assert collapsed > spread
|
||||||
|
assert collapsed > 1.2
|
||||||
|
assert float(moe.last_z_loss.detach()) > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_aux_loss_updates_router_only():
|
||||||
|
torch.manual_seed(1)
|
||||||
|
moe = LatentMoE(
|
||||||
|
hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24,
|
||||||
|
aux_loss_coef=1.0, z_loss_coef=1.0,
|
||||||
|
)
|
||||||
|
moe(torch.randn(2, 8, 32))
|
||||||
|
aux, z_loss = moe_router_losses(moe)
|
||||||
|
(aux + z_loss).backward()
|
||||||
|
assert moe.router.weight.grad is not None
|
||||||
|
assert moe.router.weight.grad.abs().sum() > 0
|
||||||
|
for expert in moe.experts:
|
||||||
|
assert expert.w_g.weight.grad is None
|
||||||
|
assert expert.w_u.weight.grad is None
|
||||||
|
assert expert.w_o.weight.grad is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_router_losses_sums_layers():
|
||||||
|
cfg = K3Config(
|
||||||
|
hidden_size=32, num_hidden_layers=2, num_heads=4, head_dim=8,
|
||||||
|
chunk_size=4, vocab_size=64, moe_latent_size=16, moe_d_ff=16,
|
||||||
|
n_routed=4, kv_lora_rank=16, q_lora_rank=32, qk_nope_head_dim=8, v_head_dim=8,
|
||||||
|
moe_aux_loss_coef=1.0, moe_z_loss_coef=1.0,
|
||||||
|
)
|
||||||
|
model = CausalLM(cfg)
|
||||||
|
model(torch.randint(0, cfg.vocab_size, (2, 8)))
|
||||||
|
aux, z_loss = moe_router_losses(model)
|
||||||
|
layers = [module for module in model.modules() if isinstance(module, LatentMoE)]
|
||||||
|
assert len(layers) == 2
|
||||||
|
torch.testing.assert_close(aux, layers[0].last_aux_loss + layers[1].last_aux_loss)
|
||||||
|
torch.testing.assert_close(z_loss, layers[0].last_z_loss + layers[1].last_z_loss)
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_aux_backward_with_checkpoint():
|
||||||
|
cfg = K3Config(
|
||||||
|
hidden_size=32, num_hidden_layers=2, num_heads=4, head_dim=8,
|
||||||
|
chunk_size=4, vocab_size=64, moe_latent_size=16, moe_d_ff=16,
|
||||||
|
n_routed=4, kv_lora_rank=16, q_lora_rank=32, qk_nope_head_dim=8, v_head_dim=8,
|
||||||
|
gradient_checkpointing=True, moe_aux_loss_coef=1.0, moe_z_loss_coef=1.0,
|
||||||
|
)
|
||||||
|
model = CausalLM(cfg)
|
||||||
|
model.train()
|
||||||
|
tokens = torch.randint(0, cfg.vocab_size, (2, 8))
|
||||||
|
task = model(tokens, labels=tokens)
|
||||||
|
aux, z_loss = moe_router_losses(model)
|
||||||
|
(task + aux + z_loss).backward()
|
||||||
|
router_grad = sum(
|
||||||
|
param.grad.abs().sum().item()
|
||||||
|
for name, param in model.named_parameters()
|
||||||
|
if "router" in name and param.grad is not None
|
||||||
|
)
|
||||||
|
assert router_grad > 0
|
||||||
|
|
||||||
|
|
||||||
def test_k3_causal_future_does_not_change_past_logits():
|
def test_k3_causal_future_does_not_change_past_logits():
|
||||||
|
|||||||
@@ -9,13 +9,21 @@ def test_warmup_then_cosine_floor():
|
|||||||
assert abs(end - 0.1) < 1e-6
|
assert abs(end - 0.1) < 1e-6
|
||||||
|
|
||||||
|
|
||||||
def test_horizon_prefers_the_earlier_stop():
|
def test_horizon_max_tokens_overrides_micro_cap():
|
||||||
# 8.2M tokens @ batch 2 seq 2048 acc 8 -> 250 opt
|
# 8.2M tokens @ batch 2 seq 2048 acc 8 -> 250 opt even if --steps is larger
|
||||||
opt_from_tokens = total_opt_steps(
|
opt_from_tokens = total_opt_steps(
|
||||||
max_tokens=8_192_000, max_micro=10_000, batch=2, seq_len=2048, grad_acc=8
|
max_tokens=8_192_000, max_micro=10_000, batch=2, seq_len=2048, grad_acc=8
|
||||||
)
|
)
|
||||||
assert opt_from_tokens == 250
|
assert opt_from_tokens == 250
|
||||||
|
# 1B-token run must not inherit the default --steps 2000 cap (250 opt)
|
||||||
|
opt_1b = total_opt_steps(
|
||||||
|
max_tokens=10**9, max_micro=2000, batch=2, seq_len=2048, grad_acc=8
|
||||||
|
)
|
||||||
|
assert opt_1b == 30518
|
||||||
|
|
||||||
|
|
||||||
|
def test_horizon_micro_when_tokens_unset():
|
||||||
opt_from_micro = total_opt_steps(
|
opt_from_micro = total_opt_steps(
|
||||||
max_tokens=10**12, max_micro=2000, batch=2, seq_len=2048, grad_acc=8
|
max_tokens=None, max_micro=2000, batch=2, seq_len=2048, grad_acc=8
|
||||||
)
|
)
|
||||||
assert opt_from_micro == 250
|
assert opt_from_micro == 250
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -1,4 +1,11 @@
|
|||||||
from kda.training.data import IGNORE_INDEX, collate_sft, encode_sft_row, load_sft_rows
|
from kda.training.data import (
|
||||||
|
IGNORE_INDEX,
|
||||||
|
collate_sft,
|
||||||
|
encode_sft_row,
|
||||||
|
fetch_opus100_enzh,
|
||||||
|
load_sft_rows,
|
||||||
|
resolve_sft_rows,
|
||||||
|
)
|
||||||
from kda.training.prompts import instruction_prompt
|
from kda.training.prompts import instruction_prompt
|
||||||
|
|
||||||
|
|
||||||
@@ -42,6 +49,36 @@ def test_collate_and_jsonl(tmp_path):
|
|||||||
assert (y == IGNORE_INDEX).any()
|
assert (y == IGNORE_INDEX).any()
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_sft_rows_reads_local_jsonl(tmp_path):
|
||||||
|
path = tmp_path / "bitext.jsonl"
|
||||||
|
path.write_text(
|
||||||
|
'{"src": "你好", "tgt": "Hello", "target_lang": "en"}\n',
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
rows = resolve_sft_rows(str(path))
|
||||||
|
assert rows == [{"src": "你好", "tgt": "Hello", "target_lang": "en"}]
|
||||||
|
|
||||||
|
|
||||||
|
def test_fetch_opus_skips_eval_sentences(tmp_path, monkeypatch):
|
||||||
|
eval_dir = tmp_path / "eval"
|
||||||
|
eval_dir.mkdir()
|
||||||
|
(eval_dir / "zh2en.src.txt").write_text("禁止句\n", encoding="utf-8")
|
||||||
|
|
||||||
|
class _DS:
|
||||||
|
def __iter__(self):
|
||||||
|
yield {"translation": {"en": "Hello", "zh": "你好"}}
|
||||||
|
yield {"translation": {"en": "skip", "zh": "禁止句"}}
|
||||||
|
yield {"translation": {"en": "Thanks", "zh": "谢谢"}}
|
||||||
|
|
||||||
|
fake = type("datasets", (), {"load_dataset": staticmethod(lambda *a, **k: _DS())})
|
||||||
|
monkeypatch.setitem(__import__("sys").modules, "datasets", fake)
|
||||||
|
rows = fetch_opus100_enzh(10, cache_dir=tmp_path / "sft", eval_dir=eval_dir)
|
||||||
|
srcs = {r["src"] for r in rows}
|
||||||
|
assert "禁止句" not in srcs
|
||||||
|
assert "你好" in srcs and "Hello" in srcs
|
||||||
|
assert "谢谢" in srcs and "Thanks" in srcs
|
||||||
|
|
||||||
|
|
||||||
def test_toy_sft_file_parses():
|
def test_toy_sft_file_parses():
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
+199
-42
@@ -17,11 +17,12 @@ from dataclasses import asdict
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from kda.layers.latent_moe import moe_route_frac
|
from kda.layers.latent_moe import LatentMoE, moe_route_frac, moe_router_losses
|
||||||
from kda.models.causal_lm import CausalLM
|
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 = {
|
||||||
@@ -35,9 +36,12 @@ _TOY_TRAIN = {
|
|||||||
"warmup": 50,
|
"warmup": 50,
|
||||||
"grad_acc": 1,
|
"grad_acc": 1,
|
||||||
"eval_every": 100,
|
"eval_every": 100,
|
||||||
|
"log_every": 10,
|
||||||
|
"ckpt_every": 100,
|
||||||
|
"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,
|
||||||
@@ -46,7 +50,10 @@ _B500M_TRAIN = {
|
|||||||
"lr": 3e-4,
|
"lr": 3e-4,
|
||||||
"warmup": 64,
|
"warmup": 64,
|
||||||
"grad_acc": 8,
|
"grad_acc": 8,
|
||||||
"eval_every": 100,
|
"eval_every": 500,
|
||||||
|
"log_every": 20,
|
||||||
|
"ckpt_every": 1000,
|
||||||
|
"gen_every": 2000,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -60,11 +67,14 @@ def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None:
|
|||||||
group["lr"] = lr
|
group["lr"] = lr
|
||||||
|
|
||||||
|
|
||||||
def _init_swanlab(cfg: K3Config, args: argparse.Namespace):
|
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:
|
||||||
@@ -72,9 +82,8 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace):
|
|||||||
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.
|
run_id = resume_id or os.environ.get("SWANLAB_RUN_ID")
|
||||||
project = os.environ.pop("SWANLAB_PROJECT", None) or "kda"
|
init_kw = dict(
|
||||||
return swanlab.init(
|
|
||||||
project=project,
|
project=project,
|
||||||
name=f"{args.preset}-{cfg.attnres}",
|
name=f"{args.preset}-{cfg.attnres}",
|
||||||
config={
|
config={
|
||||||
@@ -91,8 +100,24 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace):
|
|||||||
"langs": args.langs,
|
"langs": args.langs,
|
||||||
"kda_backend": cfg.kda_backend,
|
"kda_backend": cfg.kda_backend,
|
||||||
"gradient_checkpointing": cfg.gradient_checkpointing,
|
"gradient_checkpointing": cfg.gradient_checkpointing,
|
||||||
|
"moe_aux_loss_coef": cfg.moe_aux_loss_coef,
|
||||||
|
"moe_z_loss_coef": cfg.moe_z_loss_coef,
|
||||||
|
"eval_every": args.eval_every,
|
||||||
|
"log_every": args.log_every,
|
||||||
|
"ckpt_every": args.ckpt_every,
|
||||||
|
"gen_every": args.gen_every,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
if run_id:
|
||||||
|
init_kw["id"] = run_id
|
||||||
|
init_kw["resume"] = True
|
||||||
|
print(f"swanlab resume id={run_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
|
||||||
@@ -121,6 +146,10 @@ def _payload(
|
|||||||
"tokens": tokens,
|
"tokens": tokens,
|
||||||
"chunk_index": chunk_index,
|
"chunk_index": chunk_index,
|
||||||
"best_heldout": best_heldout,
|
"best_heldout": best_heldout,
|
||||||
|
"batch": args.batch,
|
||||||
|
"seq_len": args.seq_len,
|
||||||
|
"grad_acc": args.grad_acc,
|
||||||
|
"swanlab_id": getattr(args, "swanlab_id", None),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -148,6 +177,27 @@ def _heldout_loss(
|
|||||||
return sum(losses) / max(len(losses), 1)
|
return sum(losses) / max(len(losses), 1)
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_moe_coefs(model, cfg: K3Config) -> None:
|
||||||
|
for module in model.modules():
|
||||||
|
if isinstance(module, LatentMoE):
|
||||||
|
module.aux_loss_coef = cfg.moe_aux_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:
|
||||||
@@ -192,6 +242,24 @@ def main() -> None:
|
|||||||
)
|
)
|
||||||
p.add_argument("--grad-acc", type=int, default=train_defaults["grad_acc"])
|
p.add_argument("--grad-acc", type=int, default=train_defaults["grad_acc"])
|
||||||
p.add_argument("--eval-every", type=int, default=train_defaults["eval_every"])
|
p.add_argument("--eval-every", type=int, default=train_defaults["eval_every"])
|
||||||
|
p.add_argument(
|
||||||
|
"--log-every",
|
||||||
|
type=int,
|
||||||
|
default=train_defaults["log_every"],
|
||||||
|
help="swanlab scalar period in micro-steps",
|
||||||
|
)
|
||||||
|
p.add_argument(
|
||||||
|
"--ckpt-every",
|
||||||
|
type=int,
|
||||||
|
default=train_defaults["ckpt_every"],
|
||||||
|
help="write _last/_best this many micro-steps (1B default 1000)",
|
||||||
|
)
|
||||||
|
p.add_argument(
|
||||||
|
"--gen-every",
|
||||||
|
type=int,
|
||||||
|
default=train_defaults["gen_every"],
|
||||||
|
help="sample prefixes this often; 0 disables",
|
||||||
|
)
|
||||||
p.add_argument(
|
p.add_argument(
|
||||||
"--langs",
|
"--langs",
|
||||||
default="zh,en",
|
default="zh,en",
|
||||||
@@ -199,13 +267,24 @@ def main() -> None:
|
|||||||
)
|
)
|
||||||
p.add_argument("--heldout-frac", type=float, default=0.01)
|
p.add_argument("--heldout-frac", type=float, default=0.01)
|
||||||
p.add_argument("--resume", default=None, help="checkpoint to continue from")
|
p.add_argument("--resume", default=None, help="checkpoint to continue from")
|
||||||
|
p.add_argument(
|
||||||
|
"--swanlab-id",
|
||||||
|
default=None,
|
||||||
|
help="resume this SwanLab run (URL /runs/<id>); default: id stored in ckpt",
|
||||||
|
)
|
||||||
|
p.add_argument(
|
||||||
|
"--swanlab-new",
|
||||||
|
action="store_true",
|
||||||
|
help="start a new SwanLab run even when --resume",
|
||||||
|
)
|
||||||
p.add_argument("--gen-prefix", action="append", default=None)
|
p.add_argument("--gen-prefix", action="append", default=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",
|
||||||
@@ -225,7 +304,20 @@ def main() -> None:
|
|||||||
dest="grad_checkpoint",
|
dest="grad_checkpoint",
|
||||||
action="store_false",
|
action="store_false",
|
||||||
)
|
)
|
||||||
|
p.add_argument(
|
||||||
|
"--moe-aux-coef",
|
||||||
|
type=float,
|
||||||
|
default=None,
|
||||||
|
help="Switch/GShard aux loss weight (default 0.01; 0 disables)",
|
||||||
|
)
|
||||||
|
p.add_argument(
|
||||||
|
"--moe-z-coef",
|
||||||
|
type=float,
|
||||||
|
default=None,
|
||||||
|
help="router z-loss weight (default 0.001; 0 disables)",
|
||||||
|
)
|
||||||
args = p.parse_args()
|
args = p.parse_args()
|
||||||
|
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
|
||||||
if args.gen_prefix is None:
|
if args.gen_prefix is None:
|
||||||
args.gen_prefix = ["人工智能的发展", "The history of computing"]
|
args.gen_prefix = ["人工智能的发展", "The history of computing"]
|
||||||
|
|
||||||
@@ -243,15 +335,16 @@ def main() -> None:
|
|||||||
raise SystemExit(
|
raise SystemExit(
|
||||||
"KDA training needs bf16; this GPU does not support it (avoid V100 fp16)"
|
"KDA training needs bf16; this GPU does not support it (avoid V100 fp16)"
|
||||||
)
|
)
|
||||||
|
if device == "cuda":
|
||||||
|
torch.backends.cuda.matmul.allow_tf32 = True
|
||||||
|
torch.backends.cudnn.allow_tf32 = True
|
||||||
|
torch.set_float32_matmul_precision("high")
|
||||||
|
|
||||||
print(f"loading tokenizer {args.tokenizer} ...")
|
print(f"loading tokenizer {args.tokenizer} ...")
|
||||||
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
|
|
||||||
|
|
||||||
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)
|
||||||
@@ -279,24 +372,32 @@ 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
|
|
||||||
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))
|
||||||
chunk_index = int(payload.get("chunk_index", 0))
|
chunk_index = int(payload.get("chunk_index", 0))
|
||||||
best_heldout = float(payload.get("best_heldout", best_heldout))
|
best_heldout = float(payload.get("best_heldout", best_heldout))
|
||||||
|
if not args.swanlab_new and not args.swanlab_id:
|
||||||
|
args.swanlab_id = payload.get("swanlab_id") or args.swanlab_id
|
||||||
else:
|
else:
|
||||||
model = CausalLM(cfg).to(device)
|
model = CausalLM(cfg).to(device)
|
||||||
|
|
||||||
tracker = _init_swanlab(cfg, args)
|
_apply_moe_coefs(model, cfg)
|
||||||
|
tracker = _init_swanlab(
|
||||||
|
cfg,
|
||||||
|
args,
|
||||||
|
resume_id=None if args.swanlab_new else args.swanlab_id,
|
||||||
|
)
|
||||||
n = sum(p.numel() for p in model.parameters())
|
n = sum(p.numel() for p in model.parameters())
|
||||||
print(
|
print(
|
||||||
f"preset={args.preset} model={n:,} params ({n / 1e6:.1f}M) on {device} "
|
f"preset={args.preset} model={n:,} params ({n / 1e6:.1f}M) on {device} "
|
||||||
@@ -305,6 +406,7 @@ def main() -> None:
|
|||||||
print(
|
print(
|
||||||
f"vocab={cfg.vocab_size} tied={cfg.tie_word_embeddings} "
|
f"vocab={cfg.vocab_size} tied={cfg.tie_word_embeddings} "
|
||||||
f"layers={cfg.layer_types()} attnres={cfg.attnres} langs={langs} "
|
f"layers={cfg.layer_types()} attnres={cfg.attnres} langs={langs} "
|
||||||
|
f"moe_aux={cfg.moe_aux_loss_coef:g} moe_z={cfg.moe_z_loss_coef:g}"
|
||||||
)
|
)
|
||||||
if args.max_tokens is None:
|
if args.max_tokens is None:
|
||||||
print(
|
print(
|
||||||
@@ -312,7 +414,11 @@ def main() -> None:
|
|||||||
f"({args.steps * tpm:,} tokens); pass --max-tokens for a real run"
|
f"({args.steps * tpm:,} tokens); pass --max-tokens for a real run"
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
print(f"token budget: {args.max_tokens:,} cosine horizon {horizon} opt steps")
|
print(
|
||||||
|
f"token budget: {args.max_tokens:,} cosine horizon {horizon} opt steps "
|
||||||
|
f"log/{args.log_every} eval/{args.eval_every} ckpt/{args.ckpt_every} "
|
||||||
|
f"gen/{args.gen_every}"
|
||||||
|
)
|
||||||
|
|
||||||
train_chunks, held_chunks, n_ids = load_pretrain_chunks(
|
train_chunks, held_chunks, n_ids = load_pretrain_chunks(
|
||||||
tok,
|
tok,
|
||||||
@@ -325,9 +431,22 @@ def main() -> None:
|
|||||||
print(
|
print(
|
||||||
f"packed tokens {n_ids:,} -> {train_chunks.size(0)} train / "
|
f"packed tokens {n_ids:,} -> {train_chunks.size(0)} train / "
|
||||||
f"{held_chunks.size(0)} held-out chunks of [{args.batch}, {args.seq_len}] "
|
f"{held_chunks.size(0)} held-out chunks of [{args.batch}, {args.seq_len}] "
|
||||||
|
f"{tpm} tok/micro"
|
||||||
)
|
)
|
||||||
if train_chunks.size(0) == 0:
|
if train_chunks.size(0) == 0:
|
||||||
raise SystemExit("no training chunks; raise --limit or lower --batch/--seq-len")
|
raise SystemExit("no training chunks; raise --limit or lower --batch/--seq-len")
|
||||||
|
if args.resume:
|
||||||
|
old_batch = payload.get("batch")
|
||||||
|
old_seq = payload.get("seq_len")
|
||||||
|
if old_batch is not None and (
|
||||||
|
int(old_batch) != args.batch or int(old_seq or args.seq_len) != args.seq_len
|
||||||
|
):
|
||||||
|
print(
|
||||||
|
f"warning: resume pack [{old_batch}, {old_seq}] -> "
|
||||||
|
f"[{args.batch}, {args.seq_len}]; reset chunk_index 0 "
|
||||||
|
f"(tokens/opt_step kept)"
|
||||||
|
)
|
||||||
|
chunk_index = 0
|
||||||
|
|
||||||
optim = torch.optim.AdamW(
|
optim = torch.optim.AdamW(
|
||||||
model.parameters(),
|
model.parameters(),
|
||||||
@@ -355,16 +474,23 @@ 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 and tokens >= args.max_tokens:
|
if args.max_tokens is not None:
|
||||||
|
if tokens >= args.max_tokens:
|
||||||
break
|
break
|
||||||
if micro_step >= args.steps:
|
elif micro_step >= args.steps:
|
||||||
break
|
break
|
||||||
x, y = x.to(device), y.to(device)
|
x, y = x.to(device), y.to(device)
|
||||||
scale = lr_scale(opt_step, args.warmup, horizon)
|
scale = lr_scale(opt_step, args.warmup, horizon)
|
||||||
_set_lr(optim, args.lr * scale)
|
_set_lr(optim, args.lr * scale)
|
||||||
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_bf16):
|
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_bf16):
|
||||||
loss = model(x, labels=y) / args.grad_acc
|
task = model(x, labels=y)
|
||||||
|
aux, z_loss = moe_router_losses(model)
|
||||||
|
loss = (task + aux + z_loss) / args.grad_acc
|
||||||
loss.backward()
|
loss.backward()
|
||||||
do_step = (micro_step + 1) % args.grad_acc == 0
|
do_step = (micro_step + 1) % args.grad_acc == 0
|
||||||
grad_norm = None
|
grad_norm = None
|
||||||
@@ -374,27 +500,38 @@ def main() -> None:
|
|||||||
optim.zero_grad(set_to_none=True)
|
optim.zero_grad(set_to_none=True)
|
||||||
opt_step += 1
|
opt_step += 1
|
||||||
|
|
||||||
raw_loss = loss.item() * args.grad_acc
|
raw_loss = float(task.detach())
|
||||||
tokens += tpm
|
tokens += tpm
|
||||||
micro_step += 1
|
micro_step += 1
|
||||||
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 = {"train/loss": raw_loss, "train/lr": lr_now, "train/tokens": tokens}
|
metrics = {
|
||||||
|
"train/loss": raw_loss,
|
||||||
|
"train/lr": lr_now,
|
||||||
|
"train/tokens": tokens,
|
||||||
|
"moe/aux": float(aux.detach()),
|
||||||
|
"moe/z": float(z_loss.detach()),
|
||||||
|
}
|
||||||
if grad_norm is not None:
|
if grad_norm is not None:
|
||||||
metrics["train/grad_norm"] = grad_norm
|
metrics["train/grad_norm"] = grad_norm
|
||||||
elapsed = time.perf_counter() - t0
|
elapsed = time.perf_counter() - t0
|
||||||
if elapsed > 0:
|
if elapsed > 0:
|
||||||
metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed
|
metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed
|
||||||
|
|
||||||
log_now = (
|
ended = (args.max_tokens is not None and tokens >= args.max_tokens) or (
|
||||||
micro_step % args.eval_every == 0
|
args.max_tokens is None and micro_step >= args.steps
|
||||||
or micro_step == 1
|
|
||||||
or (args.max_tokens is not None and tokens >= args.max_tokens)
|
|
||||||
or micro_step >= args.steps
|
|
||||||
)
|
)
|
||||||
if log_now:
|
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
|
||||||
|
ckpt_now = micro_step % args.ckpt_every == 0 or ended
|
||||||
|
gen_now = args.gen_every > 0 and (
|
||||||
|
micro_step % args.gen_every == 0 or micro_step == 1 or ended
|
||||||
|
)
|
||||||
|
|
||||||
|
if eval_now:
|
||||||
held = _heldout_loss(model, held_chunks, device, use_bf16)
|
held = _heldout_loss(model, held_chunks, device, use_bf16)
|
||||||
if held is not None:
|
if held is not None:
|
||||||
metrics["heldout/loss"] = held
|
metrics["heldout/loss"] = held
|
||||||
@@ -403,8 +540,27 @@ def main() -> None:
|
|||||||
f"micro {micro_step:6d} opt {opt_step:6d} tok {tokens:,} "
|
f"micro {micro_step:6d} opt {opt_step:6d} tok {tokens:,} "
|
||||||
f"loss {raw_loss:.4f} lr {lr_now:.2e}"
|
f"loss {raw_loss:.4f} lr {lr_now:.2e}"
|
||||||
+ (f" held {held:.4f}" if held is not None else "")
|
+ (f" held {held:.4f}" if held is not None else "")
|
||||||
|
+ f" aux {metrics['moe/aux']:.4f} z {metrics['moe/z']:.4f}"
|
||||||
)
|
)
|
||||||
if micro_step % (args.eval_every * 2) == 0 or micro_step <= args.eval_every:
|
if held is not None and held < best_heldout:
|
||||||
|
best_heldout = held
|
||||||
|
payload = _payload(
|
||||||
|
cfg,
|
||||||
|
model,
|
||||||
|
optim,
|
||||||
|
args,
|
||||||
|
micro_step=micro_step,
|
||||||
|
opt_step=opt_step,
|
||||||
|
tokens=tokens,
|
||||||
|
chunk_index=next_chunk,
|
||||||
|
best_heldout=best_heldout,
|
||||||
|
)
|
||||||
|
_save(_sibling(args.out, "_best"), payload)
|
||||||
|
print(
|
||||||
|
f" best held-out {best_heldout:.4f} -> {_sibling(args.out, '_best')}"
|
||||||
|
)
|
||||||
|
del payload
|
||||||
|
if gen_now:
|
||||||
for prefix in args.gen_prefix:
|
for prefix in args.gen_prefix:
|
||||||
sample = gen_sample(prefix)
|
sample = gen_sample(prefix)
|
||||||
print(f" gen[{prefix[:16]}]: {sample}")
|
print(f" gen[{prefix[:16]}]: {sample}")
|
||||||
@@ -415,6 +571,7 @@ def main() -> None:
|
|||||||
{f"gen/{prefix[:24]}": swanlab.Text(sample)},
|
{f"gen/{prefix[:24]}": swanlab.Text(sample)},
|
||||||
step=micro_step,
|
step=micro_step,
|
||||||
)
|
)
|
||||||
|
if ckpt_now:
|
||||||
payload = _payload(
|
payload = _payload(
|
||||||
cfg,
|
cfg,
|
||||||
model,
|
model,
|
||||||
@@ -423,18 +580,12 @@ 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)
|
||||||
if held is not None and held < best_heldout:
|
del payload
|
||||||
best_heldout = held
|
if tracker is not None and (log_now or eval_now):
|
||||||
payload["best_heldout"] = best_heldout
|
|
||||||
_save(_sibling(args.out, "_best"), payload)
|
|
||||||
print(
|
|
||||||
f" best held-out {best_heldout:.4f} -> {_sibling(args.out, '_best')}"
|
|
||||||
)
|
|
||||||
if tracker is not None:
|
|
||||||
tracker.log(metrics, step=micro_step)
|
tracker.log(metrics, step=micro_step)
|
||||||
|
|
||||||
payload = _payload(
|
payload = _payload(
|
||||||
@@ -445,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)
|
||||||
@@ -455,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__":
|
||||||
|
|||||||
+239
-42
@@ -1,9 +1,10 @@
|
|||||||
"""Instruction SFT for zh↔en translation. Prompt template matches eval_mt.
|
"""Instruction SFT for zh↔en translation. Prompt template matches eval_mt.
|
||||||
|
|
||||||
用法:
|
用法:
|
||||||
uv run python train_sft.py --ckpt ckpts/k3_wiki.pt --data data/sft/train.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 data/sft/opus.jsonl \\
|
uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data opus-100 \\
|
||||||
--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
|
||||||
@@ -14,14 +15,16 @@ from dataclasses import asdict
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from kda.layers.latent_moe import moe_router_losses
|
||||||
from kda.training.data import (
|
from kda.training.data import (
|
||||||
IGNORE_INDEX,
|
IGNORE_INDEX,
|
||||||
iter_sft_batches,
|
iter_sft_batches,
|
||||||
load_sft_rows,
|
resolve_sft_rows,
|
||||||
load_tokenizer,
|
load_tokenizer,
|
||||||
)
|
)
|
||||||
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
|
||||||
|
|
||||||
|
|
||||||
@@ -30,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,
|
||||||
@@ -51,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
|
||||||
@@ -68,10 +106,66 @@ 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("--data", required=True, help="jsonl {src,tgt,target_lang} or TSV")
|
p.add_argument(
|
||||||
|
"--resume",
|
||||||
|
nargs="?",
|
||||||
|
const="__last__",
|
||||||
|
default=None,
|
||||||
|
help="resume SFT; default path is <out>_last",
|
||||||
|
)
|
||||||
|
p.add_argument(
|
||||||
|
"--data",
|
||||||
|
default="opus-100",
|
||||||
|
help="local jsonl/tsv, or 'opus-100' to stream Helsinki-NLP/opus-100 en-zh",
|
||||||
|
)
|
||||||
|
p.add_argument(
|
||||||
|
"--limit",
|
||||||
|
type=int,
|
||||||
|
default=100_000,
|
||||||
|
help="OPUS source pairs to pull (each becomes zh2en + en2zh unless --one-dir)",
|
||||||
|
)
|
||||||
|
p.add_argument(
|
||||||
|
"--one-dir",
|
||||||
|
action="store_true",
|
||||||
|
help="only zh→en rows when pulling OPUS",
|
||||||
|
)
|
||||||
p.add_argument("--out", default="ckpts/k3_sft.pt")
|
p.add_argument("--out", default="ckpts/k3_sft.pt")
|
||||||
p.add_argument("--tokenizer", default=None)
|
p.add_argument("--tokenizer", default=None)
|
||||||
p.add_argument("--batch", type=int, default=4)
|
p.add_argument("--batch", type=int, default=4)
|
||||||
@@ -82,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"
|
||||||
@@ -95,17 +203,23 @@ 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)
|
||||||
rows = load_sft_rows(args.data)
|
rows = resolve_sft_rows(
|
||||||
|
args.data,
|
||||||
|
limit=args.limit,
|
||||||
|
both_dirs=not args.one_dir,
|
||||||
|
)
|
||||||
if not rows:
|
if not rows:
|
||||||
raise SystemExit(f"no SFT rows in {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
|
||||||
@@ -118,36 +232,98 @@ 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)
|
||||||
_set_lr(optim, args.lr * lr_scale(opt_step, args.warmup, horizon))
|
_set_lr(optim, args.lr * lr_scale(opt_step, args.warmup, horizon))
|
||||||
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_bf16):
|
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_bf16):
|
||||||
loss = model(x, labels=y, ignore_index=IGNORE_INDEX) / args.grad_acc
|
task = model(x, labels=y, ignore_index=IGNORE_INDEX)
|
||||||
|
aux, z_loss = moe_router_losses(model)
|
||||||
|
loss = (task + aux + z_loss) / args.grad_acc
|
||||||
loss.backward()
|
loss.backward()
|
||||||
if (step + 1) % args.grad_acc == 0:
|
if (step + 1) % args.grad_acc == 0:
|
||||||
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
||||||
optim.step()
|
optim.step()
|
||||||
optim.zero_grad(set_to_none=True)
|
optim.zero_grad(set_to_none=True)
|
||||||
opt_step += 1
|
opt_step += 1
|
||||||
raw = loss.item() * args.grad_acc
|
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(
|
||||||
{"sft/loss": raw, "sft/lr": optim.param_groups[0]["lr"]},
|
{
|
||||||
|
"sft/loss": raw,
|
||||||
|
"sft/lr": optim.param_groups[0]["lr"],
|
||||||
|
"moe/aux": float(aux.detach()),
|
||||||
|
"moe/z": float(z_loss.detach()),
|
||||||
|
},
|
||||||
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
|
||||||
print(f"step {step:4d} sft loss {raw:.4f} lr {optim.param_groups[0]['lr']:.2e}")
|
ckpt_now = args.ckpt_every > 0 and step > 0 and (
|
||||||
|
step % args.ckpt_every == 0 or step == max_micro - 1
|
||||||
|
)
|
||||||
|
if eval_now:
|
||||||
|
print(
|
||||||
|
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()
|
||||||
srcs, refs = _read_lines(args.src), _read_lines(args.ref)
|
srcs, refs = _read_lines(args.src), _read_lines(args.ref)
|
||||||
@@ -163,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(
|
||||||
{
|
{
|
||||||
@@ -172,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()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user