Open-source code for training and scaling Energy-Steered Models
Keywords: Energy-Steered Models, language models, pretraining, scaling laws
๐ Updates โข ๐ Quick Start โข ๐ Scaling Law โข โ๏ธ Pretrained Checkpoints โข ๐ฌ Chat Demo
English ยท ็ฎไฝไธญๆ
- 2026-10-07 โ โจโจ Full codebase released.
Install the GPU, Hugging Face, and web-chat dependencies once:
uv sync --extra gpu --extra hf --extra web
source .venv/bin/activateDownload and tokenize the training data (run once before the first training job;
requires network). Prepares climbmix, dclm, and fineweb by default; use
--datasets to prepare only a subset:
bash runs/prepare_data.sh # default: climbmix + dclm + fineweb
bash runs/prepare_data.sh --datasets dclm # only the named datasetsRun a complete training job by setting the token budget:
TARGET_TOTAL_TOKENS=100000000 bash runs/train.shData preparation needs network access, so run it once on a networked cpu-worker. Training itself is offline, but must run where the data is ready.
On DCLM, validation BPB falls as training tokens and model size increase. Zero-shot accuracy scales with model size across ClimbMix, DCLM, and FineWeb, and improves with training tokens.
![]() DCLM pretraining |
![]() Parameter scaling |
![]() Zero-shot model-size scaling |
![]() Zero-shot token scaling |
Use configs/train.yaml as the default configuration and set
TARGET_TOTAL_TOKENS to choose the training budget.
Pretrained models are published in the
OpenESM Hugging Face collection.
The collection includes 160M, 520M, and 1B models trained on DCLM, FineWeb, and
ClimbMix; one 160M OWT model; and a FineWeb SFT model named
ESM-FineWeb-1B-CHAT.
Load a published model directly by its Hugging Face ID. The first call downloads and caches the weights and tokenizer:
import torch
from esm.modeling_esm import load_checkpoint
model_id = "guan-wang/ESM-FineWeb-1B-CHAT"
model, tokenizer, hparams, device = load_checkpoint(model_id, device="cuda")
tokens = tokenizer.encode("Hello, ESM.", append=tokenizer.get_bos_token_id())
input_ids = torch.tensor([tokens], dtype=torch.long, device=device)
with torch.no_grad():
logits = model(input_ids)
print(logits.shape)Replace the model ID with any repository listed in the collection.
Run the browser chat with the published SFT model:
bash runs/chat.sh --web \
--checkpoint guan-wang/ESM-FineWeb-1B-CHAT \
--host 0.0.0.0 \
--port 8000Open http://localhost:8000 in a browser on the server machine. If the server runs on a remote GPU node, use your cluster's port forwarding to expose port 8000 on your local machine.
The PyTorch CPU and CUDA indexes are configured in pyproject.toml. Choose
one extra per environment.
The shell launchers are deliberately thin. They always run from the project root, so relative paths are stable in local shells and distributed jobs.
# Pretraining: training length comes from TARGET_TOTAL_TOKENS.
TARGET_TOTAL_TOKENS=100000000 bash runs/train.sh
# Fine-tuning: initialize from a published Hugging Face model.
TARGET_TOTAL_TOKENS=100000000 bash runs/sft.sh \
--dataset_name esm_sft \
--finetuning_model_ckpt guan-wang/ESM-DCLM-1B
# Position-wise BPB evaluation.
MODEL_DIR="$(python -c 'from huggingface_hub import snapshot_download; print(snapshot_download(repo_id="guan-wang/ESM-DCLM-1B"))')"
bash runs/eval.sh \
--checkpoint guan-wang/ESM-DCLM-1B \
--dataset dclm \
--data-dir /path/to/data \
--tokenizer-path "${MODEL_DIR}" \
--output_root . \
--run_name eval-dclm
# QA evaluation.
bash runs/qa.sh \
--checkpoint guan-wang/ESM-DCLM-1B \
--tokenizer-path "${MODEL_DIR}" \
--eval-bundle /path/to/eval_bundle \
--output_root . \
--run_name qa-dclm
# Interactive generation.
bash runs/chat.sh --checkpoint guan-wang/ESM-FineWeb-1B-CHATChat, validation evaluation, QA evaluation, and SFT initialization accept
published Hugging Face model IDs. BPB evaluation also needs the local snapshot
directory for token_bytes.pt. Zero-shot evaluation currently expects a local
Lightning checkpoint.
Training reads configs/train.yaml; evaluation reads configs/eval.yaml.
Explicit command-line arguments take precedence over YAML values. Training
length is controlled by TARGET_TOTAL_TOKENS. The effective gradient
accumulation is derived from global_batch_size, device_batch_size, and the
distributed world size, so max_steps and accumulate_grad_batches are not
independent training controls.
Every run uses the same output layout:
outputs/<run-name>/
โโโ checkpoints/
โโโ logs/
โ โโโ wandb/
โโโ results/
Data and tokenizer assets are not included in the repository. Pass their locations through the configuration or command line. Supported pretraining and BPB datasets include DCLM, FineWeb, ClimbMix, and OWT.
load_checkpoint() accepts a Hugging Face model ID and downloads the
Transformers weights and tokenizer automatically. For example:
import torch
from esm.modeling_esm import load_checkpoint
model, tokenizer, hparams, device = load_checkpoint(
"guan-wang/ESM-FineWeb-1B", device="cuda"
)
tokens = tokenizer.encode("hello", append=tokenizer.get_bos_token_id())
input_ids = torch.tensor([tokens], dtype=torch.long, device=device)
with torch.no_grad():
logits = model(input_ids)
print(logits.shape)For BPB evaluation, download the same model repository locally and pass that
directory as --tokenizer-path; it contains the tokenizer byte table used by
the metric. The loader also supports local Transformers directories and
Lightning checkpoints produced by this training code.
openesm/
โโโ assets/ # Repository branding assets
โ โโโ openesm-github-title.png # README title banner
โโโ configs/ # Default run configurations
โ โโโ eval.yaml # Evaluation defaults
โ โโโ train.yaml # Pretraining defaults
โโโ esm/ # Model, data, and training library
โ โโโ __init__.py # Package initialization
โ โโโ common.py # Shared constants and utilities
โ โโโ config.py # Training and evaluation config parsing
โ โโโ configuration_esm.py # Transformers model configuration
โ โโโ core_eval.py # Core benchmark evaluation
โ โโโ dataloader.py # Streaming data loaders
โ โโโ dataset.py # Shared dataset utilities
โ โโโ dataset_sft.py # Supervised fine-tuning datasets
โ โโโ disk_aware_checkpoint.py # Checkpoint save/load utilities
โ โโโ logger.py # Training metric and artifact logging
โ โโโ metrics.py # Loss and BPB metrics
โ โโโ modeling_esm.py # ESM architecture, loading, and inference
โ โโโ optim.py # Muon and AdamW optimizers and schedules
โ โโโ pretrain_dataset.py # Pretraining data pipeline
โ โโโ tokenizer.py # ESM tokenizer and tokenizer loading
โ โโโ trainer.py # Lightning training and evaluation module
โโโ figs/ # README figures and demo images
โ โโโ chat.jpg # Web chat interface screenshot
โ โโโ dclm_best_val_bpb_vs_depth.png # DCLM parameter scaling plot
โ โโโ qa_acc_vs_model_size.png # Zero-shot accuracy by model size
โ โโโ scaling_pretrain_dclm_validation_bpb.png # DCLM validation BPB by tokens
โ โโโ zero_shot_acc_vs_training_tokens.png # Zero-shot accuracy by tokens
โโโ runs/ # Shell entry points for common workflows
โ โโโ chat.sh # Launch interactive chat
โ โโโ eval.sh # Run validation evaluation
โ โโโ prepare_data.sh # Prepare and tokenize datasets
โ โโโ qa.sh # Run task-based QA evaluation
โ โโโ sft.sh # Fine-tune from a checkpoint
โ โโโ train.sh # Pretrain an ESM model
โ โโโ zeroshot.sh # Run zero-shot benchmarks
โโโ scripts/ # Python command-line entry points
โ โโโ chat.py # Serve terminal or web chat
โ โโโ eval.py # Evaluate validation loss and BPB
โ โโโ export_hf.py # Export a checkpoint to Transformers format
โ โโโ prepare_data.py # Build training and evaluation data
โ โโโ qa.py # Evaluate QA task suites
โ โโโ sft.py # Run supervised fine-tuning
โ โโโ train.py # Run pretraining
โ โโโ zeroshot.py # Run zero-shot benchmark evaluation
โ โโโ zeroshot_datasets.py # Define zero-shot benchmark datasets
โโโ tasks/ # Dataset/task adapters for evaluation
โ โโโ common.py # Shared task interfaces and helpers
โ โโโ customjson.py # Load custom JSON evaluation tasks
โ โโโ gsm8k.py # GSM8K math reasoning task
โ โโโ mmlu.py # MMLU multiple-choice task
โ โโโ smoltalk.py # SmolTalk conversation task
โ โโโ spellingbee.py # SpellingBee word-generation task
โโโ tests/ # Lightweight unit and compatibility tests
โ โโโ test_config_and_boundaries.py # Validate config and data boundaries
โ โโโ test_dataloader_lightweight_state.py # Check lightweight loader state
โ โโโ test_dataset_sft_mask.py # Check SFT loss masking
โ โโโ test_hf_configuration.py # Check Transformers configuration
โ โโโ test_modeling_esm_rope.py # Check rotary embeddings
โโโ .gitignore # Excludes local data, outputs, and caches
โโโ pyproject.toml # Package metadata and dependency groups
โโโ README.md # English project guide
โโโ README_zh.md # Simplified Chinese project guide
โโโ uv.lock # Locked dependency versions
uv sync --extra cpu --group dev
pytest -q
python -m scripts.train --help
python -m scripts.eval --help
python -m scripts.qa --help
python -m scripts.chat --help




