Text Generation
Transformers
Safetensors
mixture-of-experts
looped-transformer

How to Loop MoE: Flatten the Experts, Untie the Attention

Weights of the four 100B-token models from the paper How to Loop MoE: Flatten the Experts, Untie the Attention.

Each model is in its own folder:

Folder Looped block Experts per looped layer
Foil-1 1 layer, looped 16 times 64
Foil-2 2 layers, looped 8 times 32
Foil-3 4 layers, looped 4 times 16
Base 8 layers, looped 2 times 8

All four models have:

  • 563.2M parameters;
  • one MoE layer before the loop and one after it, each with 8 experts;
  • top-2 routing, hidden size 1024, 16 attention heads;
  • separate attention weights for every pass through the loop (untied attention);
  • 18 layers in total once the loop is unrolled;
  • context length 4096;
  • the SmolLM2 tokenizer (vocabulary 49,152).

They were trained on 100B tokens of FineWeb-Edu (sample-100BT).

Usage

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

repo, model_name = "ShourenWSR/how-to-loop-moe", "Foil-1"   # or Foil-2, Foil-3, Base
tokenizer = AutoTokenizer.from_pretrained(repo, subfolder=model_name, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(
    repo, subfolder=model_name, trust_remote_code=True, dtype=torch.bfloat16
).eval()

ids = tokenizer("The capital of France is", return_tensors="pt").input_ids
print(tokenizer.decode(model.generate(ids, max_new_tokens=20, do_sample=False)[0]))

trust_remote_code=True is required: the architecture is defined in modeling_loop_lm.py. The same file is at the repository root and in each folder; with subfolder=, transformers loads the code from the root.

Labels are not shifted inside the model. When you pass labels to forward, they must already be the next-token targets: labels[..., t] == input_ids[..., t + 1], with the last position set to -100. Passing labels=input_ids, as with most Hugging Face causal language models, gives a meaningless loss. The returned loss also includes the router auxiliary losses. To get the language-modelling loss alone, compute it from logits:

logits = model(input_ids=ids).logits.float()
loss = torch.nn.functional.cross_entropy(logits[0, :-1], ids[0, 1:])

No padding masks. The model does not use attention_mask, and it raises an error if the mask contains padding. Run one sequence at a time, or pad batches on the right and omit attention_mask; right padding does not change the outputs for the real tokens. Batched generation with left padding is not supported.

Intermediate checkpoints

Intermediate checkpoints of each run are in the per-run repositories: Foil-1, Foil-2, Foil-3, Base. The weights here are the final step of those runs (step 254,313).

License

Apache-2.0. The modeling code is ported from the release accompanying Sparse Layers are Critical to Scaling Looped Language Models (arXiv:2605.09165); see NOTICE.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Dataset used to train ShourenWSR/how-to-loop-moe

Papers for ShourenWSR/how-to-loop-moe