# How to Perform Supervised Fine-Tuning (SFT) on a Pretrained LLM: A Complete Implementation Guide

> Learn to perform Supervised Fine-Tuning SFT on a pretrained LLM with this implementation guide. Discover prompt masking and masked cross-entropy loss for efficient training.

- Repository: [Fareed Khan/train-llm-from-scratch](https://github.com/FareedKhan-dev/train-llm-from-scratch)
- Tags: how-to-guide
- Published: 2026-06-11

---

**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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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:

```bash
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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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:

1. Shifts the tokens to create targets (next-token prediction)
2. Applies the mask to zero out loss contributions from prompt tokens
3. Computes the mean cross-entropy only over assistant tokens
4. 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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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_ckpt` at specified intervals

Run distributed training across multiple GPUs:

```bash
PYTHONPATH=. torchrun --standalone --nproc_per_node=2 scripts/train_sft.py

```

For single-GPU training:

```bash
PYTHONPATH=. python scripts/train_sft.py

```

## Configuration and Hyperparameters

All SFT hyperparameters are centralized in [`config/post_training_config.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/prepare_sft_data.py), then pack them into fixed-length HDF5 sequences with `pack_examples`.
- **Masked Loss**: The `sft_loss` function in [`src/post_training/sft.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/post_training/sft.py) computes cross-entropy only on assistant tokens, ignoring prompt tokens through the loss mask.
- **Efficient Loading**: The `get_sft_batch_iterator` in [`data_loader/sft_dataset.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/data_loader/sft_dataset.py) streams packed data from HDF5 with DDP support.
- **Training Orchestration**: [`scripts/train_sft.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_sft.py) handles distributed training, evaluation, and checkpointing with cosine learning rate scheduling.
- **Configuration**: `SFTConfig` in [`config/post_training_config.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/config/post_training_config.py) centralizes 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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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.