commit 285bbd7aacf425cb72f48d87d42554c0b47abe51 Author: Owen Qwen Date: Sun May 31 23:30:56 2026 -0500 Inital commit - includes v1 and v2 diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..4ccc42e --- /dev/null +++ b/.gitignore @@ -0,0 +1,8 @@ +__pycache__/ +*.py[cod] +.venv/ +.env +checkpoints/ +runs/ +wandb/ +.DS_Store diff --git a/README.md b/README.md new file mode 100644 index 0000000..7935b2d --- /dev/null +++ b/README.md @@ -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 diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..f041f75 --- /dev/null +++ b/pyproject.toml @@ -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"] diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..5fcd91b --- /dev/null +++ b/requirements.txt @@ -0,0 +1,4 @@ +torch>=2.2 +transformers>=4.40 +datasets>=2.19 +tqdm>=4.66 diff --git a/src/verysimplemoe.egg-info/PKG-INFO b/src/verysimplemoe.egg-info/PKG-INFO new file mode 100644 index 0000000..d60bed5 --- /dev/null +++ b/src/verysimplemoe.egg-info/PKG-INFO @@ -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 diff --git a/src/verysimplemoe.egg-info/SOURCES.txt b/src/verysimplemoe.egg-info/SOURCES.txt new file mode 100644 index 0000000..be3d4b2 --- /dev/null +++ b/src/verysimplemoe.egg-info/SOURCES.txt @@ -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 \ No newline at end of file diff --git a/src/verysimplemoe.egg-info/dependency_links.txt b/src/verysimplemoe.egg-info/dependency_links.txt new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/src/verysimplemoe.egg-info/dependency_links.txt @@ -0,0 +1 @@ + diff --git a/src/verysimplemoe.egg-info/entry_points.txt b/src/verysimplemoe.egg-info/entry_points.txt new file mode 100644 index 0000000..dbc820e --- /dev/null +++ b/src/verysimplemoe.egg-info/entry_points.txt @@ -0,0 +1,3 @@ +[console_scripts] +verysimplemoe-generate = verysimplemoe.generate:main +verysimplemoe-train = verysimplemoe.train:main diff --git a/src/verysimplemoe.egg-info/requires.txt b/src/verysimplemoe.egg-info/requires.txt new file mode 100644 index 0000000..5fcd91b --- /dev/null +++ b/src/verysimplemoe.egg-info/requires.txt @@ -0,0 +1,4 @@ +torch>=2.2 +transformers>=4.40 +datasets>=2.19 +tqdm>=4.66 diff --git a/src/verysimplemoe.egg-info/top_level.txt b/src/verysimplemoe.egg-info/top_level.txt new file mode 100644 index 0000000..0ae12d6 --- /dev/null +++ b/src/verysimplemoe.egg-info/top_level.txt @@ -0,0 +1 @@ +verysimplemoe diff --git a/src/verysimplemoe/__init__.py b/src/verysimplemoe/__init__.py new file mode 100644 index 0000000..9a538f9 --- /dev/null +++ b/src/verysimplemoe/__init__.py @@ -0,0 +1,3 @@ +from .model import ARCH_PRESETS, MoEConfig, SimpleMoELanguageModel + +__all__ = ["ARCH_PRESETS", "MoEConfig", "SimpleMoELanguageModel"] diff --git a/src/verysimplemoe/generate.py b/src/verysimplemoe/generate.py new file mode 100644 index 0000000..3635c41 --- /dev/null +++ b/src/verysimplemoe/generate.py @@ -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() diff --git a/src/verysimplemoe/model.py b/src/verysimplemoe/model.py new file mode 100644 index 0000000..09699a6 --- /dev/null +++ b/src/verysimplemoe/model.py @@ -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 diff --git a/src/verysimplemoe/train.py b/src/verysimplemoe/train.py new file mode 100644 index 0000000..08cc7a1 --- /dev/null +++ b/src/verysimplemoe/train.py @@ -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()