Supervised fine-tuning is the cheapest stage of post-training. It is also the stage where GPU time is most often wasted, because SFT data does not look like pretraining data. Pretraining streams a near-endless corpus cut into full-length blocks, so every slot in every batch holds a real token. SFT works on a few thousand to a few million conversations whose lengths vary by two orders of magnitude, and it computes the loss on only some of their tokens. The defaults that suit pretraining therefore leave most of an SFT batch empty.

This article is about the GPU side of an SFT step: where the time and memory go, and how to stop paying for padding. It covers the length distribution, batching strategies, padding-free attention, the logits tensor and token-weighted loss across ranks. The method itself (chat templates, masking prompt tokens, teaching the model to stop) is covered in SFT, in depth, and the per-parameter memory budget is in full fine-tuning in depth. Library details below were checked against the TRL v1.14.1 documentation in October 2026. TRL changes defaults between releases, so pin the version you run.

What an SFT step spends compute on

A training step costs roughly six floating-point operations per parameter per token that passes through the model: two in the forward pass and four in the backward pass. Attention adds a term that grows with the square of sequence length. The rule that matters is that the GPU does this work for every slot in the batch tensor, including padding. A pad token runs through every linear layer just like a real token does. The loss mask zeroes its contribution to the gradient, but it does not give back the compute already spent on it.

So the throughput metric for SFT is real tokens per second, the tokens that are not padding, along with the share of those that carry a loss. A run that logs only samples per second cannot tell a fast configuration from one that is mostly padding.

Measure the length distribution first

Measure the length distribution before choosing a batching strategy. Tokenise the dataset with the exact chat template you will train with and record the length of each example. To make the effect concrete, we ran a small simulation: 20,000 synthetic examples with lognormal lengths (median 450 tokens, capped at 4,096), which is a common shape for chat data. The sample came out with a median of 454 tokens, a p90 of 1,439, a p99 of 3,585 and a mean of 669. These are modelled numbers, not a real dataset, but real instruction sets usually show the same long right tail.

With random batches of eight padded to the longest member, only 36.8 percent of slots held real tokens. Nearly two thirds of the linear-layer compute went to padding. Attention was worse: summed over batches, padded attention did 4.70 times the work of attending over each example's own length. Sorting examples by length inside windows of 512 and then batching raised the real-token share to 96.0 percent. Best-fit-decreasing packing into 4,096-token rows filled more than 99.9 percent of slots (3,267 rows for 13.4 million tokens).

StrategyReal-token share (model)What it costs you
Pad to longest, random batches36.8%Most compute spent on padding; memory peaks follow the p99 example
Length-grouped batches96.0%Batches are no longer random; long batches arrive together and spike memory
Packing, padding-freeover 99.9%Needs boundary-aware attention; changes how many examples make up a step

Packing without cross-contamination

Packing concatenates several examples into one fixed-length row. Done naively, it is a correctness bug: a causal mask over the whole row lets each example attend to the ones before it. The fix is to tell the attention kernel where each example starts and ends.

Variable-length attention kernels, such as the varlen entry points in FlashAttention, take the batch as one flat sequence plus a vector of cumulative sequence boundaries, usually called cu_seqlens. Each query then attends only to keys inside its own example, and no compute is spent on cross-example pairs. In the Hugging Face stack the boundaries are carried by position_ids that restart at zero for each example. That also gives every example correct rotary positions, so a packed example sees the same positions it would see alone. The underlying kernel is explained in FlashAttention in depth.

TRL exposes this as two options. padding_free=True flattens each batch into one sequence with no padding, and the docs state it is supported only with FlashAttention 2 or 3. packing=True groups examples into rows of max_length tokens. With the default packing_strategy="bfd", which is best-fit decreasing, padding-free mode is switched on regardless of the flag. BFD truncates any example longer than a row. "bfd_split" splits such examples across rows instead, and "wrapped" fills rows aggressively and cuts examples mid-sequence. For SFT, where a cut response teaches the model to stop mid-thought, use BFD. Set max_length above your p99, then count how many examples it truncates.

Same eight examples, three ways to put them on the GPUPad to longest (36.8% real tokens in the length model)Length-grouped batches (96.0% real)Packed, padding-free: one flat row, boundaries passed to the attention kernelcu_seqlens = [0, a, a+b, ...] position_ids restart at 0 for every exampleVarlen attentioneach token sees only its own exampleLinear layerscost tracks real tokensLoss on target tokenschunked lm_head + CE
Padding, length grouping and padding-free packing for the same examples. Bar widths are illustrative; the percentages come from the length model in the text.

