LatentMoE: K3 sigmoid routing and Switch aux/z-loss
Route with σ(W_r x), Top-k(s+b), then L1-normalize over the selected set. Add Switch/GShard aux and router z-loss into train_k3 and train_sft. Wiki parquet URLs honor HF_ENDPOINT for mirrored downloads.
This commit is contained in:
@@ -9,13 +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 参数结构不变).
|
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
|
||||||
@@ -24,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)
|
||||||
@@ -49,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)]
|
||||||
)
|
)
|
||||||
@@ -67,6 +75,8 @@ class LatentMoE(nn.Module):
|
|||||||
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_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:
|
||||||
@@ -79,6 +89,8 @@ 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(
|
def _routed_u(
|
||||||
@@ -133,15 +145,43 @@ class LatentMoE(nn.Module):
|
|||||||
u_flat = u_flat.index_add(0, tok, weighted)
|
u_flat = u_flat.index_add(0, tok, weighted)
|
||||||
return u_flat.view(B, T, ell)
|
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, 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)
|
||||||
|
else:
|
||||||
|
zero = logits.new_zeros(())
|
||||||
|
self.last_aux_loss = zero
|
||||||
|
self.last_z_loss = zero
|
||||||
|
|
||||||
u = self._routed_u(z, ids, probs)
|
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]
|
||||||
@@ -162,3 +202,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()
|
||||||
|
|||||||
@@ -53,6 +53,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"
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
@@ -91,7 +94,7 @@ def _wiki_files(lang: str, n_shards: int) -> list[str]:
|
|||||||
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)]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+144
-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$ 上,稀疏执行,见 \S7.4):
|
||||||
\[
|
\[
|
||||||
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,132 @@ 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 = σ(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 / Σ_{j∈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{① permute-dispatch}(按专家排序 + pad 到 $[R, C, \ell]$)→
|
||||||
|
\textbf{② padded bmm}(专家参数堆成 batch 维,三次 batched GEMM 一次算完 $R$ 个专家)→
|
||||||
|
\textbf{③ 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}{ll}
|
||||||
|
\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 Σ 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.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} 一起反传。aux/z 只更新 router 参数,
|
||||||
|
不碰专家权重(\texttt{ids} 已 \texttt{.detach()})。
|
||||||
|
|
||||||
\subsection{形状总览}
|
\subsection{形状总览}
|
||||||
|
|
||||||
\begin{center}
|
\begin{center}
|
||||||
@@ -140,12 +237,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 +256,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 防路由塌缩。
|
||||||
|
|||||||
@@ -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$ & 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, \ell} & dispatch 后 pad 到容量 $C$ \\
|
||||||
|
$C$ & 标量 & 最大专家负载(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}
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -89,12 +89,9 @@ def test_preset_0_5b_schedule():
|
|||||||
|
|
||||||
|
|
||||||
def _dense_moe_forward(moe: LatentMoE, x: torch.Tensor) -> torch.Tensor:
|
def _dense_moe_forward(moe: LatentMoE, x: torch.Tensor) -> torch.Tensor:
|
||||||
"""Old dense path: run every routed expert, then gather top-k."""
|
"""Dense path: run every routed expert, then gather K3 sigmoid-norm top-k."""
|
||||||
z = moe.down(x)
|
z = moe.down(x)
|
||||||
logits = moe.router(x)
|
ids, probs = moe._route(moe.router(x))
|
||||||
topk = torch.topk(logits, moe.top_k, dim=-1)
|
|
||||||
ids = topk.indices
|
|
||||||
probs = F.softmax(topk.values, dim=-1)
|
|
||||||
all_out = torch.stack([expert(z) for expert in moe.experts])
|
all_out = torch.stack([expert(z) for expert in moe.experts])
|
||||||
B, T, _ = x.shape
|
B, T, _ = x.shape
|
||||||
all_out = all_out.permute(1, 2, 0, 3).reshape(B * T, moe.n_routed, moe.latent_size)
|
all_out = all_out.permute(1, 2, 0, 3).reshape(B * T, moe.n_routed, moe.latent_size)
|
||||||
@@ -113,20 +110,22 @@ def test_moe_router_activates_topk_only():
|
|||||||
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
|
||||||
@@ -188,6 +187,79 @@ def test_moe_unselected_experts_have_zero_grad():
|
|||||||
assert torch.equal(param.grad, torch.zeros_like(param.grad))
|
assert torch.equal(param.grad, torch.zeros_like(param.grad))
|
||||||
|
|
||||||
|
|
||||||
|
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():
|
||||||
torch.manual_seed(51)
|
torch.manual_seed(51)
|
||||||
cfg = K3Config(hidden_size=64, num_hidden_layers=4, num_heads=4, head_dim=8,
|
cfg = K3Config(hidden_size=64, num_hidden_layers=4, num_heads=4, head_dim=8,
|
||||||
|
|||||||
+45
-5
@@ -17,7 +17,7 @@ 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
|
||||||
@@ -91,6 +91,8 @@ 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,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
@@ -148,6 +150,13 @@ 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 _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:
|
||||||
@@ -225,6 +234,18 @@ 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()
|
||||||
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"]
|
||||||
@@ -252,6 +273,10 @@ def main() -> None:
|
|||||||
cfg.attnres_block_size = args.attnres_block_size
|
cfg.attnres_block_size = args.attnres_block_size
|
||||||
if args.grad_checkpoint is not None:
|
if args.grad_checkpoint is not None:
|
||||||
cfg.gradient_checkpointing = args.grad_checkpoint
|
cfg.gradient_checkpointing = args.grad_checkpoint
|
||||||
|
if args.moe_aux_coef is not None:
|
||||||
|
cfg.moe_aux_loss_coef = args.moe_aux_coef
|
||||||
|
if args.moe_z_coef is not None:
|
||||||
|
cfg.moe_z_loss_coef = args.moe_z_coef
|
||||||
|
|
||||||
langs = [part.strip() for part in args.langs.split(",") if part.strip()]
|
langs = [part.strip() for part in args.langs.split(",") if part.strip()]
|
||||||
tpm = tokens_per_micro(args.batch, args.seq_len)
|
tpm = tokens_per_micro(args.batch, args.seq_len)
|
||||||
@@ -283,6 +308,10 @@ def main() -> None:
|
|||||||
cfg.attnres_block_size = args.attnres_block_size
|
cfg.attnres_block_size = args.attnres_block_size
|
||||||
if args.grad_checkpoint is not None:
|
if args.grad_checkpoint is not None:
|
||||||
cfg.gradient_checkpointing = args.grad_checkpoint
|
cfg.gradient_checkpointing = args.grad_checkpoint
|
||||||
|
if args.moe_aux_coef is not None:
|
||||||
|
cfg.moe_aux_loss_coef = args.moe_aux_coef
|
||||||
|
if args.moe_z_coef is not None:
|
||||||
|
cfg.moe_z_loss_coef = args.moe_z_coef
|
||||||
model.gradient_checkpointing = cfg.gradient_checkpointing
|
model.gradient_checkpointing = cfg.gradient_checkpointing
|
||||||
model.to(device)
|
model.to(device)
|
||||||
payload = torch.load(args.resume, map_location="cpu", weights_only=False)
|
payload = torch.load(args.resume, map_location="cpu", weights_only=False)
|
||||||
@@ -296,6 +325,7 @@ def main() -> None:
|
|||||||
else:
|
else:
|
||||||
model = CausalLM(cfg).to(device)
|
model = CausalLM(cfg).to(device)
|
||||||
|
|
||||||
|
_apply_moe_coefs(model, cfg)
|
||||||
tracker = _init_swanlab(cfg, args)
|
tracker = _init_swanlab(cfg, args)
|
||||||
n = sum(p.numel() for p in model.parameters())
|
n = sum(p.numel() for p in model.parameters())
|
||||||
print(
|
print(
|
||||||
@@ -304,7 +334,8 @@ 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(
|
||||||
@@ -364,7 +395,9 @@ def main() -> None:
|
|||||||
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,14 +407,20 @@ 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
|
||||||
|
|
||||||
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
|
||||||
@@ -403,6 +442,7 @@ 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 micro_step % (args.eval_every * 2) == 0 or micro_step <= args.eval_every:
|
||||||
for prefix in args.gen_prefix:
|
for prefix in args.gen_prefix:
|
||||||
|
|||||||
+14
-4
@@ -14,6 +14,7 @@ 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,
|
||||||
@@ -131,23 +132,32 @@ def main() -> None:
|
|||||||
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:
|
||||||
best = raw
|
best = raw
|
||||||
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:
|
if 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}")
|
print(
|
||||||
|
f"step {step:4d} sft loss {raw:.4f} 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)
|
||||||
|
|||||||
Reference in New Issue
Block a user