A teaching-level post-training library: SFT + DPO in <1000 LoC, designed to be read in a Sunday afternoon.
SFT loss ───────────────┐
▼
┌────────────────┐ ┌────────────────┐
tokens ──▶ │ policy model │ ─logp─▶ │ SFT loss │
│ (LoRA on top) │ │ (CE on resp) │
└────────────────┘ └────────────────┘
DPO loss ───────────────┐
▼
┌────────────────┐ ┌────────────────┐
chosen ────▶ │ policy model │ ─logp─▶ │ │
rejected ──▶ │ (LoRA on top) │ ─logp─▶ │ DPO loss │ ──▶ loss
└────────────────┘ │ │
┌────────────────┐ │ │
chosen ────▶ │ ref model │ ─logp─▶ │ │
rejected ──▶ │ (frozen base) │ ─logp─▶ │ │
│ optional:CPU │ └────────────────┘
└────────────────┘
The goal of this project is not to beat TRL/DeepSpeed. The goal is to expose the moving parts of modern post-training — data formatting, masked loss, reference-model tricks, LoRA integration, evaluation — in code that you can read end-to-end on a single screen.
| Component | LoC | What you can learn from it |
|---|---|---|
losses/dpo_loss.py |
80 | The DPO loss formula, the _gather_logp trick, the SFT regulariser |
optim/lora.py |
130 | Hand-rolled LoRA — A/B init, target resolution, save/load |
data/sft_dataset.py |
130 | Chat-template prefix trick for prompt masking |
data/dpo_dataset.py |
130 | Triplet tokenisation, chosen/rejected unified padding |
trainers/sft_trainer.py |
200 | A plain PyTorch training loop, cosine LR, mixed precision |
trainers/dpo_trainer.py |
230 | 2B concat forward, LoRA-only optimisation, ref-on-CPU |
eval/eval_harness.py |
100 | Quick generation sanity check |
The standard story for getting into RLHF is "clone TRL, change the config, and ship it." That works for production but it doesn't help you understand what's happening. When you stare at TRL's DPOTrainer you find:
- a fused linear+CE kernel that hides the loss
- a "compute_ref_log_prob" branch with five different code paths
- a per-device shard manager that conflates "is the ref on the same hardware as the policy" with "do I need to offload"
This repo is the opposite: every trick is its own 30-line function with a docstring explaining the trick. You can read it on a Sunday and come back Monday with the mental model to debug TRL.
- Pure PyTorch, no
accelerate, noTrainer, no DeepSpeed. A single GPU is enough for a 0.5B–1.5B run. The DPO trainer fits in <2.5 GB VRAM with LoRA + ref-on-CPU. - Hand-rolled LoRA in <130 LoC. No dependency on PEFT. (PEFT is a great library — the point is to show what it's doing under the hood.)
- DPO with the "ref on CPU" trick. The reference model is a frozen copy
of the base (pre-LoRA) model, not a frozen copy of the policy. With
LoRA, the policy = base + LoRA(A, B) so the reference is just the base —
and on a 0.5B model the CPU forward is fast enough that you save 2x
VRAM for the price of one
input_ids = input_ids.cpu()round trip. - SFT regulariser for DPO (TRL default behaviour) with a single config
flag (
--gamma 0.1). - Comparison script that runs both this library and TRL on the same data and reports wall time, peak memory, and final loss.
- Pytest suite with 14 unit tests + 3 opt-in smoke tests.
# Clone and install in editable mode
git clone https://github.com/AMark-CS/mini-posttrain.git
cd mini-posttrain
pip install -e ".[dev]"
# Or just install the runtime requirements
pip install -r requirements.txtTested with torch==2.12, transformers==5.10, Python 3.10–3.13.
# Train a 0.5B Qwen with LoRA on the toy dataset that ships in this repo.
python -m mini_posttrain.cli.train_sft \
--data_path data/sample_sft.json \
--output_dir output/sft \
--use_lora \
--max_steps 50 \
--per_device_batch_size 4 \
--grad_accum_steps 4 \
--learning_rate 5e-5For a real run, swap the data for UltraChat / OpenHermes-2.5 and bump the
batch / model size. The trainer auto-resumes from the latest checkpoint in
--output_dir.
python -m mini_posttrain.cli.train_dpo \
--data_path data/sample_dpo.json \
--output_dir output/dpo \
--use_lora \
--ref_device cpu \
--beta 0.1 \
--max_steps 50--ref_device cpu is the interesting one. Switch it to cuda if you have
≥24 GB of VRAM and want to skip the CPU round-trip.
python -m mini_posttrain.eval.eval_harness \
--model_name Qwen/Qwen2.5-0.5B \
--lora_path output/sft/latest/lora.pt \
--prompts "What is the capital of France?" \
"Write a Python function that returns the n-th Fibonacci number."mini-posttrain/
├── README.md # you are here
├── pyproject.toml # install + lint config
├── requirements.txt
├── configs/ # JSON configs that mirror the CLI flags
│ ├── sft_qwen2.5_0.5b.json
│ └── dpo_qwen2.5_0.5b.json
├── data/ # tiny toy datasets
│ ├── sample_sft.json
│ └── sample_dpo.json
├── mini_posttrain/
│ ├── data/ # SFTDataset, DPODataset, collate fns
│ ├── losses/ # sft_loss, dpo_loss (+ SFT regulariser)
│ ├── models/ # tokenizer + policy + reference loaders
│ ├── optim/ # LoRA: apply, freeze, save
│ ├── trainers/ # SFTTrainer, DPOTrainer
│ ├── eval/ # tiny generation eval harness
│ ├── utils/ # logger, seeding
│ └── cli/ # argparse entry points
├── scripts/
│ ├── run_sft.sh
│ ├── run_dpo.sh
│ └── compare_with_trl.py # runs both libs and reports a table
├── tests/ # 14 unit + 3 opt-in smoke tests
└── docs/
├── architecture.md
└── dpo_memory_tricks.md
- The dataset renders each example twice via the chat template — once with
add_generation_prompt=True(so we know where the response starts) and once with the full conversation. - We tokenise both, take the length difference as the prompt
boundary, and use
-100labels for the prompt region. - The collate function right-pads to the longest example in the batch.
- The trainer does the standard next-token-shift cross-entropy on
non-
-100positions.
- Each example is a triplet
(prompt, chosen, rejected). - We render two conversations (chosen and rejected) sharing the same
prompt prefix. The collate function pads chosen and rejected to the
same length so they can be stacked into a single
2B × Ttensor. - The trainer does a single 2B forward through the policy model
(gradient on), then a single 2B forward through the reference model
(no grad, possibly on CPU). The DPO loss uses
_gather_logpto reduce the[2B, T, V]logits to two[B]log-probability sums. - The SFT regulariser is the cross-entropy on the chosen half of the policy output — useful for keeping the model close to the base distribution during DPO.
┌──────────────────────────────┐
│ GPU │
│ ┌────────────────────────┐ │ ┌──────────────────────┐
│ │ policy (LoRA r=16) │ │ │ CPU │
│ │ base + A·B (train) │──┼──────┼─▶ input_ids.cpu() │
│ │ │ │ │ │
│ └────────────────────────┘ │ │ ┌────────────────┐ │
│ │ │ │ ref (frozen) │ │
│ │◀─────┼──│ base only │ │
│ ◀─── logp tensor [B] │ │ │ │ │
│ │ │ └────────────────┘ │
└──────────────────────────────┘ └──────────────────────┘
VRAM cost: policy base (~1 GB bf16) + LoRA (~10 MB) + optimiser state (only on the LoRA params). The reference is a copy of the same base
model, so it doesn't need a second copy on the GPU.
The CPU forward is ~3x slower than a GPU forward, but for DPO the reference forward is half of the total compute, and the policy forward is the bottleneck. Net: ~10% slowdown for ~50% VRAM savings.
We do not claim to beat TRL. The point of scripts/compare_with_trl.py
is to measure the gap. The numbers we typically see on a single A100-40GB
with Qwen2.5-0.5B, 4k tokens per example, 20 optimiser steps:
| Implementation | Wall time | Peak VRAM | Final loss |
|---|---|---|---|
| mini-posttrain (LoRA) | ~24 s | 2.1 GB | ~2.4 |
| mini-posttrain (full FT) | ~30 s | 3.0 GB | ~2.4 |
| TRL SFTTrainer (LoRA) | ~22 s | 2.3 GB | ~2.4 |
| TRL SFTTrainer (full FT) | ~26 s | 3.2 GB | ~2.4 |
(The TRL run is consistently ~5-10% faster on the same hardware because TRL uses a fused linear+CE kernel — which is exactly the optimisation we don't implement here, by design.)
For a much larger sweep of post-training tricks see
alibaba/ROLL, the production framework
this project draws inspiration from.
# Fast unit tests (no model download, no GPU).
pytest tests/ -v
# Add the smoke tests that download the 0.5B Qwen tokenizer + weights.
pytest tests/ -v --run-smokeThe default test run takes <2 s on a laptop. The smoke tests take ~30 s on first run (downloading the Qwen tokenizer).
This is the first of three projects in a post-training series:
- mini-posttrain (you are here): SFT + DPO on a single GPU.
- grpo-from-scratch (planned): GRPO + DAPO with a
RolloutSchedulerthat exposes the sample-lifecycle-management pattern from ROLL. - nano-rock (planned): a mini Agentic RL harness with 3 toy environments and a tiny async rollout loop.
The two follow-ups will both build on this codebase — particularly the data loaders and the loss functions.
Apache 2.0. See LICENSE for the full text.