The logits tensor and chunked loss

The last matrix multiply in a language model maps each hidden state to a score for every vocabulary entry. Modern vocabularies are large. Take a 4,096-wide model with a 128,256-entry vocabulary and a micro-batch of 32,768 tokens (eight rows of 4,096). The logits alone are 8.4 GB in bfloat16. A loss that upcasts to float32 adds another 16.8 GB, and the backward pass needs a gradient of the same shape. The float32 copy alone is larger than the roughly 16 GB of bfloat16 weights of an 8B model, and it is why SFT runs often fail with out-of-memory errors at the loss, not inside the layers.

Two observations shrink it. First, most positions have label -100 (prompt and padding), and their logits are never needed. If 35 percent of tokens are response tokens, computing the projection only for those cuts the bfloat16 logits from 8.4 GB to 2.9 GB. Second, cross-entropy can be computed over chunks of tokens, so the full tensor never exists at once. TRL v1.14.1 does both by default: loss_type="chunked_nll" drops ignored positions before the lm_head matmul and processes the loss in chunks. A fused alternative is the Liger Kernel's linear cross-entropy, enabled with use_liger_kernel=True. The two are not compatible. With Liger on, TRL resolves the loss to plain "nll" and lets the fused kernel do the work. The idea in plain PyTorch:

def chunked_response_loss(hidden, lm_head_weight, labels, chunk=4096):
    """hidden: [T, d] flat tokens; labels: [T] with -100 for ignored positions.
    Returns the summed loss and the number of target tokens (normalise later)."""
    # shift: position t predicts token t+1
    h, y = hidden[:-1], labels[1:]
    keep = y != -100
    h, y = h[keep], y[keep]                 # project only positions that carry a loss
    total = h.new_zeros((), dtype=torch.float32)
    for i in range(0, h.shape[0], chunk):
        logits = (h[i:i + chunk] @ lm_head_weight.T).float()   # [chunk, V] lives briefly
        total = total + torch.nn.functional.cross_entropy(logits, y[i:i + chunk], reduction="sum")
    return total, keep.sum()

Under plain autograd this loop still keeps each chunk's logits for the backward pass. Production kernels compute the gradient inside the loop or recompute the logits, so treat this as an explanation, not a drop-in. On packed batches, ignore the label at each example's first position so the shift never pairs two examples.

Weighting the loss by tokens across ranks

Once batches vary in how many target tokens they carry, the loss has to be weighted by tokens. Averaging per micro-batch and then averaging those means makes a token in a short batch count more than a token in a long one. The correct objective is the sum of token losses over the whole optimizer step divided by the number of target tokens in that step, across every accumulation micro-batch and every data-parallel rank. The method-level argument is in the SFT method article. The GPU-level consequence is that you need the global token count before you scale the first micro-batch's gradient, which means one extra all-reduce per step:

for step_batches in loader.optimizer_steps(accum=ga_steps):      # list of micro-batches
    local = sum((b["labels"][..., 1:] != -100).sum() for b in step_batches)
    global_tokens = local.clone()
    torch.distributed.all_reduce(global_tokens)                  # one small collective per step
    for b in step_batches:
        loss_sum, _ = model_loss_sum(model, b)                   # summed, not averaged
        # DDP/FSDP average gradients over world_size, so multiply it back in
        (loss_sum * world_size / global_tokens).backward()
    clip_grad_norm_(model.parameters(), 1.0)
    optimizer.step(); scheduler.step(); optimizer.zero_grad(set_to_none=True)

In TRL and the Transformers Trainer, average_tokens_across_devices defaults to True and performs the cross-rank part of this. If you write your own loop, or a custom compute_loss_func, you own it. A quick check: train the same data with gradient accumulation 1 and 8 at equal global batch size. The loss curves should nearly overlap. If they separate, the normalisation is per micro-batch.

Memory: what batching controls

Weights, gradients and optimizer state scale with parameters and are set by your choice of full fine-tuning, LoRA or QLoRA. Activations scale with tokens per micro-batch, and that is the part SFT batching controls. Three settings matter most. Gradient checkpointing recomputes layer activations in the backward pass instead of storing them. It defaults to True in SFTConfig (unlike the plain Trainer) and costs roughly one extra forward pass. activation_offloading=True moves saved activations to CPU memory, trading PCIe bandwidth for HBM. The token budget per micro-batch sets the peak.

With padding, the peak is set by the longest batch, which arrives at an unpredictable step, so a run can pass a 200-step smoke test and fail at step 3,000. With packing, every micro-batch has the same token count, so the peak at step one is the peak for the whole run. For the full memory budget and sharding choices, see the full fine-tuning budget.

