AAIE-Distilled MoE

A 617.91M-parameter (222.73M active per token) Mixture-of-Experts transformer trained on general web text via knowledge distillation against a Qwen/Qwen2.5-3B teacher. This is the raw base checkpoint β€” it has not seen any instruction-tuning or task-specific fine-tuning, so it behaves like a classic base language model: it continues text plausibly, but will not reliably follow instructions phrased as questions or requests.

⚠️ This checkpoint is from an unfinished, deliberately-stopped run β€” read "Training status" below before using this for anything beyond experimentation. For a fully-converged model from this same project, see the companion AAIE-Distilled Dense model instead (dense architecture, same distillation recipe, trained to completion).

Training status β€” read this first

Training was stopped by request at step 41,860 of a 100,000-step target (41.9%) β€” this is not a converged model. It was stopped cleanly (not a crash), and validation loss/router health were both stable at the time of stopping, but the model had already been in a noisy, flat validation-perplexity plateau (roughly 17–23 ppl) since around step 20,000, with no clear further improvement through the final checkpoint. See this project's MOE_TRAINING_REPORT.md (or the copy in the training repo root) for the full writeup, including why this plateau is the main open question for anyone continuing this run, and a matched-step comparison against the dense variant showing the two architectures are close enough to be within each other's noise band at this point in training β€” no confident "MoE beats/loses to dense" claim should be drawn from the current checkpoints.

Also note: this run's teacher (Qwen/Qwen2.5-3B) is larger than the dense sibling model's teacher (Qwen/Qwen2.5-1.5B) β€” any comparison between the two architectures also carries this confound and isn't a clean architecture-only comparison.

Architecture

Same GQA-attention backbone (8 query heads / 2 KV heads) + RoPE + RMSNorm + tied embeddings as the dense sibling model, but the feed-forward layer is replaced with an 8-expert, top-2-routed Mixture-of-Experts FFN: each token is scored by a router against all 8 experts, sent to its top-2, and those experts' outputs are combined by (renormalized) router weight. Only the selected experts run per token, so compute per token scales with top-2, not all 8.

Value
Layers / d_model 20 / 512
Experts 8 total, top-2 routed, expert d_ff=2144 each
Total parameters 617.91M
Active parameters per token 222.73M (36% of total β€” fewer than the 354M-parameter dense sibling model)
Tokenizer same as Qwen/Qwen2.5-0.5B (151,936 vocab)

Two auxiliary losses (load-balancing + router z-loss, Zoph et al. 2022 ST-MoE) plus router-input jitter were used during training to keep routing balanced β€” the router never collapsed over the full training run (zero collapse warnings, max_expert_frac stayed in the healthy 0.18–0.33 range against a 0.5 warning threshold the entire time). Router jitter is a training-only regularizer and is not part of this exported inference module.

Data

Dataset Size Role
HuggingFaceFW/fineweb-edu (sample-10BT subset) ~5.48B tokens seen (41,500 steps Γ— 131,072 tokens/step) of a 9.9B-token train split General web-text language modeling
Qwen/Qwen2.5-3B (frozen teacher) β€” Distillation signal: loss = 0.5 * CE(labels) + 0.5 * KD(teacher logits, T=2.0)

Training settings

Value
Base random initialization
Optimizer AdamW, weight decay 0.1
LR schedule warmup 1,000 steps β†’ 3e-4 peak, cosine decay toward 3e-5 by step 100,000 (schedule target not reached β€” training stopped early, see above)
Batch micro-batch 2 Γ— grad-accum 64 Γ— seq-len 1024 = 131,072 tokens/step
Steps completed 41,500 (checkpointed) of 100,000 target
Grad clip 1.0
Sequence length 1024
Hardware single A100 (40GB), SLURM cluster
MoE aux losses load-balancing coef 0.1, router z-loss coef 0.001, router jitter 0.02 (train-time only)

Usage

This is a base model β€” prompt it as a text continuation, not an instruction:

from transformers import AutoModelForCausalLM, AutoTokenizer

repo = "your-username/aaie-moe-pretrain"  # after pushing, see push_to_hub.py
tokenizer = AutoTokenizer.from_pretrained(repo)
model = AutoModelForCausalLM.from_pretrained(repo, trust_remote_code=True, device_map="auto")

inputs = tokenizer("The purpose of a database index is", return_tensors="pt").to(model.device)
out = model.generate(**inputs, max_new_tokens=60, use_cache=True, do_sample=True, temperature=0.8)
print(tokenizer.decode(out[0, inputs["input_ids"].shape[1]:], skip_special_tokens=True))

trust_remote_code=True is required (custom architecture, see modeling_aaiemoe.py). use_cache=True (the default) enables real KV-caching for .generate() β€” see GQAAttentionCached in modeling_aaiemoe.py. The MoE routing forward pass (MoEFFN) runs per-token top-2 expert selection at every call, same as during training (minus the training-only jitter/aux-loss terms).

Known limitations

  • Not converged: stopped at 41.9% of its planned training budget, in the middle of an unresolved validation-perplexity plateau. Expect noisier, less reliable outputs than a fully-trained model of comparable size.
  • Not instruction-tuned: asking it a question or giving it an instruction will often just produce a plausible-sounding continuation of the prompt text rather than an answer/response β€” this is expected base-model behavior on top of being additionally undertrained.
  • Qualitative spot-check (LLM-judge, base checkpoint, pre-instruction-tuning) scored this model's outputs lower than the dense sibling's (1.50/10 vs 3.05/10 mean, on a 20-example IT/CS assignment-feedback rubric) β€” but that comparison used a dense checkpoint trained to full completion against this checkpoint mid-training, so it mostly reflects relative training progress rather than an architecture verdict. See MOE_TRAINING_REPORT.md Β§5.4 for the full caveat.
  • Router health metrics (no collapse, balanced expert usage) were monitored throughout training and were fine at every checkpoint β€” if you fine-tune this further and see degraded outputs, routing collapse is unlikely to be the cause based on this run's history, but worth re-checking regardless (log max_expert_frac/entropy per forward pass if you do).
Downloads last month
310
Safetensors
Model size
0.6B params
Tensor type
F32
Β·
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support