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
+7 -4
View File
@@ -19,12 +19,15 @@ import torch
from .prompts import instruction_prompt
WIKI_SHARD_TOTAL = {"zh": 6, "en": 41}
WIKI_BASE = (
"https://huggingface.co/datasets/wikimedia/wikipedia/resolve/main/20231101.{lang}"
)
WIKI_PATH = "datasets/wikimedia/wikipedia/resolve/main/20231101.{lang}"
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):
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")
total = WIKI_SHARD_TOTAL[lang]
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)]