A TRL configuration that applies it

A configuration that applies all of the above with TRL, using only parameters documented for v1.14.1. The model id and dataset are placeholders. attn_implementation is passed through to from_pretrained.

import torch
from trl import SFTConfig, SFTTrainer
from datasets import load_dataset

args = SFTConfig(
    output_dir="out/support-sft",
    model_init_kwargs={"dtype": torch.bfloat16, "attn_implementation": "flash_attention_2"},
    max_length=4096,                 # above the measured p99; log how many examples BFD truncates
    packing=True,                    # packing_strategy defaults to "bfd" -> padding-free on
    assistant_only_loss=True,        # needs {% generation %} markers in the chat template
    per_device_train_batch_size=4,   # packed rows of 4,096 tokens per GPU
    gradient_accumulation_steps=4,
    learning_rate=2e-5,
    num_train_epochs=2,
    gradient_checkpointing=True,     # already the SFTConfig default; stated for the run record
    logging_steps=10,
    include_num_input_tokens_seen="non_padding",   # count real tokens, not padded slots
)
trainer = SFTTrainer(model="your-org/base-8b", args=args,
                     train_dataset=load_dataset("json", data_files="train.jsonl", split="train"))
trainer.train()

With packing, per_device_train_batch_size counts packed rows, not examples. The number of examples per optimizer step now varies, and the number of optimizer steps per epoch drops. When you compare against an unpacked baseline, compare at equal tokens seen, not equal steps. Learning-rate schedules tuned in steps need the same translation.

Worked example: reading throughput honestly

Suppose an 8B model on eight GPUs logs 3,000 real tokens per second per GPU after packing. The training work is about 6 × 8×109 = 4.8×1010 FLOPs per token, so each GPU is sustaining roughly 1.4×1014 FLOP/s, or 144 TFLOP/s, before attention and recomputation. Divide by the dense BF16 peak on your GPU's datasheet to get model FLOPs utilisation. These throughput figures are hypothetical. The method is the point: compute utilisation from real tokens, not from slots.

Switch the same run to random padded batches. At the same slots per second, only 36.8 percent are real, so real throughput falls to about 1,100 tokens per second and the epoch takes roughly 2.7 times as long, before counting the extra attention work.

Failure modes

  • Cross-example attention. Packing with an attention implementation that ignores boundaries. Loss looks fine and the model quietly learns from neighbouring examples. Check the attention backend the run actually loaded, and run a two-example test where changing the first example must not change the second one's loss.
  • Positions not reset. A custom collator that concatenates rows without restarting position_ids gives later examples positions they never see at inference.
  • Silent truncation. BFD truncates over-length examples, and max_length defaults to 1,024 in v1.14.1. Long responses lose their endings, and with them the end-of-turn token.
  • Loss that is not comparable. Switching batching or normalisation changes the reported loss. Compare configurations on a fixed eval set.
  • Logit OOM. A custom loss that materialises full float32 logits undoes chunked loss. Profile peak memory around the loss call.

Trade-offs

ChoiceGainCost
Packing (BFD) + padding-freeNear-full slots, flat memory peakFlashAttention 2/3 dependency; step semantics change
Length grouping, no packingMost of the gain with any attention backendCorrelated batches; memory spikes on long groups
Chunked lossRemoves the full logits tensorA little extra kernel overhead
Fused Liger lossFused projection and lossAnother dependency; not combinable with chunked_nll
Gradient checkpointingActivation memory falls sharplyAbout one extra forward pass of compute

What to do next

  1. Tokenise your dataset with the exact training template and log median, p90, p99 and maximum length.
  2. Log real tokens per second and target tokens per step, not only samples per second.
  3. Turn on packing with BFD and FlashAttention 2 or 3. Run the two-example isolation test before the real run.
  4. Set max_length above p99 and count how many examples are truncated.
  5. Keep the chunked loss default unless you deliberately switch to Liger. Profile peak memory at the loss.
  6. Check token-weighted normalisation: losses at accumulation 1 and 8 should overlap.
  7. Compare against your baseline at equal tokens seen, and pick the checkpoint on a held-out eval. Continue with the TRL deep dive for the trainer side.
Key takeaway: An SFT step pays for every slot in the batch, padding included, so measure your length distribution before choosing how to batch. In the length model here, random padded batches were only 36.8 percent real tokens, length grouping reached 96.0 percent and BFD packing over 99.9 percent. Pack only with boundary-aware attention, keep the logits from materialising, and weight the loss by target tokens across the whole step.