Skip to content
Merged
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
99 changes: 98 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ This repo supports two training backends:
## Table of Contents

- [Quick Start](#quick-start)
- [SFT on a Checkpoint with Singularity](#-sft-on-a-checkpoint-with-singularity)
- [Project Structure](#-project-structure)
- [Design Philosophy](#-design-philosophy)
- [Feature Guide](#-feature-guide)
Expand Down Expand Up @@ -102,6 +103,101 @@ For cluster environments, use the submission script. It auto-generates a SLURM b
python scripts/submit.py --config configs/trl/sft.yaml
```

For the full tokenize-then-train workflow in a container, see [SFT on a Checkpoint with Singularity](#-sft-on-a-checkpoint-with-singularity).

## 🚀 SFT on a Checkpoint with Singularity

This guide fine-tunes a given checkpoint with SFT on a SLURM cluster, with training inside a Singularity (or Apptainer) container. It takes two jobs, both submitted from the login node with the same config:

1. **[Tokenize the datasets](#step-1-tokenize-the-datasets)**: a `--tokenize-only` job on 1 GPU. It loads, filters, tokenizes, and packs the data, writes the result to the Hugging Face datasets cache, and exits.
2. **[Train](#step-2-train)**: the full job. It finds the processed data in the cache, skips preprocessing, and trains.

Tokenizing first keeps the multi-node allocation from sitting idle during CPU-bound preprocessing, and it surfaces data and chat-template problems in a small job.

The example config is [`configs/trl/prelude-sft.yaml`](configs/trl/prelude-sft.yaml). It fine-tunes a 9B checkpoint on LUMI, with the tokenizer from a separate repo. For your own run, copy it and replace the checkpoint, data, container paths, and SLURM account.

### Before you start

The login node needs only the base dependencies from [Installation](#installation), because the training stack lives in the container. Run every `submit.py` command from the repository root, inside that environment. Relative paths in the config (`container.env_file`, `paths.output_base`) resolve against the root, and `submit.py` copies the code from it.

#### Configure the container

The fields in `prelude-sft.yaml` that the container run depends on:

- **`container.image`**: the job runs `accelerate launch scripts/train.py` in this image through `singularity exec`. The image must hold the Python packages from `pyproject.toml`, with a PyTorch build for the cluster's GPUs. The `post_training` code does not come from the image; see `run_name` below.
- **`container.path`**: the job sets `PATH` inside the container to exactly this value, so it must contain the directory with `python` and `accelerate`. This image keeps them in `/opt/venv/bin`. The default is `/usr/local/bin:/usr/bin:/bin`.
- **`container.bind_mounts`**: Singularity `--bind` specs. Bind every host path the job reads or writes: the run directory under `paths.output_base`, the Hugging Face cache, and any local checkpoint or dataset. Bind each path as `src` alone, so it keeps the same path inside the container:
- `submit.py` resolves `paths.output_base` to its real path, following symlinks. Bind that real path; here, the repository under `/pfs/lustrep3/...`.
- The frozen config refers to the prefetched checkpoint and tokenizer by their host paths in the Hugging Face cache.
- **`container.env_file`**: a shell file that sets the Hugging Face cache; see [the next section](#write-the-env-file).
- **`run_name`**: at submission, `submit.py` copies `src/post_training/` and `scripts/` into the run directory, and the job runs that copy. With a fixed `run_name`, Step 2 reuses the copy from Step 1, so both jobs run the same transforms and chat templates. The Step 2 submission review warns that the frozen source "will NOT be replaced"; that is expected. To pick up a code change, delete `<run_dir>/src` and `<run_dir>/scripts`, then run Step 1 again.

#### Write the env file

The job sources `container.env_file` on the host before it starts the container, then passes the Hugging Face cache variables into the container. The repository ships `env/jupiter.env` as an example. Create one for your cluster, such as `env/lumi.env`:

```bash
export HF_HOME=/scratch/<project>/<user>/hf_cache
export HF_HUB_CACHE=$HF_HOME/hub
export HUGGINGFACE_HUB_CACHE=$HF_HOME/hub
export HF_DATASETS_CACHE=$HF_HOME/datasets
```

- Export `HF_HOME`, `HF_HUB_CACHE`, and `HUGGINGFACE_HUB_CACHE`. The job script runs with `set -u`, so a missing one stops it with `unbound variable`. `HF_DATASETS_CACHE` defaults to `$HF_HOME/datasets`.
- Use `export NAME=value` lines. `submit.py` reads these lines before it prefetches, so the login node downloads into the cache that the job reads.
- Keep `HF_HOME` inside a bind mount.

### Step 1: Tokenize the datasets

```bash
python scripts/submit.py --config configs/trl/prelude-sft.yaml --tokenize-only
```

On the login node, `submit.py`:

1. reads the Hugging Face cache variables from the env file,
2. downloads the checkpoint, tokenizer, and datasets into that cache (`prefetch_assets: true`, the default),
3. prints a submission review and asks for confirmation (`--confirm` skips it),
4. freezes the config and code into the run directory, and submits the job on 1 node with 1 GPU. The other `slurm` values (account, partition, CPUs, memory, wall time) stay as configured.

In the container, the job loads the tokenizer and chat template, then loads and filters the datasets. It builds the trainer, which loads the checkpoint, then tokenizes and packs the data. It prints one decoded sample and exits.

Preprocessing is CPU-bound. Keep `data.num_proc` and `sft.dataset_num_proc` at or below `slurm.cpus_per_task`, and give the job enough wall time. `slurm.*` overrides do not change the processed data, so Step 1 can use its own:

```bash
python scripts/submit.py --config configs/trl/prelude-sft.yaml --tokenize-only 'slurm.wall_time="08:00:00"'
```

> [!NOTE]
> Quote `slurm.wall_time` on the command line as shown. Unquoted, `24:00:00` parses as the integer `86400`, which SLURM reads as minutes.

Before Step 2, read `<run_dir>/slurm/slurm-<id>.out`:

- It shows the `Tokenized dataset preview` block and `--tokenize-only set — exiting after trainer initialization.` Check that the preview follows the chat template's format.
- A warning `... rows, ... with an all-zero assistant mask` means more than 1% of the rows were dropped. A warning that rows are cut `PART-WAY THROUGH their supervised span` means those rows train on truncated answers. Raise `sft.max_seq_length`, or set `sft.truncated_span_action: drop`, then run Step 1 again.
- A `ValueError` stops the job if the chat template lacks `{% generation %}` markers or if no row survives the filter.

### Step 2: Train

```bash
python scripts/submit.py --config configs/trl/prelude-sft.yaml
```

Use the same config and the same overrides as Step 1, except for `slurm.*`. `submit.py` renders `<run_dir>/slurm/job.sh` again without `--tokenize-only` and submits it on all nodes. In the container, each preprocessing stage finds its output in the datasets cache and loads it, and training starts. Before the wall time runs out, the job requeues itself and resumes from the latest checkpoint in `<run_dir>/checkpoints/`.

The cache is hit only when every input to the data pipeline is unchanged. Between the two steps, keep these identical:

| Keep identical | Why |
|---|---|
| `data.*` | datasets, weights, transforms, seed, and chat template |
| `sft.max_seq_length`, `sft.packing`, `sft.truncated_span_action` | row filtering, truncation, and packing |
| `model.name_or_path`, `model.revision`, `model.tokenizer_name_or_path`, `model.tokenizer_revision` | the tokenizer |
| `container.image` | the library versions that compute the cache keys |
| `container.env_file` | the cache location (`HF_DATASETS_CACHE`) |
| `run_name` | the frozen transforms and chat templates |

To confirm the cache hit, open the training job's `<run_dir>/slurm/slurm-<id>.err`: the `Tokenizing train dataset` and `Packing train dataset` progress bars must not appear. If they do, an input in the table changed, or `datasets` warned in Step 1 that a function `couldn't be hashed properly`. Either way, the training job processes the data again from scratch.

## 📂 Project Structure

```text
Expand Down Expand Up @@ -268,6 +364,7 @@ Templates that are safe for SFT today:
|------|--------|-------|
| `olmo3-instruct-sft` | `allenai/OLMo-3-7B-Instruct-SFT` (HF Hub) | Use to reproduce the Instruct-SFT recipe. |
| `olmo3-think-sft` | `allenai/Olmo-3-7B-Think-SFT` (HF Hub) | Use to reproduce the Think-SFT recipe. |
| `qwen3` | `Qwen/Qwen3-8B` (HF Hub) | Assistant turns whose `<think>` block the template strips stay out of the loss. |

Templates that are *not* safe for SFT (kept for inference / DPO compatibility):

Expand Down Expand Up @@ -333,7 +430,7 @@ You must specify exactly one determining factor for training duration in the `tr
- **Debug**: `debug.enabled: true`
Forces `report_to: none`, uses a separate output directory, and allows overwriting existing runs.
- **Tokenize only**: `--tokenize-only` (CLI flag on `train.py` / `submit.py`)
Exits immediately after the trainer is initialized — dataset loading, tokenization, and packing all run, but the training loop is never entered. Useful for pretokenizing the dataset before committing to a full run. When passed to `submit.py`, the job is automatically constrained to 1 node and 1 GPU.
Exits immediately after the trainer is initialized — dataset loading, tokenization, and packing all run, but the training loop is never entered. Useful for pretokenizing the dataset before committing to a full run. When passed to `submit.py`, the job is automatically constrained to 1 node and 1 GPU. See [SFT on a Checkpoint with Singularity](#-sft-on-a-checkpoint-with-singularity) for the full workflow.

```bash
python scripts/submit.py --config configs/trl/sft.yaml --tokenize-only
Expand Down
121 changes: 121 additions & 0 deletions configs/trl/prelude-sft.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,121 @@

method: sft
backend: trl
run_name: oellm-9b-256k-theta64m-prelude-anneal300b-sft # auto-generated from model + datasets if null
offline: false

# Container (remove or set image: null for bare-metal)
container:
image: /scratch/project_465002530/containers/post-training-rocm7.2.4-py3.12-torch2.9.1-trl1.7.0-olmo-patched.sif
bind_mounts:
- /pfs/lustrep3/scratch/project_465002530/users/krishnak/post-training/ # replace with your own path to the post-training repo
path: /opt/venv/bin:/opt/rocm/bin:/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin
env_file: env/lumi.env # make sure you have this file in the post-training repo, or change to your own env file

# -- Model -------------------------------------------------------------------
model:
name_or_path: birgermoell/oellm-9b-256k-theta64m-prelude-anneal300b
attn_implementation: flash_attention_2 # No flash attention 3 for rocm yet, so use flash attention 2 for now
dtype: bfloat16
tokenizer_name_or_path: openeurollm/tokenizer-256k
tokenizer_revision: qwen3-tokens # This tokenizer has dedicated tokens for qwen3 chat template, which is used in the data section below. If you use a different tokenizer, make sure to change the chat_template in the data section accordingly.

# -- Training hyper-parameters -----------------------------------------------
training:
num_train_epochs: 2
learning_rate: 8.0e-5
effective_batch_size: 32 # per_device * grad_accum * world_size
per_device_train_batch_size: 1
warmup_ratio: 0.03
adam_beta2: 0.95
lr_scheduler_type: "linear"
gradient_checkpointing: true
bf16: true
seed: 42
use_liger_kernel: true

# -- SFT method parameters ---------------------------------------------------
sft:
max_seq_length: 32768
packing: true
dataset_num_proc: 32

# -- Checkpointing -----------------------------------------------------------
checkpointing:
save_steps: 250 # How frequently to save the full checkpoints (with optimizer states)
save_total_limit: 3 # Full checkpoints to keep
inference_checkpoint_steps: 250 # Minimal inference model interval (set to null to disable)
inference_checkpoint_path: "inference_checkpoints" # Relative to run dir

# -- Data mix ----------------------------------------------------------------
data:
chat_template: qwen3 # Name from chat template registry
num_proc: 32 # null = auto-detect, capped at 32
datasets:
- name: "dolci-instruct-sft"
path: "allenai/Dolci-Instruct-SFT" # HuggingFace dataset path
split: "train"
weight: 1.0
transform: null # null = already conversational

# -- DeepSpeed ---------------------------------------------------------------
deepspeed:
bf16:
enabled: auto
zero_optimization:
stage: 2
overlap_comm: false # Enabling overlap_comm has caused issues before, so we disable it for now. If you want to enable it, set this to true and test carefully.
contiguous_gradients: true
reduce_scatter: true
gradient_clipping: 1.0
train_micro_batch_size_per_gpu: "auto"
gradient_accumulation_steps: "auto"
train_batch_size: "auto"
optimizer:
type: AdamW
params:
lr: "auto"
betas: "auto"
eps: "auto"
weight_decay: "auto"

# -- Accelerate launch flags (explicit multi-node control) -------------------
accelerate:
mixed_precision: "bf16"
use_deepspeed: true
deepspeed_multinode_launcher: "standard" # "standard" | "pdsh" | etc.
same_network: true # All nodes on same network
rdzv_backend: "static" # "static" | "c10d" | "etcd"
dynamo_backend: "inductor" # "inductor" | "no" | etc.

# -- Logging & tracking ------------------------------------------------------
logging:
report_to:
- "wandb"
- "tensorboard"
wandb_project: "sft-training"
logging_steps: 10
include_num_input_tokens_seen: "non_padding"

# -- SLURM -------------------------------------------------------------------
slurm: # LUMI Cluster SLURM configuration
account: "project_465002530"
partition: "standard-g" # Test on dev-g first
num_nodes: 4
gpus_per_node: 8
cpus_per_task: 56
wall_time: "36:00:00"
job_name: "Prelude-SFT"
signal_time_seconds: 300 # SIGUSR1 sent this many seconds before timeout to trigger self-healing
max_failures: 1 # Self-healing retry limit
mem: "256G"

# -- Debug mode --------------------------------------------------------------
debug:
enabled: false
override_existing: false

# -- Output paths -------------------------------------------------------------
paths:
output_base: "outputs"
debug_base: "outputs/debug"
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ dependencies = [
"wandb",
"huggingface-hub",
"psutil",
"datasets==5.0.0",
]

[project.optional-dependencies]
Expand All @@ -21,7 +22,6 @@ trl = [
"trl",
"deepspeed",
"transformers",
"datasets>4.5.0",
"accelerate",
"kernels",
"flash_attn",
Expand Down
Loading
Loading