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:
dela
2026-08-25 19:50:07 +08:00
parent d1da0816f2
commit 7a12f61de1
8 changed files with 381 additions and 82 deletions
+14 -4
View File
@@ -14,6 +14,7 @@ from dataclasses import asdict
import torch
from kda.layers.latent_moe import moe_router_losses
from kda.training.data import (
IGNORE_INDEX,
iter_sft_batches,
@@ -131,23 +132,32 @@ def main() -> None:
x, y = x.to(device), y.to(device)
_set_lr(optim, args.lr * lr_scale(opt_step, args.warmup, horizon))
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()
if (step + 1) % args.grad_acc == 0:
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optim.step()
optim.zero_grad(set_to_none=True)
opt_step += 1
raw = loss.item() * args.grad_acc
raw = float(task.detach())
if raw < best:
best = raw
if tracker is not None:
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,
)
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:
model.eval()
srcs, refs = _read_lines(args.src), _read_lines(args.ref)