Skip to content

Latest commit

ย 

History

9 Commits

Folders and files

NameName
Last commit message
Last commit date
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 

Repository files navigation

OpenESM

GitHub HuggingFace

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 ยท ็ฎ€ไฝ“ไธญๆ–‡


๐ŸŽ‰ Updates

  • 2026-10-07 โ€” โœจโœจ Full codebase released.

๐Ÿš€ Quick Start

Install the GPU, Hugging Face, and web-chat dependencies once:

uv sync --extra gpu --extra hf --extra web
source .venv/bin/activate

Download 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 datasets

Run a complete training job by setting the token budget:

TARGET_TOTAL_TOKENS=100000000 bash runs/train.sh

Data 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.

๐Ÿ“ˆ Scaling Law

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 validation BPB versus training tokens
DCLM pretraining
DCLM validation BPB versus model size
Parameter scaling
Average zero-shot accuracy versus model size across training datasets
Zero-shot model-size scaling
Zero-shot accuracy versus training tokens
Zero-shot token scaling

Use configs/train.yaml as the default configuration and set TARGET_TOTAL_TOKENS to choose the training budget.

โš™๏ธ Pretrained Checkpoints

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.

๐Ÿ’ฌ Chat Demo

ESM Chat web interface

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 8000

Open 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.

Run

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-CHAT

Chat, 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 from Hugging Face

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.

Repository layout

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

Development

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