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:
+45
-5
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user