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:
@@ -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)]
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user