Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 12 additions & 5 deletions .dockerignore
Original file line number Diff line number Diff line change
@@ -1,14 +1,21 @@
# Data directories - these will be accessed via volume mount
data/
results/
logs_to_keep/
# Generated data payloads are accessed via a volume mount. Keep the root
# compatibility modules and the packaged speedrunning_plms.data source code.
/data/*
!/data/
!/data/*.py
/results/
/logs_to_keep/

# Cache directories
.cache/
.pytest_cache/
__pycache__/
*.pyc
*.pyo
*.pyd
*.egg-info/
build/
dist/

# Git files
.git/
Expand Down Expand Up @@ -45,4 +52,4 @@ logs/

# Temporary files
tmp/
temp/
temp/
12 changes: 7 additions & 5 deletions Dockerfile
Original file line number Diff line number Diff line change
Expand Up @@ -34,14 +34,16 @@ WORKDIR /app
COPY requirements.txt .

RUN pip install --upgrade pip setuptools && \
pip install -r requirements.txt -U && \
pip install --force-reinstall torch torchvision --index-url https://download.pytorch.org/whl/cu128 -U && \
pip install numpy==1.26.4

pip install torch torchvision --index-url https://download.pytorch.org/whl/cu128 -U && \
pip install -r requirements.txt

# 5️⃣ Copy the rest of the source
COPY . .

# Install the package and its repository test tooling. Runtime dependencies
# were installed above, so this also validates the package metadata in-image.
RUN pip install -e ".[test]"

# 6️⃣ Change working directory to where the volume will be mounted
WORKDIR /workspace

Expand Down Expand Up @@ -69,4 +71,4 @@ RUN mkdir -p \
VOLUME ["/workspace"]

# 8️⃣ Default command – override in `docker run … python train.py`
CMD ["bash"]
CMD ["bash"]
6 changes: 6 additions & 0 deletions MANIFEST.in
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
include LICENSE
include README.md
include requirements.txt
recursive-include example_yamls *.yaml
recursive-include evaluation *.py *.json
recursive-include tests *.py
61 changes: 55 additions & 6 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -166,6 +166,43 @@ flowchart TB

## Getting Started

### Python package

Install the reusable model, data, optimizer, and training modules from the
repository root:

```bash
python -m pip install .
```

Install optional experiment tracking and launch dependencies with:

```bash
python -m pip install ".[training]"
```

Install benchmark dependencies with `python -m pip install ".[evaluation]"`.

Models saved with `save_pretrained()` include the custom model code and
canonical Transformers `AutoClass` metadata. A local checkpoint can be loaded
directly with `PLM.from_pretrained(path)`. A model repository can be loaded as
custom code after inspecting its source and pinning an immutable revision:

```python
from transformers import AutoModelForMaskedLM

model = AutoModelForMaskedLM.from_pretrained(
"organization/model-name",
trust_remote_code=True,
revision="full-hub-commit-sha",
code_revision="full-hub-commit-sha",
)
```

The model follows the standard masked-language-model interface. Batched
`input_ids`, `attention_mask`, and optional `labels` return a
`MaskedLMOutput` with `loss` and `logits`.

### Quick Start

On many popular HPC platforms will be missing Python headers `Python.h` which break `torch.compile`. To fix this, run the following code:
Expand Down Expand Up @@ -217,11 +254,17 @@ sudo docker run --gpus all --shm-size=128g -v ${PWD}:/workspace speedrun_plm \
torchrun --standalone --nproc_per_node=NUM_GPUS_ON_YOUR_SYSTEM train.py
```

Some key arguments for `train.py` include
Some key arguments for `train.py` include:

`--hf_token YOUR_HUGGINGFACE_TOKEN`, a Huggingface write token is required to save your models to Huggingface hub
`--wandb_token YOUR_WANDB_TOKEN`, is required for Weights and Biases (WANDB) logging
`--yaml_path YOUR_YAML_FILE`, points to an experimental set up with more settings. See `example_yamls/default.yaml` for inspiration
- `--push_to_hub --hf_model_name ORGANIZATION/MODEL` explicitly enables final model publication. Publication is disabled by default.
- `--hf_token YOUR_HUGGINGFACE_TOKEN` authenticates an opted-in Hub publication.
- `--wandb_token YOUR_WANDB_TOKEN` enables Weights and Biases logging.
- `--yaml_path YOUR_YAML_FILE` points to an experiment configuration. See `example_yamls/default.yaml`.

When publication is enabled, training uploads one complete artifact containing
weights, configuration, remote code, and its runtime requirements only after
training and final evaluation succeed. No code-only artifact is uploaded at
startup.

See [Command-line Argument](#command-line-arguments) for the full list of argument.

Expand Down Expand Up @@ -271,7 +314,7 @@ This script will automatically:
| Argument | Type | Default | Description |
|----------|------|---------|-------------|
| `--yaml_path` | str | None | Path to YAML file with experiment configuration. CLI arguments override YAML. |
| `--hf_token` | str | None | HuggingFace token (required for model saving/uploading). Prompted if not provided. |
| `--hf_token` | str | None | Hugging Face token for an explicitly enabled publication. |
| `--wandb_token` | str | None | Weights & Biases API token (for experiment tracking). Prompted if not provided. |
| `--log_name` | str | None | Name for the log file and wandb run. If not set, a random UUID is used. |
| `--bugfix` | flag | False | Use small batch size and max length for debugging. |
Expand Down Expand Up @@ -314,7 +357,8 @@ This script will automatically:
| `--lr_hidden` | float | 0.05 | Learning rate for hidden layers (Muon). |
| `--muon_momentum_warmup_steps` | int | 300 | Steps for Muon momentum warmup (0.85 → 0.95). |
| `--eval_every` | int | 1000 | Evaluate on validation set every N steps. |
| `--hf_model_name` | str | "lhallee/speedrun" | HuggingFace model name for saving. |
| `--push_to_hub` | flag | False | Publish one complete final model artifact after successful training and evaluation. |
| `--hf_model_name` | str | None | Hugging Face destination repository used with `--push_to_hub`. |
| `--save_every` | int | None | Save checkpoint every N steps (if set). |
| `--num_workers` | int | 4 | Number of workers for optimized dataloader. |
| `--prefetch_factor` | int | 2 | Prefetch factor for optimized dataloader. |
Expand All @@ -323,6 +367,11 @@ This script will automatically:

## Performance Benchmarks

`evaluation/benchmark_esm.py` loads every model, remote-code module,
tokenizer, and dataset from the full commit SHA recorded in
`evaluation/benchmark_manifest.json`. Update that manifest intentionally when
changing benchmark inputs so result provenance remains reproducible.

### Recommended Configuration

Batch sizes of 8×64×1024 (524,288) or 4×64×1024 (262,144) tokens have demonstrated excellent performance. We recommend a local batch size of 64×1024 (65,536) tokens for 80GB VRAM systems, with adjustments for smaller configurations.
Expand Down
8 changes: 8 additions & 0 deletions data/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
import sys
from pathlib import Path

_SRC = Path(__file__).resolve().parents[1] / "src"
if str(_SRC) not in sys.path:
sys.path.insert(0, str(_SRC))

from speedrunning_plms.data import * # noqa: F401,F403
41 changes: 15 additions & 26 deletions data/create_og90_splits.py
Original file line number Diff line number Diff line change
@@ -1,33 +1,22 @@
import argparse
from datasets import load_dataset, DatasetDict
import sys
from pathlib import Path

parser = argparse.ArgumentParser()
parser.add_argument('--hf_token', type=str, default=None)
_SRC = Path(__file__).resolve().parents[1] / "src"
if str(_SRC) not in sys.path:
sys.path.insert(0, str(_SRC))

args = parser.parse_args()
from speedrunning_plms.data.splits import build_og_prot90_splits, login_if_token, push_splits

if args.hf_token:
import huggingface_hub
huggingface_hub.login(token=args.hf_token)

data = load_dataset('tattabio/OG_prot90', split='train').remove_columns('id').shuffle(seed=11)
#data = data.cast_column('sequence', Value(dtype='string'))
print(data)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--hf_token", type=str, default=None)
args = parser.parse_args()
login_if_token(args.hf_token)
data = build_og_prot90_splits()
push_splits(data, "Synthyra/og_prot90")

data = data.train_test_split(test_size=20000, seed=22)

train = data['train']
valid = data['test']
valid = valid.train_test_split(test_size=10000, seed=33)
test = valid['test']
valid = valid['train']

data = DatasetDict({
'train': train,
'valid': valid,
'test': test
})

print(data)

data.push_to_hub('Synthyra/og_prot90')
if __name__ == "__main__":
main()
41 changes: 15 additions & 26 deletions data/create_omgprot50_splits.py
Original file line number Diff line number Diff line change
@@ -1,33 +1,22 @@
import argparse
from datasets import load_dataset, DatasetDict
import sys
from pathlib import Path

parser = argparse.ArgumentParser()
parser.add_argument('--hf_token', type=str, default=None)
_SRC = Path(__file__).resolve().parents[1] / "src"
if str(_SRC) not in sys.path:
sys.path.insert(0, str(_SRC))

args = parser.parse_args()
from speedrunning_plms.data.splits import build_omg_prot50_splits, login_if_token, push_splits

if args.hf_token:
import huggingface_hub
huggingface_hub.login(token=args.hf_token)

data = load_dataset('tattabio/OMG_prot50', split='train').remove_columns('id').shuffle(seed=11)
#data = data.cast_column('sequence', Value(dtype='string'))
print(data)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--hf_token", type=str, default=None)
args = parser.parse_args()
login_if_token(args.hf_token)
data = build_omg_prot50_splits()
push_splits(data, "Synthyra/omg_prot50")

data = data.train_test_split(test_size=20000, seed=22)

train = data['train']
valid = data['test']
valid = valid.train_test_split(test_size=10000, seed=33)
test = valid['test']
valid = valid['train']

data = DatasetDict({
'train': train,
'valid': valid,
'test': test
})

print(data)

data.push_to_hub('Synthyra/omg_prot50')
if __name__ == "__main__":
main()
43 changes: 15 additions & 28 deletions data/create_uniref50_splits.py
Original file line number Diff line number Diff line change
@@ -1,35 +1,22 @@
import argparse
from datasets import load_dataset, DatasetDict, concatenate_datasets
import sys
from pathlib import Path

parser = argparse.ArgumentParser()
parser.add_argument('--hf_token', type=str, default=None)
_SRC = Path(__file__).resolve().parents[1] / "src"
if str(_SRC) not in sys.path:
sys.path.insert(0, str(_SRC))

args = parser.parse_args()
from speedrunning_plms.data.splits import build_uniref50_splits, login_if_token, push_splits

if args.hf_token:
import huggingface_hub
huggingface_hub.login(token=args.hf_token)

data = load_dataset('agemagician/uniref50_09012025').remove_columns('id').remove_columns('name').shuffle(seed=11)
data = data.rename_column('text', 'sequence')
print(data)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--hf_token", type=str, default=None)
args = parser.parse_args()
login_if_token(args.hf_token)
data = build_uniref50_splits()
push_splits(data, "Synthyra/uniref50")

data = concatenate_datasets([data['train'], data['validation'], data['test']])

data = data.train_test_split(test_size=20000, seed=22)

train = data['train']
valid = data['test']
valid = valid.train_test_split(test_size=10000, seed=33)
test = valid['test']
valid = valid['train']

data = DatasetDict({
'train': train,
'valid': valid,
'test': test
})

print(data)

data.push_to_hub('Synthyra/uniref50')
if __name__ == "__main__":
main()
Loading