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
+45 -5
View File
@@ -17,7 +17,7 @@ from dataclasses import asdict
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.k3_config import K3Config
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,
"kda_backend": cfg.kda_backend,
"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:
@@ -148,6 +150,13 @@ def _heldout_loss(
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:
frac = moe_route_frac(model)
if frac is None:
@@ -225,6 +234,18 @@ def main() -> None:
dest="grad_checkpoint",
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()
if args.gen_prefix is None:
args.gen_prefix = ["人工智能的发展", "The history of computing"]
@@ -252,6 +273,10 @@ def main() -> 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
langs = [part.strip() for part in args.langs.split(",") if part.strip()]
tpm = tokens_per_micro(args.batch, args.seq_len)
@@ -283,6 +308,10 @@ def main() -> 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
model.gradient_checkpointing = cfg.gradient_checkpointing
model.to(device)
payload = torch.load(args.resume, map_location="cpu", weights_only=False)
@@ -296,6 +325,7 @@ def main() -> None:
else:
model = CausalLM(cfg).to(device)
_apply_moe_coefs(model, cfg)
tracker = _init_swanlab(cfg, args)
n = sum(p.numel() for p in model.parameters())
print(
@@ -304,7 +334,8 @@ def main() -> None:
)
print(
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:
print(
@@ -364,7 +395,9 @@ def main() -> None:
scale = lr_scale(opt_step, args.warmup, horizon)
_set_lr(optim, args.lr * scale)
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()
do_step = (micro_step + 1) % args.grad_acc == 0
grad_norm = None
@@ -374,14 +407,20 @@ def main() -> None:
optim.zero_grad(set_to_none=True)
opt_step += 1
raw_loss = loss.item() * args.grad_acc
raw_loss = float(task.detach())
tokens += tpm
micro_step += 1
lr_now = optim.param_groups[0]["lr"]
if raw_loss < best_train:
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:
metrics["train/grad_norm"] = grad_norm
elapsed = time.perf_counter() - t0
@@ -403,6 +442,7 @@ def main() -> None:
f"micro {micro_step:6d} opt {opt_step:6d} tok {tokens:,} "
f"loss {raw_loss:.4f} lr {lr_now:.2e}"
+ (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:
for prefix in args.gen_prefix: