feat: add MoE auxiliary loss metrics

- Propagates MoE load-balancing loss through model outputs
- Logs task, auxiliary, and weighted losses across strategies
- Computes only explicitly requested callback metrics
- Preserves tensor compute_loss API and adds regression tests
This commit is contained in:
2026-08-02 06:30:43 +08:00
parent 0fc1b1bd46
commit 1c7369f293
14 changed files with 370 additions and 68 deletions
+12 -3
View File
@@ -81,6 +81,14 @@ Where $\rho_t = \pi_\theta(a_t|s_t) / \pi_{\text{old}}(a_t|s_t)$ is the per-toke
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`.
### MoE Load Balancing
MoE layers add a differentiable load-balancing term based on mean router probabilities and top-k expert assignment frequency. The training objective is:
$$ L = L_{\text{task}} + \lambda_{\text{MoE}} L_{\text{aux}} $$
`TrainConfig.moe_aux_loss_coef` controls $\lambda_{\text{MoE}}$ (default `0.01`). The unweighted and weighted auxiliary losses are logged separately.
## Training Loop Internals
Two-level loop: **epoch****batch**. Optimizer step fires every `grad_accum_steps` batches.
@@ -92,9 +100,10 @@ on_train_begin
for batch in dataloader:
on_batch_begin
with executor.accumulate(model):
loss = strategy.compute_loss(batch)
context.loss = loss.item()
stand_loss = loss / executor.grad_accum_steps
loss_output = strategy(batch)
context.loss = loss_output["loss"].item()
context.metrics = loss_output["metrics"]
stand_loss = loss_output["loss"] / executor.grad_accum_steps
executor.backward(stand_loss)
context.consumed_samples += (
context.config.batch_per_device * context.world_size
+6 -3
View File
@@ -54,9 +54,10 @@ on_train_begin
for batch in dataloader:
on_batch_begin
with executor.accumulate(model):
loss = strategy.compute_loss(batch)
context.loss = loss.item()
stand_loss = loss / executor.grad_accum_steps
loss_output = strategy(batch)
context.loss = loss_output["loss"].item()
context.metrics = loss_output["metrics"]
stand_loss = loss_output["loss"] / executor.grad_accum_steps
executor.backward(stand_loss)
context.consumed_samples += (
context.config.batch_per_device * context.world_size
@@ -88,6 +89,8 @@ on_train_end
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`).
Strategies return `{"loss": Tensor, "metrics": Dict[str, Tensor]}` when called by the trainer. Built-in metrics include the task-specific loss and, for MoE models, `moe_aux_loss` plus `moe_aux_loss_weighted`. Direct `compute_loss(batch)` calls continue to return a single loss tensor.
## Strategies
### SEQ (Pre-training)