Inital commit - includes v1 and v2

This commit is contained in:
2026-05-31 23:30:56 -05:00
commit 285bbd7aac
14 changed files with 1165 additions and 0 deletions
+8
View File
@@ -0,0 +1,8 @@
__pycache__/
*.py[cod]
.venv/
.env
checkpoints/
runs/
wandb/
.DS_Store
+152
View File
@@ -0,0 +1,152 @@
# VerySimpleMoE
A working PyTorch implementation of a tiny decoder-only Mixture-of-Experts language model trained from HuggingFace FineWeb.
## Architectures
The original architecture is preserved as the `v1` preset:
- 12 experts total per MoE layer
- 6 active experts per token (`top-k=6` router)
- each expert has exactly 500,000 parameters
- expert MLP: `Linear(500 -> 500, bias=False)`, GELU, `Linear(500 -> 500, bias=False)`
- params: `500*500 + 500*500 = 500,000`
- GPT-style causal self-attention before the MoE block
- learned top-k router plus load-balancing auxiliary loss
- GPT-2 tokenizer by default
The new research architecture is available as `v2-32x1m`:
- 32 experts total per MoE layer
- 4 active experts per token by default (`top-k=4` router)
- each expert has exactly 1,000,000 parameters
- expert MLP: `Linear(500 -> 1000, bias=False)`, GELU, `Linear(1000 -> 500, bias=False)`
- params: `500*1000 + 1000*500 = 1,000,000`
- router noise during training for exploration
- router z-loss to keep router logits stable
- optional phased expert training so only a subset of experts are trainable/routeable at a time
The default still uses `n_layers=1`, which means the model has one MoE block. If you raise `--n-layers`, each layer gets its own full expert set.
## Install
```bash
python3 -m venv .venv
source .venv/bin/activate
pip install -e .
```
## Train v1 on FineWeb
This streams FineWeb, so it does not download the full dataset first.
```bash
verysimplemoe-train \
--arch v1 \
--dataset-name HuggingFaceFW/fineweb \
--dataset-config sample-10BT \
--out-dir checkpoints/verysimplemoe-v1 \
--max-steps 1000 \
--batch-size 8 \
--grad-accum-steps 4 \
--block-size 256
```
## Train v2: 32 experts, 1M params per expert
Recommended laptop-friendly command:
```bash
verysimplemoe-train \
--arch v2-32x1m \
--out-dir checkpoints/verysimplemoe-v2-32x1m \
--train-experts-per-phase 16 \
--expert-phase-steps 500 \
--batch-size 4 \
--grad-accum-steps 4 \
--block-size 256
```
Phased expert training means:
- only 16 of the 32 experts are eligible for routing in a given phase
- only those 16 experts have gradients enabled
- the expert optimizer is rebuilt per phase, so Adam state is kept only for the current expert subset
- the phase window overlaps by default using a half-window stride, e.g. `0-15`, `8-23`, `16-31`, `24-31 + 0-7`
Useful router options:
```bash
--active-experts 4 # top-k experts per token
--router-noise-std 0.1 # train-time router exploration
--router-z-loss-coef 1e-4 # router logit stabilization
--aux-loss-coef 0.01 # load-balancing loss
```
## Tiny CPU smoke test
```bash
verysimplemoe-train --device cpu --max-steps 5 --batch-size 1 --grad-accum-steps 1 --block-size 64
```
For a tiny v2 smoke test:
```bash
verysimplemoe-train \
--device cpu \
--arch v2-32x1m \
--train-experts-per-phase 16 \
--max-steps 5 \
--batch-size 1 \
--grad-accum-steps 1 \
--block-size 64
```
Useful speed flags on NVIDIA GPUs:
```bash
verysimplemoe-train --amp --compile --max-steps 10000 --batch-size 16
```
Note: with phased expert training, `torch.compile` may recompile when the active expert phase changes.
## Resume after interruption
If training stops after a checkpoint save begins, resume from the checkpoint directory:
```bash
verysimplemoe-train \
--resume-from checkpoints/verysimplemoe-v2-32x1m \
--out-dir checkpoints/verysimplemoe-v2-32x1m \
--max-steps 1000 \
--batch-size 4 \
--grad-accum-steps 4 \
--block-size 256 \
--train-experts-per-phase 16
```
`--max-steps` is the final target step count, not additional steps.
## Generate with the final model
```bash
verysimplemoe-generate \
--checkpoint checkpoints/verysimplemoe-v2-32x1m \
--prompt "The future of open language models is" \
--max-new-tokens 120 \
--temperature 0.8 \
--top-k 50
```
You can also run modules without installing scripts:
```bash
PYTHONPATH=src python -m verysimplemoe.train --max-steps 100
PYTHONPATH=src python -m verysimplemoe.generate --prompt "Hello"
```
## Files
- `src/verysimplemoe/model.py` - model, MoE router, experts, generation
- `src/verysimplemoe/train.py` - FineWeb streaming trainer, phased expert training, checkpointing
- `src/verysimplemoe/generate.py` - checkpoint loader and text generation CLI
+22
View File
@@ -0,0 +1,22 @@
[build-system]
requires = ["setuptools>=68", "wheel"]
build-backend = "setuptools.build_meta"
[project]
name = "verysimplemoe"
version = "0.1.0"
description = "A very small PyTorch MoE language model trainer for FineWeb"
requires-python = ">=3.10"
dependencies = [
"torch>=2.2",
"transformers>=4.40",
"datasets>=2.19",
"tqdm>=4.66",
]
[project.scripts]
verysimplemoe-train = "verysimplemoe.train:main"
verysimplemoe-generate = "verysimplemoe.generate:main"
[tool.setuptools.packages.find]
where = ["src"]
+4
View File
@@ -0,0 +1,4 @@
torch>=2.2
transformers>=4.40
datasets>=2.19
tqdm>=4.66
+9
View File
@@ -0,0 +1,9 @@
Metadata-Version: 2.4
Name: verysimplemoe
Version: 0.1.0
Summary: A very small PyTorch MoE language model trainer for FineWeb
Requires-Python: >=3.10
Requires-Dist: torch>=2.2
Requires-Dist: transformers>=4.40
Requires-Dist: datasets>=2.19
Requires-Dist: tqdm>=4.66
+12
View File
@@ -0,0 +1,12 @@
README.md
pyproject.toml
src/verysimplemoe/__init__.py
src/verysimplemoe/generate.py
src/verysimplemoe/model.py
src/verysimplemoe/train.py
src/verysimplemoe.egg-info/PKG-INFO
src/verysimplemoe.egg-info/SOURCES.txt
src/verysimplemoe.egg-info/dependency_links.txt
src/verysimplemoe.egg-info/entry_points.txt
src/verysimplemoe.egg-info/requires.txt
src/verysimplemoe.egg-info/top_level.txt
@@ -0,0 +1 @@
@@ -0,0 +1,3 @@
[console_scripts]
verysimplemoe-generate = verysimplemoe.generate:main
verysimplemoe-train = verysimplemoe.train:main
+4
View File
@@ -0,0 +1,4 @@
torch>=2.2
transformers>=4.40
datasets>=2.19
tqdm>=4.66
+1
View File
@@ -0,0 +1 @@
verysimplemoe
+3
View File
@@ -0,0 +1,3 @@
from .model import ARCH_PRESETS, MoEConfig, SimpleMoELanguageModel
__all__ = ["ARCH_PRESETS", "MoEConfig", "SimpleMoELanguageModel"]
+65
View File
@@ -0,0 +1,65 @@
from __future__ import annotations
import argparse
import json
from pathlib import Path
import torch
from transformers import AutoTokenizer
from .model import MoEConfig, SimpleMoELanguageModel
def get_device(device_arg: str) -> torch.device:
if device_arg != "auto":
return torch.device(device_arg)
if torch.cuda.is_available():
return torch.device("cuda")
if torch.backends.mps.is_available():
return torch.device("mps")
return torch.device("cpu")
def load_model(checkpoint_dir: Path, device: torch.device, tokenizer_name: str | None = None) -> tuple[SimpleMoELanguageModel, object]:
with (checkpoint_dir / "config.json").open() as f:
config = MoEConfig.from_dict(json.load(f))
# Older interrupted checkpoints may not include tokenizer files because the
# first crash happened before tokenizer.save_pretrained(). Fall back to GPT-2.
tokenizer = AutoTokenizer.from_pretrained(tokenizer_name or checkpoint_dir)
model = SimpleMoELanguageModel(config)
state = torch.load(checkpoint_dir / "model.pt", map_location=device)
model.load_state_dict(state)
model.to(device)
model.eval()
return model, tokenizer
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Run a trained VerySimpleMoE checkpoint")
p.add_argument("--checkpoint", type=Path, default=Path("checkpoints/verysimplemoe"))
p.add_argument("--tokenizer", default=None, help="Tokenizer name/path. Use gpt2 for interrupted checkpoints without tokenizer files.")
p.add_argument("--prompt", default="Once upon a time")
p.add_argument("--max-new-tokens", type=int, default=100)
p.add_argument("--temperature", type=float, default=0.8)
p.add_argument("--top-k", type=int, default=50)
p.add_argument("--device", default="auto")
return p.parse_args()
def main() -> None:
args = parse_args()
device = get_device(args.device)
model, tokenizer = load_model(args.checkpoint, device, args.tokenizer)
input_ids = tokenizer.encode(args.prompt, return_tensors="pt").to(device)
with torch.no_grad():
output_ids = model.generate(
input_ids,
max_new_tokens=args.max_new_tokens,
temperature=args.temperature,
top_k=args.top_k,
)
print(tokenizer.decode(output_ids[0], skip_special_tokens=True))
if __name__ == "__main__":
main()
+326
View File
@@ -0,0 +1,326 @@
from __future__ import annotations
from dataclasses import asdict, dataclass
from typing import Any, Optional, Sequence
import torch
import torch.nn as nn
import torch.nn.functional as F
@dataclass
class MoEConfig:
"""Configuration for a tiny decoder-only MoE language model.
The default values are the original v1 architecture. Keep these defaults
stable so old experiments/checkpoints remain reproducible.
v1 experts have exactly 500,000 trainable parameters:
Linear(500 -> 500, bias=False) + Linear(500 -> 500, bias=False)
= 500 * 500 + 500 * 500 = 500,000
"""
arch: str = "v1"
vocab_size: int = 50257
block_size: int = 256
n_layers: int = 1
d_model: int = 500
n_heads: int = 10
dropout: float = 0.1
n_experts: int = 12
active_experts: int = 6
expert_hidden_size: int = 500
aux_loss_coef: float = 0.01
router_noise_std: float = 0.0
router_z_loss_coef: float = 0.0
def to_dict(self) -> dict:
return asdict(self)
@classmethod
def from_dict(cls, data: dict) -> "MoEConfig":
return cls(**data)
ARCH_PRESETS: dict[str, dict[str, int | float | str]] = {
"v1": {
"arch": "v1",
"d_model": 500,
"n_heads": 10,
"n_experts": 12,
"active_experts": 6,
"expert_hidden_size": 500,
"router_noise_std": 0.0,
"router_z_loss_coef": 0.0,
},
"v2-32x1m": {
"arch": "v2-32x1m",
"d_model": 500,
"n_heads": 10,
"n_experts": 32,
"active_experts": 4,
"expert_hidden_size": 1000,
"router_noise_std": 0.1,
"router_z_loss_coef": 1e-4,
},
}
class Expert(nn.Module):
def __init__(self, d_model: int, hidden_size: int, dropout: float):
super().__init__()
# No biases: with d_model=hidden_size=500 this is exactly 500k params.
self.net = nn.Sequential(
nn.Linear(d_model, hidden_size, bias=False),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(hidden_size, d_model, bias=False),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.net(x)
class TopKMoE(nn.Module):
def __init__(self, config: MoEConfig):
super().__init__()
if config.active_experts > config.n_experts:
raise ValueError("active_experts must be <= n_experts")
self.n_experts = config.n_experts
self.active_experts = config.active_experts
self.aux_loss_coef = config.aux_loss_coef
self.router_noise_std = config.router_noise_std
self.router_z_loss_coef = config.router_z_loss_coef
self.router = nn.Linear(config.d_model, config.n_experts, bias=False)
self.experts = nn.ModuleList(
Expert(config.d_model, config.expert_hidden_size, config.dropout)
for _ in range(config.n_experts)
)
def expert_parameter_counts(self) -> list[int]:
return [sum(p.numel() for p in expert.parameters()) for expert in self.experts]
def forward(
self,
x: torch.Tensor,
active_expert_ids: Optional[Sequence[int] | torch.Tensor] = None,
collect_router_stats: bool = False,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str, torch.Tensor] | None]:
bsz, seq_len, d_model = x.shape
flat_x = x.reshape(-1, d_model)
router_logits = self.router(flat_x) # [tokens, experts]
if active_expert_ids is None:
eligible_expert_ids = torch.arange(self.n_experts, device=flat_x.device)
eligible_logits = router_logits
else:
eligible_expert_ids = torch.as_tensor(active_expert_ids, dtype=torch.long, device=flat_x.device)
if eligible_expert_ids.numel() < self.active_experts:
raise ValueError("active_expert_ids must contain at least active_experts entries")
eligible_logits = router_logits.index_select(dim=-1, index=eligible_expert_ids)
if self.training and self.router_noise_std > 0:
eligible_logits = eligible_logits + torch.randn_like(eligible_logits) * self.router_noise_std
eligible_probs = F.softmax(eligible_logits, dim=-1)
topk_probs, topk_local_idx = torch.topk(eligible_probs, self.active_experts, dim=-1)
topk_idx = eligible_expert_ids.index_select(0, topk_local_idx.reshape(-1)).reshape_as(topk_local_idx)
# Normalize selected expert weights so the mixed expert output scale is stable.
topk_probs = topk_probs / topk_probs.sum(dim=-1, keepdim=True).clamp_min(1e-9)
flat_out = torch.zeros_like(flat_x)
for expert_id, expert in enumerate(self.experts):
token_pos, choice_pos = torch.where(topk_idx == expert_id)
if token_pos.numel() == 0:
continue
expert_in = flat_x.index_select(0, token_pos)
expert_out = expert(expert_in)
weights = topk_probs[token_pos, choice_pos].unsqueeze(-1).to(expert_out.dtype)
flat_out.index_add_(0, token_pos, expert_out * weights)
# Switch Transformer-style load balancing term, computed only over the
# currently eligible experts. This matters when training a 16-of-32
# expert phase: the router should balance the phase, not unreachable
# experts.
n_eligible = eligible_expert_ids.numel()
importance_local = eligible_probs.mean(dim=0)
one_hot_local = F.one_hot(topk_local_idx, num_classes=n_eligible).float().sum(dim=1)
load_local = one_hot_local.mean(dim=0) / self.active_experts
load_balance_loss = n_eligible * torch.sum(importance_local * load_local) * self.aux_loss_coef
if self.router_z_loss_coef > 0:
router_z_loss = torch.mean(torch.logsumexp(eligible_logits, dim=-1).pow(2)) * self.router_z_loss_coef
else:
router_z_loss = x.new_zeros(())
router_stats = None
if collect_router_stats:
importance = torch.zeros(self.n_experts, device=flat_x.device, dtype=importance_local.dtype)
load = torch.zeros(self.n_experts, device=flat_x.device, dtype=load_local.dtype)
importance.index_copy_(0, eligible_expert_ids, importance_local.detach())
load.index_copy_(0, eligible_expert_ids, load_local.detach())
entropy = -(eligible_probs.detach() * eligible_probs.detach().clamp_min(1e-9).log()).sum(dim=-1).mean()
router_stats = {
"eligible_expert_ids": eligible_expert_ids.detach().cpu(),
"importance": importance.detach().cpu(),
"load": load.detach().cpu(),
"entropy": entropy.detach().cpu(),
}
return flat_out.reshape(bsz, seq_len, d_model), load_balance_loss, router_z_loss, router_stats
class CausalSelfAttention(nn.Module):
def __init__(self, config: MoEConfig):
super().__init__()
if config.d_model % config.n_heads != 0:
raise ValueError("d_model must be divisible by n_heads")
self.n_heads = config.n_heads
self.head_dim = config.d_model // config.n_heads
self.qkv = nn.Linear(config.d_model, 3 * config.d_model)
self.proj = nn.Linear(config.d_model, config.d_model)
self.dropout_p = config.dropout
self.resid_dropout = nn.Dropout(config.dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
bsz, seq_len, d_model = x.shape
qkv = self.qkv(x)
q, k, v = qkv.chunk(3, dim=-1)
q = q.view(bsz, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
k = k.view(bsz, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
v = v.view(bsz, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
y = F.scaled_dot_product_attention(
q,
k,
v,
attn_mask=None,
dropout_p=self.dropout_p if self.training else 0.0,
is_causal=True,
)
y = y.transpose(1, 2).contiguous().view(bsz, seq_len, d_model)
return self.resid_dropout(self.proj(y))
class Block(nn.Module):
def __init__(self, config: MoEConfig):
super().__init__()
self.ln1 = nn.LayerNorm(config.d_model)
self.attn = CausalSelfAttention(config)
self.ln2 = nn.LayerNorm(config.d_model)
self.moe = TopKMoE(config)
def forward(
self,
x: torch.Tensor,
active_expert_ids: Optional[Sequence[int] | torch.Tensor] = None,
collect_router_stats: bool = False,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str, torch.Tensor] | None]:
x = x + self.attn(self.ln1(x))
moe_out, load_balance_loss, router_z_loss, router_stats = self.moe(
self.ln2(x),
active_expert_ids=active_expert_ids,
collect_router_stats=collect_router_stats,
)
x = x + moe_out
return x, load_balance_loss, router_z_loss, router_stats
class SimpleMoELanguageModel(nn.Module):
def __init__(self, config: MoEConfig):
super().__init__()
self.config = config
self.token_embedding = nn.Embedding(config.vocab_size, config.d_model)
self.position_embedding = nn.Embedding(config.block_size, config.d_model)
self.drop = nn.Dropout(config.dropout)
self.blocks = nn.ModuleList(Block(config) for _ in range(config.n_layers))
self.ln_f = nn.LayerNorm(config.d_model)
self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False)
self.lm_head.weight = self.token_embedding.weight # weight tying
self.apply(self._init_weights)
def _init_weights(self, module: nn.Module) -> None:
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
def expert_parameter_counts(self) -> list[list[int]]:
return [block.moe.expert_parameter_counts() for block in self.blocks]
def forward(
self,
input_ids: torch.Tensor,
labels: Optional[torch.Tensor] = None,
active_expert_ids: Optional[Sequence[int] | torch.Tensor] = None,
collect_router_stats: bool = False,
) -> dict[str, Any]:
bsz, seq_len = input_ids.shape
if seq_len > self.config.block_size:
raise ValueError(f"Sequence length {seq_len} exceeds block_size {self.config.block_size}")
positions = torch.arange(0, seq_len, device=input_ids.device).unsqueeze(0)
x = self.token_embedding(input_ids) + self.position_embedding(positions)
x = self.drop(x)
load_balance_loss = x.new_zeros(())
router_z_loss = x.new_zeros(())
router_stats = []
for block in self.blocks:
x, block_load_balance, block_router_z, block_stats = block(
x,
active_expert_ids=active_expert_ids,
collect_router_stats=collect_router_stats,
)
load_balance_loss = load_balance_loss + block_load_balance
router_z_loss = router_z_loss + block_router_z
if collect_router_stats and block_stats is not None:
router_stats.append(block_stats)
x = self.ln_f(x)
logits = self.lm_head(x)
loss = None
lm_loss = None
if labels is not None:
lm_loss = F.cross_entropy(logits.view(-1, logits.size(-1)), labels.reshape(-1))
loss = lm_loss + load_balance_loss + router_z_loss
aux_loss = load_balance_loss + router_z_loss
return {
"logits": logits,
"loss": loss,
"lm_loss": lm_loss,
"aux_loss": aux_loss,
"load_balance_loss": load_balance_loss,
"router_z_loss": router_z_loss,
"router_stats": router_stats,
}
@torch.no_grad()
def generate(
self,
input_ids: torch.Tensor,
max_new_tokens: int = 100,
temperature: float = 0.8,
top_k: Optional[int] = 50,
) -> torch.Tensor:
self.eval()
for _ in range(max_new_tokens):
idx_cond = input_ids[:, -self.config.block_size :]
logits = self(idx_cond)["logits"][:, -1, :]
if temperature <= 0:
next_id = torch.argmax(logits, dim=-1, keepdim=True)
else:
logits = logits / temperature
if top_k is not None and top_k > 0:
values, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits = logits.masked_fill(logits < values[:, [-1]], -float("inf"))
probs = F.softmax(logits, dim=-1)
next_id = torch.multinomial(probs, num_samples=1)
input_ids = torch.cat((input_ids, next_id), dim=1)
return input_ids
+555
View File
@@ -0,0 +1,555 @@
from __future__ import annotations
import argparse
import json
import math
import random
import time
from pathlib import Path
from typing import Iterator, Sequence
import torch
from datasets import load_dataset
from torch.utils.data import DataLoader, IterableDataset
from tqdm.auto import tqdm
from transformers import AutoTokenizer
from .model import ARCH_PRESETS, MoEConfig, SimpleMoELanguageModel
class FineWebTokenDataset(IterableDataset):
"""Streams FineWeb text and emits fixed-length next-token prediction chunks."""
def __init__(
self,
tokenizer,
block_size: int,
dataset_name: str = "HuggingFaceFW/fineweb",
dataset_config: str = "sample-10BT",
split: str = "train",
shuffle_buffer: int = 10_000,
seed: int = 1337,
text_column: str = "text",
):
self.tokenizer = tokenizer
self.block_size = block_size
self.dataset_name = dataset_name
self.dataset_config = dataset_config
self.split = split
self.shuffle_buffer = shuffle_buffer
self.seed = seed
self.text_column = text_column
def __iter__(self) -> Iterator[dict[str, torch.Tensor]]:
ds = load_dataset(
self.dataset_name,
name=self.dataset_config,
split=self.split,
streaming=True,
)
if self.shuffle_buffer > 0:
worker_info = torch.utils.data.get_worker_info()
worker_id = 0 if worker_info is None else worker_info.id
ds = ds.shuffle(buffer_size=self.shuffle_buffer, seed=self.seed + worker_id)
eos = self.tokenizer.eos_token_id
token_buffer: list[int] = []
for row in ds:
text = row.get(self.text_column)
if not text:
continue
# FineWeb rows can exceed GPT-2 tokenizer.model_max_length (1024),
# but we chunk tokens ourselves below, so suppress that irrelevant
# tokenizer warning instead of truncating the document.
token_buffer.extend(self.tokenizer.encode(text, add_special_tokens=False, verbose=False))
token_buffer.append(eos)
while len(token_buffer) >= self.block_size + 1:
chunk = token_buffer[: self.block_size + 1]
del token_buffer[: self.block_size]
x = torch.tensor(chunk[:-1], dtype=torch.long)
y = torch.tensor(chunk[1:], dtype=torch.long)
yield {"input_ids": x, "labels": y}
def set_seed(seed: int) -> None:
random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
def get_device(device_arg: str) -> torch.device:
if device_arg != "auto":
return torch.device(device_arg)
if torch.cuda.is_available():
return torch.device("cuda")
if torch.backends.mps.is_available():
return torch.device("mps")
return torch.device("cpu")
def save_checkpoint(
out_dir: Path,
model: SimpleMoELanguageModel,
tokenizer,
trainer_state: dict,
args: argparse.Namespace,
) -> None:
out_dir.mkdir(parents=True, exist_ok=True)
torch.save(model.state_dict(), out_dir / "model.pt")
torch.save(trainer_state, out_dir / "trainer_state.pt")
with (out_dir / "config.json").open("w") as f:
json.dump(model.config.to_dict(), f, indent=2)
with (out_dir / "train_args.json").open("w") as f:
# argparse contains pathlib.Path values such as --out-dir.
json.dump(vars(args), f, indent=2, default=str)
tokenizer.save_pretrained(out_dir)
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Train VerySimpleMoE on FineWeb")
p.add_argument("--arch", choices=sorted(ARCH_PRESETS), default="v1", help="Architecture preset. v1 preserves the original model.")
p.add_argument("--out-dir", type=Path, default=Path("checkpoints/verysimplemoe"))
p.add_argument("--resume-from", type=Path, default=None, help="Resume model/optimizer state from a checkpoint directory")
p.add_argument("--tokenizer", default="gpt2")
p.add_argument("--dataset-name", default="HuggingFaceFW/fineweb")
p.add_argument("--dataset-config", default="sample-10BT")
p.add_argument("--split", default="train")
p.add_argument("--text-column", default="text")
p.add_argument("--max-steps", type=int, default=1000)
p.add_argument("--batch-size", type=int, default=8)
p.add_argument("--grad-accum-steps", type=int, default=4)
p.add_argument("--lr", type=float, default=3e-4)
p.add_argument("--weight-decay", type=float, default=0.1)
p.add_argument("--warmup-steps", type=int, default=100)
p.add_argument("--max-grad-norm", type=float, default=1.0)
p.add_argument("--num-workers", type=int, default=0)
p.add_argument("--shuffle-buffer", type=int, default=10_000)
p.add_argument("--seed", type=int, default=1337)
p.add_argument("--device", default="auto", help="auto, cpu, cuda, cuda:0, or mps")
p.add_argument("--amp", action="store_true", help="Use bfloat16 autocast on CUDA")
p.add_argument("--compile", action="store_true", help="torch.compile the model")
p.add_argument("--save-every", type=int, default=500)
p.add_argument("--log-every", type=int, default=10)
# Model overrides. Leave unset to use the selected architecture preset.
p.add_argument("--block-size", type=int, default=None)
p.add_argument("--n-layers", type=int, default=None)
p.add_argument("--d-model", type=int, default=None)
p.add_argument("--n-heads", type=int, default=None)
p.add_argument("--dropout", type=float, default=None)
p.add_argument("--n-experts", type=int, default=None)
p.add_argument("--active-experts", type=int, default=None, help="Top-k experts selected per token")
p.add_argument("--expert-hidden-size", type=int, default=None)
p.add_argument("--aux-loss-coef", type=float, default=None)
p.add_argument("--router-noise-std", type=float, default=None)
p.add_argument("--router-z-loss-coef", type=float, default=None)
# Phased expert training. This restricts the router to a moving subset of
# experts and only keeps optimizer state for that subset.
p.add_argument("--train-experts-per-phase", type=int, default=0, help="0/all disables phased expert training")
p.add_argument("--expert-phase-steps", type=int, default=500)
p.add_argument("--expert-phase-stride", type=int, default=None, help="Default: half of --train-experts-per-phase")
return p.parse_args()
def build_config(args: argparse.Namespace, tokenizer) -> MoEConfig:
if args.resume_from is not None and (args.resume_from / "config.json").exists():
with (args.resume_from / "config.json").open() as f:
config = MoEConfig.from_dict(json.load(f))
print(f"Loaded model config from {args.resume_from / 'config.json'}")
return config
values = MoEConfig().to_dict()
values.update(ARCH_PRESETS[args.arch])
values["vocab_size"] = len(tokenizer)
overrides = {
"block_size": args.block_size,
"n_layers": args.n_layers,
"d_model": args.d_model,
"n_heads": args.n_heads,
"dropout": args.dropout,
"n_experts": args.n_experts,
"active_experts": args.active_experts,
"expert_hidden_size": args.expert_hidden_size,
"aux_loss_coef": args.aux_loss_coef,
"router_noise_std": args.router_noise_std,
"router_z_loss_coef": args.router_z_loss_coef,
}
for key, value in overrides.items():
if value is not None:
values[key] = value
return MoEConfig(**values)
def get_phase_active_expert_ids(
step: int,
n_experts: int,
train_experts_per_phase: int,
phase_steps: int,
phase_stride: int | None,
) -> list[int] | None:
if train_experts_per_phase <= 0 or train_experts_per_phase >= n_experts:
return None
if phase_steps <= 0:
raise ValueError("expert_phase_steps must be > 0")
stride = phase_stride if phase_stride is not None else max(1, train_experts_per_phase // 2)
phase = max(0, step - 1) // phase_steps
start = (phase * stride) % n_experts
return [(start + i) % n_experts for i in range(train_experts_per_phase)]
def set_expert_trainability(model: SimpleMoELanguageModel, active_expert_ids: Sequence[int] | None) -> None:
active = None if active_expert_ids is None else set(active_expert_ids)
for block in model.blocks:
for expert_id, expert in enumerate(block.moe.experts):
requires_grad = active is None or expert_id in active
for param in expert.parameters():
param.requires_grad = requires_grad
def expert_parameters(model: SimpleMoELanguageModel, active_expert_ids: Sequence[int] | None = None) -> list[torch.nn.Parameter]:
active = None if active_expert_ids is None else set(active_expert_ids)
params: list[torch.nn.Parameter] = []
for block in model.blocks:
for expert_id, expert in enumerate(block.moe.experts):
if active is None or expert_id in active:
params.extend(p for p in expert.parameters() if p.requires_grad)
return params
def shared_parameters(model: SimpleMoELanguageModel) -> list[torch.nn.Parameter]:
expert_param_ids = {id(param) for block in model.blocks for expert in block.moe.experts for param in expert.parameters()}
return [param for param in model.parameters() if id(param) not in expert_param_ids and param.requires_grad]
def trainable_parameters(model: SimpleMoELanguageModel) -> list[torch.nn.Parameter]:
return [param for param in model.parameters() if param.requires_grad]
def make_optimizer(params: Sequence[torch.nn.Parameter], args: argparse.Namespace) -> torch.optim.Optimizer | None:
params = [param for param in params if param.requires_grad]
if not params:
return None
return torch.optim.AdamW(params, lr=args.lr, weight_decay=args.weight_decay, betas=(0.9, 0.95))
def set_optimizer_lr(optimizer: torch.optim.Optimizer | None, lr: float) -> None:
if optimizer is None:
return
for group in optimizer.param_groups:
group["lr"] = lr
def zero_optimizer(optimizer: torch.optim.Optimizer | None) -> None:
if optimizer is not None:
optimizer.zero_grad(set_to_none=True)
def step_optimizer(optimizer: torch.optim.Optimizer | None) -> None:
if optimizer is not None:
optimizer.step()
def summarize_router_stats(router_stats: list[dict[str, torch.Tensor]]) -> dict[str, float] | None:
if not router_stats:
return None
entropies: list[float] = []
loads: list[torch.Tensor] = []
for stats in router_stats:
entropies.append(float(stats["entropy"]))
eligible = stats["eligible_expert_ids"].long()
loads.append(stats["load"].index_select(0, eligible))
load_values = torch.cat(loads)
return {
"entropy": sum(entropies) / len(entropies),
"load_min": float(load_values.min()),
"load_max": float(load_values.max()),
"dead": float((load_values == 0).sum()),
"eligible": float(load_values.numel()),
}
def make_trainer_state(
step: int,
use_phased_experts: bool,
optimizer: torch.optim.Optimizer | None,
shared_optimizer: torch.optim.Optimizer | None,
expert_optimizer: torch.optim.Optimizer | None,
active_expert_ids: Sequence[int] | None,
) -> dict:
if use_phased_experts:
return {
"step": step,
"optimizers": {
"shared": None if shared_optimizer is None else shared_optimizer.state_dict(),
"experts": None if expert_optimizer is None else expert_optimizer.state_dict(),
},
"active_expert_ids": None if active_expert_ids is None else list(active_expert_ids),
}
return {"step": step, "optimizer": None if optimizer is None else optimizer.state_dict()}
def main() -> None:
args = parse_args()
set_seed(args.seed)
device = get_device(args.device)
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
config = build_config(args, tokenizer)
if args.train_experts_per_phase > 0 and args.train_experts_per_phase < config.active_experts:
raise ValueError("train_experts_per_phase must be >= active_experts")
if config.active_experts > config.n_experts:
raise ValueError("active_experts must be <= n_experts")
raw_model = SimpleMoELanguageModel(config).to(device)
trainer_state = None
start_step = 0
if args.resume_from is not None:
model_path = args.resume_from / "model.pt"
trainer_state_path = args.resume_from / "trainer_state.pt"
if not model_path.exists():
raise FileNotFoundError(f"Missing checkpoint model: {model_path}")
print(f"Resuming from {args.resume_from}")
raw_model.load_state_dict(torch.load(model_path, map_location=device))
if trainer_state_path.exists():
trainer_state = torch.load(trainer_state_path, map_location=device)
start_step = int(trainer_state.get("step", 0))
print(f"Loaded trainer state at step {start_step}")
else:
print("No trainer_state.pt found; resuming model weights only")
expert_counts = raw_model.expert_parameter_counts()
expert_param_count = expert_counts[0][0] if expert_counts and expert_counts[0] else 0
print(f"Device: {device}")
print(f"Architecture: {config.arch}")
print(f"Total parameters: {sum(p.numel() for p in raw_model.parameters()):,}")
print(f"MoE: {config.n_experts} experts, top-{config.active_experts} active per token")
print(f"Expert parameters per expert: {expert_param_count:,}")
print(f"Expert parameters per MoE layer: {expert_counts}")
if config.d_model == 500 and config.expert_hidden_size == 500:
assert all(count == 500_000 for layer in expert_counts for count in layer)
if config.d_model == 500 and config.expert_hidden_size == 1000:
assert all(count == 1_000_000 for layer in expert_counts for count in layer)
use_phased_experts = 0 < args.train_experts_per_phase < config.n_experts
current_active_expert_ids = get_phase_active_expert_ids(
start_step + 1,
config.n_experts,
args.train_experts_per_phase,
args.expert_phase_steps,
args.expert_phase_stride,
)
set_expert_trainability(raw_model, current_active_expert_ids if use_phased_experts else None)
if use_phased_experts:
print(
f"Phased expert training: {args.train_experts_per_phase}/{config.n_experts} experts per phase, "
f"phase_steps={args.expert_phase_steps}, active_expert_ids={current_active_expert_ids}"
)
if args.compile:
print("Note: torch.compile may recompile when the expert phase changes.")
dataset = FineWebTokenDataset(
tokenizer=tokenizer,
block_size=config.block_size,
dataset_name=args.dataset_name,
dataset_config=args.dataset_config,
split=args.split,
shuffle_buffer=args.shuffle_buffer,
seed=args.seed,
text_column=args.text_column,
)
loader = DataLoader(dataset, batch_size=args.batch_size, num_workers=args.num_workers)
data_iter = iter(loader)
optimizer: torch.optim.Optimizer | None = None
shared_optimizer: torch.optim.Optimizer | None = None
expert_optimizer: torch.optim.Optimizer | None = None
if use_phased_experts:
shared_optimizer = make_optimizer(shared_parameters(raw_model), args)
expert_optimizer = make_optimizer(expert_parameters(raw_model, current_active_expert_ids), args)
else:
optimizer = make_optimizer(trainable_parameters(raw_model), args)
if trainer_state is not None:
if use_phased_experts and "optimizers" in trainer_state:
optimizers_state = trainer_state["optimizers"]
if shared_optimizer is not None and optimizers_state.get("shared") is not None:
shared_optimizer.load_state_dict(optimizers_state["shared"])
print("Loaded shared optimizer state")
saved_active = trainer_state.get("active_expert_ids")
if saved_active == current_active_expert_ids and expert_optimizer is not None and optimizers_state.get("experts") is not None:
expert_optimizer.load_state_dict(optimizers_state["experts"])
print("Loaded active expert optimizer state")
elif use_phased_experts:
print("Rebuilt active expert optimizer for current phase")
elif not use_phased_experts and optimizer is not None and trainer_state.get("optimizer") is not None:
optimizer.load_state_dict(trainer_state["optimizer"])
print("Loaded optimizer state")
else:
print("Optimizer state format does not match current training mode; starting optimizers fresh")
train_model = torch.compile(raw_model) if args.compile else raw_model # type: ignore[assignment]
use_amp = args.amp and device.type == "cuda"
train_model.train()
zero_optimizer(optimizer)
zero_optimizer(shared_optimizer)
zero_optimizer(expert_optimizer)
start = time.time()
running_loss = 0.0
running_lm_loss = 0.0
running_aux_loss = 0.0
running_load_balance_loss = 0.0
running_router_z_loss = 0.0
running_router_entropy = 0.0
router_stat_count = 0
last_router_summary: dict[str, float] | None = None
progress = tqdm(range(start_step + 1, args.max_steps + 1), desc="training", initial=start_step, total=args.max_steps)
for step in progress:
if use_phased_experts:
next_active_expert_ids = get_phase_active_expert_ids(
step,
config.n_experts,
args.train_experts_per_phase,
args.expert_phase_steps,
args.expert_phase_stride,
)
if next_active_expert_ids != current_active_expert_ids:
current_active_expert_ids = next_active_expert_ids
set_expert_trainability(raw_model, current_active_expert_ids)
expert_optimizer = make_optimizer(expert_parameters(raw_model, current_active_expert_ids), args)
zero_optimizer(shared_optimizer)
zero_optimizer(expert_optimizer)
progress.write(f"Expert phase changed at step {step}: active_expert_ids={current_active_expert_ids}")
lr_scale = min(1.0, step / max(1, args.warmup_steps))
current_lr = args.lr * lr_scale
set_optimizer_lr(optimizer, current_lr)
set_optimizer_lr(shared_optimizer, current_lr)
set_optimizer_lr(expert_optimizer, current_lr)
step_loss = 0.0
step_lm_loss = 0.0
step_aux_loss = 0.0
step_load_balance_loss = 0.0
step_router_z_loss = 0.0
collect_router_stats = step % args.log_every == 0
for _ in range(args.grad_accum_steps):
batch = next(data_iter)
input_ids = batch["input_ids"].to(device, non_blocking=True)
labels = batch["labels"].to(device, non_blocking=True)
if use_amp:
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
out = train_model(
input_ids,
labels=labels,
active_expert_ids=current_active_expert_ids if use_phased_experts else None,
collect_router_stats=collect_router_stats,
)
loss = out["loss"] / args.grad_accum_steps
else:
out = train_model(
input_ids,
labels=labels,
active_expert_ids=current_active_expert_ids if use_phased_experts else None,
collect_router_stats=collect_router_stats,
)
loss = out["loss"] / args.grad_accum_steps
loss.backward()
step_loss += float(loss.detach().cpu())
step_lm_loss += float((out["lm_loss"] / args.grad_accum_steps).detach().cpu())
step_aux_loss += float((out["aux_loss"] / args.grad_accum_steps).detach().cpu())
step_load_balance_loss += float((out["load_balance_loss"] / args.grad_accum_steps).detach().cpu())
step_router_z_loss += float((out["router_z_loss"] / args.grad_accum_steps).detach().cpu())
summary = summarize_router_stats(out.get("router_stats", [])) if collect_router_stats else None
if summary is not None:
running_router_entropy += summary["entropy"]
router_stat_count += 1
last_router_summary = summary
if args.max_grad_norm > 0:
torch.nn.utils.clip_grad_norm_(trainable_parameters(raw_model), args.max_grad_norm)
step_optimizer(optimizer)
step_optimizer(shared_optimizer)
step_optimizer(expert_optimizer)
zero_optimizer(optimizer)
zero_optimizer(shared_optimizer)
zero_optimizer(expert_optimizer)
running_loss += step_loss
running_lm_loss += step_lm_loss
running_aux_loss += step_aux_loss
running_load_balance_loss += step_load_balance_loss
running_router_z_loss += step_router_z_loss
if step % args.log_every == 0:
denom = args.log_every
elapsed = max(time.time() - start, 1e-9)
toks_per_sec = (step - start_step) * args.batch_size * args.grad_accum_steps * config.block_size / elapsed
avg_loss = running_loss / denom
avg_lm_loss = running_lm_loss / denom
avg_aux_loss = running_aux_loss / denom
avg_load_balance_loss = running_load_balance_loss / denom
avg_router_z_loss = running_router_z_loss / denom
postfix = {
"loss": f"{avg_loss:.3f}",
"ppl": f"{math.exp(min(avg_lm_loss, 20)):.1f}",
"lm": f"{avg_lm_loss:.3f}",
"aux": f"{avg_aux_loss:.3f}",
"lb": f"{avg_load_balance_loss:.3f}",
"z": f"{avg_router_z_loss:.4f}",
"tok_s": f"{toks_per_sec:.0f}",
}
if router_stat_count > 0 and last_router_summary is not None:
postfix.update(
{
"r_ent": f"{running_router_entropy / router_stat_count:.2f}",
"load": f"{last_router_summary['load_min']:.2f}-{last_router_summary['load_max']:.2f}",
"dead": f"{int(last_router_summary['dead'])}/{int(last_router_summary['eligible'])}",
}
)
progress.set_postfix(**postfix)
running_loss = 0.0
running_lm_loss = 0.0
running_aux_loss = 0.0
running_load_balance_loss = 0.0
running_router_z_loss = 0.0
running_router_entropy = 0.0
router_stat_count = 0
last_router_summary = None
if step % args.save_every == 0:
trainer_state_to_save = make_trainer_state(
step,
use_phased_experts,
optimizer,
shared_optimizer,
expert_optimizer,
current_active_expert_ids if use_phased_experts else None,
)
save_checkpoint(args.out_dir, raw_model, tokenizer, trainer_state_to_save, args)
trainer_state_to_save = make_trainer_state(
args.max_steps,
use_phased_experts,
optimizer,
shared_optimizer,
expert_optimizer,
current_active_expert_ids if use_phased_experts else None,
)
save_checkpoint(args.out_dir, raw_model, tokenizer, trainer_state_to_save, args)
print(f"Saved checkpoint to {args.out_dir}")
if __name__ == "__main__":
main()