How to Perform Supervised Fine-Tuning (SFT) on a Pretrained LLM: A Complete Implementation Guide
Supervised Fine-Tuning (SFT) in the train-llm-from-scratch repository is implemented as a two-step pipeline that first packs instruction datasets into fixed-length sequences with prompt masking, then trains the model using a masked cross-entropy loss that only computes gradients on assistant responses.
Supervised Fine-Tuning transforms a base language model into a helpful assistant by training it on instruction-following datasets. In the FareedKhan-dev/train-llm-from-scratch project, this process is broken into a data preparation stage that converts public datasets into packed HDF5 format, followed by a distributed training loop that optimizes a prompt-masked loss function.
Dataset Preparation and Packing
The SFT workflow begins with converting raw instruction datasets into a format suitable for language model training. The repository supports multiple public datasets including Alpaca, Dolly, and GSM8K, each transformed into a standardized chat-style message format.
Converting and Tokenizing Instruction Data
In scripts/prepare_sft_data.py, source-specific conversion functions—gsm8k_to_messages, alpaca_to_messages, and dolly_to_messages—transform each dataset's schema into a list of messages. These messages are then encoded using the model's chat template via encode_chat, which applies the appropriate formatting and tokenization.
The conversion pipeline handles the following steps:
- Loads datasets from HuggingFace or local sources
- Applies dataset-specific logic to extract instruction-response pairs
- Formats them using the chat template to produce tokenized sequences
- Packs multiple examples into fixed-length rows to maximize training efficiency
Packing and Storage
The pack_examples function in src/post_training/sft.py takes the tokenized sequences and concatenates them into tensors of fixed length (specified by context_length). This function generates two critical outputs: tokens (the input IDs) and loss_mask (a binary mask indicating which tokens should participate in loss calculation).
The packed tensors are written to an HDF5 file (sft_packed.h5) via the write_packed function in scripts/prepare_sft_data.py. This format enables efficient random access during training and supports datasets larger than available RAM.
To prepare the SFT dataset, run:
PYTHONPATH=. HF_HOME=/ephemeral/hf_cache \
python scripts/prepare_sft_data.py \
--context_length 1024 \
--out_dir /ephemeral/data \
--dev_frac 0.02
The Masked SFT Loss Function
A key innovation in this implementation is the prompt-masked loss, which ensures the model only learns to generate assistant responses, not to reproduce the user prompts.
Loss Computation Logic
The sft_loss function in src/post_training/sft.py implements a next-token cross-entropy loss that respects the binary mask. Given input tokens and a corresponding loss mask, it:
- Shifts the tokens to create targets (next-token prediction)
- Applies the mask to zero out loss contributions from prompt tokens
- Computes the mean cross-entropy only over assistant tokens
- Returns the scalar loss value for backpropagation
This approach prevents the model from wasting capacity learning to predict user instructions, focusing all gradient updates on improving response quality.
Training Loop and Distributed Setup
The scripts/train_sft.py script orchestrates the full training process, handling distributed data parallel (DDP) setup, model loading, optimization, and checkpointing.
Data Loading and Batch Iteration
Training data is streamed via data_loader/sft_dataset.py using the get_sft_batch_iterator function. This iterator yields tuples of (tokens, loss_mask, epoch) and handles DDP sharding to ensure each process receives a unique slice of the data. The iterator efficiently reads from the HDF5 file created during preparation, supporting shuffling and batch collation.
Training Execution
The main training routine performs the following steps:
- Loads the pretrained checkpoint using
load_backbone_from_ckpt - Optionally applies Torch compilation for performance optimization
- Configures the optimizer with a cosine learning rate schedule
- Iterates through epochs, computing the masked loss for each batch
- Performs periodic evaluation on the held-out dev set using
eval_dev - Saves checkpoints via
save_stage_ckptat specified intervals
Run distributed training across multiple GPUs:
PYTHONPATH=. torchrun --standalone --nproc_per_node=2 scripts/train_sft.py
For single-GPU training:
PYTHONPATH=. python scripts/train_sft.py
Configuration and Hyperparameters
All SFT hyperparameters are centralized in config/post_training_config.py within the SFTConfig class. This configuration specifies:
- Paths to the pretrained checkpoint and output directories
- Training batch size and context length
- Learning rate, warmup steps, and cosine schedule parameters
- Evaluation frequency and checkpointing intervals
- The fraction of data reserved for validation (
dev_frac)
Modifying these values in SFTConfig allows you to adjust the training dynamics without changing the source code, making experiments reproducible and configurable via Python dataclasses.
Summary
- Data Preparation: Convert instruction datasets (Alpaca, Dolly, GSM8K) into chat format using
scripts/prepare_sft_data.py, then pack them into fixed-length HDF5 sequences withpack_examples. - Masked Loss: The
sft_lossfunction insrc/post_training/sft.pycomputes cross-entropy only on assistant tokens, ignoring prompt tokens through the loss mask. - Efficient Loading: The
get_sft_batch_iteratorindata_loader/sft_dataset.pystreams packed data from HDF5 with DDP support. - Training Orchestration:
scripts/train_sft.pyhandles distributed training, evaluation, and checkpointing with cosine learning rate scheduling. - Configuration:
SFTConfiginconfig/post_training_config.pycentralizes all hyperparameters for reproducible experiments.
Frequently Asked Questions
What is the purpose of the loss mask in SFT?
The loss mask in src/post_training/sft.py ensures that the model only receives gradients for tokens belonging to the assistant's response, not the user's prompt. This focuses the model's learning capacity on generating helpful answers rather than memorizing questions, improving instruction-following performance.
Why use HDF5 format for the SFT dataset?
HDF5 provides memory-mapped access to large datasets, allowing data_loader/sft_dataset.py to efficiently stream batches without loading the entire dataset into RAM. This is essential when working with large instruction datasets that may exceed available memory, and it enables fast random access for shuffling during training.
How does the packing strategy improve training efficiency?
The pack_examples function concatenates multiple short instruction-response pairs into a single sequence of fixed length (e.g., 1024 tokens). This minimizes padding waste and increases GPU utilization by ensuring nearly every token in the batch participates in the forward and backward passes, rather than processing many short sequences with excessive padding.
Can I use custom datasets for SFT with this codebase?
Yes, you can extend scripts/prepare_sft_data.py by implementing a new conversion function that transforms your dataset's format into the message list structure (similar to alpaca_to_messages), then tokenize it using encode_chat. The rest of the pipeline—packing, loss computation, and training—remains identical regardless of the source dataset.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →