Personalize this lesson
Adapt explanations and teaching visuals to your background and preferred voice.
In this lab's fictional access-policy dataset, incident thread T102 asks whether an 18-day-old service-account key is still usable. The authored policy requires keys older than 14 days to undergo security review, beginning with a rotation ticket. A candidate assistant replies, "The key is old, but rotation can wait until next month." The answer skips the required escalation. That plausible failure is the regression this pipeline exists to catch; filing the ticket doesn't itself approve access.
The chat-template foundation showed how structured conversations turn into token sequences. The synthetic-data foundation showed how candidate rows get filtered and verified. Supervised fine-tuning (SFT) uses accepted demonstrations to increase the likelihood of desired responses through gradient updates.[1][2] InstructGPT used supervised demonstrations before reward-model training and reinforcement learning.[3] SFT doesn't guarantee an exact answer or procedural compliance on a new case.
Training loss can drop smoothly while policy behavior stays broken. Prompt-dominated metrics, another turn from T102 leaking into validation, unintended cross-example attention, or selection by training loss can hide the failure. Each raises a different diagnostic question. Check the split, labels, and evaluation before funding an optimizer sweep; none of these symptoms alone proves the cause.
Count supervised answer tokens before launching any job. A packed row can fill every token slot in the context window while providing zero learning signal if aggressive truncation cuts off the completion and masks every remaining token to -100. Reject that row outright instead of pretending its context tokens represent training data.

Turn behavior into a run contract
Before picking a trainer, map out which boundary might let the T102 regression slip through: the learning objective, the data split, the label mask, the token budget, the checkpoint format, or the release gate. The examples below verify those contracts on CPU. They don't train or benchmark a large language model; the TRL configuration is an integration blueprint, not an unverified GPU script.
Define the behavioral contract first: cases with keys older than 14 days require a rotation-ticket escalation under this lab's policy. Compare prompting with the policy in context, retrieval, and SFT before assuming new weights are needed. If unfamiliar domain language remains a barrier, investigate continued pretraining too. The observed failure suggests candidate interventions; it doesn't establish a required training stage.
Next, trace what each update actually sees. Which examples and tokens contribute to the loss function? What example and supervised-token budget reaches a single optimizer step? What held-out metric dictates checkpoint selection? Finally, can the validated recipe run on a single GPU, or do Fully Sharded Data Parallel (FSDP) or Zero Redundancy Optimizer (ZeRO) become necessary?
Following this order keeps the trainer where it belongs: as an execution engine applying an update. The split, mask, batch budget, resume bundle, and export gate define what that update means.
Separate objective from parameterization
Suppose T102 outputs fluent prose but persistently advises waiting until next month. Is the missing ingredient raw domain text, a demonstration of procedural compliance, or a ranking across several plausible replies? Answering that question separates the learning objective from the parameterization that absorbs the gradient updates.
An objective determines the mathematical signal guiding the model. A parameterization specifies which weights in the model architecture are allowed to change. They aren't interchangeable.
| Failure you observe | Objective to investigate | Why |
|---|---|---|
| Base model has weak exposure to domain vocabulary in unlabeled corpora | continued pretraining | next-token training on domain text adapts the language distribution before behavior training[4] |
| Model can read an access policy but doesn't follow the desired escalation procedure | SFT | prompt-completion demonstrations directly teach that response behavior[1] |
| Several acceptable answers need ranking by preference | preference training; compare against the current model or an SFT baseline | comparisons express relative preference rather than a single target answer[5] |
Once the objective is SFT, make a second decision: which part of the model can carry the new behavior?
| SFT parameterization | What moves during the same supervised objective | When to test it |
|---|---|---|
| Full fine-tuning | all trainable model weights | when memory permits and adapters may be too restrictive |
| LoRA | small low-rank adapter matrices; base weights stay frozen | when iteration speed, memory, or many task variants matter[6] |
| QLoRA | LoRA adapters while the frozen base is stored in 4-bit form | when the base model doesn't fit comfortably at higher precision[7] |
LoRA and QLoRA aren't alternatives to SFT. They're parameterizations usable for an SFT run. Continued pretraining can also use adapters. Both stages can use causal next-token cross-entropy; their training distributions and supervised targets differ. Calling every adapter run SFT would hide that distinction.
Full weights or adapters
Full fine-tuning
For our response-only objective, full fine-tuning updates all eligible model parameters. Test it when weights, gradients, optimizer state, activations, and temporary buffers fit, or when an adapter baseline appears capacity-limited. More trainable parameters permit a broader update, but don't establish better held-out behavior. Full fine-tuning doesn't require response-only labels by definition; that is a separate objective choice.
LoRA
LoRA freezes selected base weights and trains low-rank adapter matrices. Target modules and rank determine the trainable count. The original paper's GPT-3 175B comparison reported up to 10,000 times fewer trainable parameters and three times lower GPU memory requirements, not universal 99% and 70% savings.[6] Activations and frozen weights still consume memory. Measure the actual run; adapters are useful candidates when many variants share one base or full updates are too costly.
QLoRA
QLoRA reduces frozen-base storage while retaining adapter updates. Its paper uses 4-bit NormalFloat (NF4), quantized scaling constants, and paged optimizers. The base weights are dequantized to a computation dtype such as BF16 for forward and backward operations; adapter gradients don't mean the stored 4-bit base receives weight updates.[7] Check your implementation's storage, compute, and adapter dtypes separately. Compare held-out behavior under measured memory and time budgets.
Reject broken demonstrations before training
For this pipeline, store a prompt, a desired completion, and a group key for leakage-resistant splitting on every SFT row. A missing completion isn't harmless metadata: it's a training example with no answer to teach.
Before you run the validator, predict its decision for each row. T102 has an answer and a group key. B550 has an empty answer. The last row has an answer but no group key. Only one row should reach tokenization.
1rows = [
2 {"thread_id": "T102", "prompt": "Stale key", "completion": "Open a rotation ticket; keys older than 14 days need review."},
3 {"thread_id": "B550", "prompt": "Expired key", "completion": ""},
4 {"prompt": "Privileged role change", "completion": "Escalate privileged-role changes."},
5]
6
7required = {"thread_id", "prompt", "completion"}
8accepted, rejected = [], []
9for index, row in enumerate(rows):
10 missing = sorted(required - row.keys())
11 if missing:
12 rejected.append((index, f"missing {missing}"))
13 elif any(not isinstance(row[key], str) or not row[key].strip() for key in required):
14 rejected.append((index, "required fields must be nonempty strings"))
15 else:
16 accepted.append(row)
17
18print("accepted_threads=", [row["thread_id"] for row in accepted])
19print("rejected=", rejected)1accepted_threads= ['T102']
2rejected= [(1, 'required fields must be nonempty strings'), (2, "missing ['thread_id']")]The output makes the contract concrete: T102 is accepted, B550 is rejected before it can become a prompt-only example, and the last row is rejected before it can contaminate a grouped split. Next, keep accepted rows from the same deployment unit together.
Split by the unit that could leak
Before formatting rows, reserve evaluation examples that training can't imitate through near duplicates. For a policy assistant, several messages from one access incident or one policy document often share facts and phrasing. Randomly splitting individual messages can put one part of the same case in training and another in evaluation.
Group by the unit your evaluation claim requires to be unseen: incident thread, policy document, tenant, or time period. Low token overlap doesn't make two turns from T102 independent. If several group identifiers connect cases, keep their connected groups together. For a chronological holdout, train on earlier periods and check windows that straddle the cutoff. The check below demonstrates thread separation only.
Choose the grouping from the deployment claim. Thread grouping tests new threads, not unseen tenants or policy documents shared by those threads. Deduplicate across groups, and reserve separate validation and final-test groups. Use validation for sweeps and checkpoint selection; inspect the final test only after freezing that choice. Repeatedly tuning on the test turns it into another validation set.
Predict the three printed sets before reading the output. T102 should stay with P771 in training, R550 should be held out, and the intersection should be empty.
1rows = [
2 {"case_id": "T102", "turn": 1, "answer": "Open a rotation ticket; keys older than 14 days need review."},
3 {"case_id": "T102", "turn": 2, "answer": "Keep the review open until security approves the rotation."},
4 {"case_id": "R550", "turn": 1, "answer": "Escalate privileged-role changes."},
5 {"case_id": "P771", "turn": 1, "answer": "Cite the session-timeout policy."},
6]
7
8eval_cases = {"R550"}
9train = [row for row in rows if row["case_id"] not in eval_cases]
10evaluation = [row for row in rows if row["case_id"] in eval_cases]
11
12train_cases = {row["case_id"] for row in train}
13held_out_cases = {row["case_id"] for row in evaluation}
14assert train_cases.isdisjoint(held_out_cases)
15
16print("train_cases=", sorted(train_cases))
17print("eval_cases=", sorted(held_out_cases))
18print("case_overlap=", train_cases & held_out_cases)1train_cases= ['P771', 'T102']
2eval_cases= ['R550']
3case_overlap= set()Two turns from incident T102 have different wording. Can one train and the other evaluate if their token overlap is low?
Answer
No. They share the incident being held out, so both turns stay in the same split. Choose grouping from the intended deployment claim; thread separation doesn't establish unseen-document or unseen-tenant performance.
Data path: template, tokenize, label, pack
Follow one row all the way through the data path. Structured conversation dictionaries must pass through a Jinja2 chat template, undergo tokenization, receive target label masks, and either get packed into fixed-length blocks or padded into batches.
Use the checkpoint's actual formatting contract. ChatML-style markers such as <|im_start|> aren't universal; templates can mix special IDs with ordinary text, spaces, or newlines. Subword splitting doesn't by itself prove a broken role boundary. Inspect the rendered template, its token IDs, and assistant masks together rather than registering invented markers and assuming the pretrained model understands them.[8]
For complete training conversations, use add_generation_prompt=False. To start a new assistant reply, add_generation_prompt=True adds a header when the template supports it; continuing a partial assistant message is a different operation. Align the template's turn terminator with serving stop settings. Missing terminators can weaken stopping supervision, but don't prove inevitable endless generation. If you format text first and tokenize later, avoid adding duplicate special tokens.[8]
Version this entire data path together. Delimiter, whitespace, or template changes can alter token IDs or labels. Record the change and rerun batch checks; previous runs remain evidence about their recorded formatting.
Predict which positions in T102 should teach the model before inspecting the labels below. The system instruction, user request, and delimiter tokens must remain visible so the response can condition on them. However, only the target answer and its turn terminator should contribute to cross-entropy loss. We set ignored positions to -100, the standard ignore_index in PyTorch cross-entropy loss.
1IGNORE_INDEX = -100
2
3tokens = [
4 ("prefix", "<system>"),
5 ("prefix", "Follow access policy."),
6 ("prefix", "<user>"),
7 ("prefix", "Key T102 is 18 days old."),
8 ("prefix", "<assistant>"),
9 ("target", "Open a rotation ticket; keys older than 14 days need review."),
10 ("target", "<end_of_turn>"),
11]
12
13labels = [
14 token if span == "target" else IGNORE_INDEX
15 for span, token in tokens
16]
17
18supervised = [
19 token
20 for (_, token), label in zip(tokens, labels)
21 if label != IGNORE_INDEX
22]
23print("supervised_tokens=", supervised)
24print("masked_positions=", sum(label == IGNORE_INDEX for label in labels))1supervised_tokens= ['Open a rotation ticket; keys older than 14 days need review.', '<end_of_turn>']
2masked_positions= 5These seven entries are readable spans rather than real subword token IDs. A real sentence expands into multiple integer IDs. The five prefix spans provide context without generating loss; the answer and terminator spans identify what to label after tokenization.
Causal language models shift targets by one position: logits at position predict the label at position . The hidden state at the assistant marker <assistant> predicts the first answer token. Setting a label to -100 removes that token's direct cross-entropy loss, but it doesn't remove gradients flowing through attention layers into earlier prompt representations. Response loss still updates the weights that process the prompt. Mask padding tokens by position, not by blindly masking every EOS ID when EOS doubles as padding.[9]
Don't confuse the label mask with the attention mask. The label mask decides which positions contribute to the loss scalar. The attention mask decides which earlier positions a token can read during the forward pass. Prompt tokens must remain visible as context even when their labels are -100; packed sequences also need attention boundaries so adjacent examples can't attend across boundaries.
The CPU check below uses real integer labels and PyTorch gradients. With all logits set to zero over a four-token vocabulary, each supervised target has cross-entropy loss . Only logit rows 1 and 2 predict the labeled answer and stop positions. Most training frameworks handle this target shift internally; don't shift labels twice.
1import math
2import torch
3import torch.nn.functional as F
4
5# Positions: prompt, assistant marker, answer, stop.
6labels = torch.tensor([-100, -100, 2, 3])
7logits = torch.zeros(4, 4, dtype=torch.float64, requires_grad=True)
8targets = labels[1:]
9assert (targets != -100).any(), "Reject an all-ignored target batch"
10loss = F.cross_entropy(logits[:-1], targets, ignore_index=-100)
11loss.backward()
12active_rows = torch.nonzero(logits.grad.abs().sum(dim=1)).flatten().tolist()
13assert active_rows == [1, 2]
14assert math.isclose(loss.item(), math.log(4))
15assert (torch.full_like(targets, -100) != -100).sum().item() == 0
16print("logit_rows_with_direct_loss=", active_rows)
17print(f"mean_target_loss={loss.item():.6f}")1logit_rows_with_direct_loss= [1, 2]
2mean_target_loss=1.386294This verifies shifted loss and direct logit gradients. Test those mechanisms on collated tensors before launching long GPU jobs.
Why mask prompt targets here? We want a response-only objective and a metric that measures the response. With 800 scored prompt positions and 30 response positions, prompt targets occupy about 96% of a full-sequence mean's terms. That isn't 96% of the gradient: individual losses and their parameter derivatives can differ greatly. Nor does full-sequence training inevitably cause echoing or degraded style.
Huerta-Enochian and Ko studied prompt-loss weights using LLaMA 1/2 7B and Alpaca-derived datasets. Fractional weights improved some short-completion benchmarks in that setting; the study attributed the effect to regularization rather than increased training stability.[10] Response-only masking is this run's deliberate choice, not a universal definition of good SFT. Compare objectives on the same held-out tasks before changing it.
Predict which number exposes the metric difference below. The supplied prefix losses are small, so adding them lowers the full-sequence average without improving the supplied answer losses:
1prefix_losses = [0.03, 0.04, 0.02, 0.05, 0.03]
2answer_losses = [2.20, 1.80]
3
4full_sequence_loss = sum(prefix_losses + answer_losses) / len(prefix_losses + answer_losses)
5response_only_loss = sum(answer_losses) / len(answer_losses)
6
7print(f"full_sequence_loss={full_sequence_loss:.3f}")
8print(f"response_only_loss={response_only_loss:.3f}")
9assert response_only_loss > full_sequence_loss1full_sequence_loss=0.596
2response_only_loss=2.000Both numbers measure token-level cross-entropy over different target sets. In this authored fixture, five low-loss prefix positions lower the combined average to while answer loss remains . For our response-only objective, the loss is mean negative log-likelihood over the labeled target set :
represents the trainable parameters, the target token, and its available prefix. Here has two positions. Default unweighted, mean-reduced PyTorch cross-entropy with ignore_index=-100 averages only active targets. Changing the mask changes the objective; compare separately reported response loss rather than treating lower full-sequence loss as evidence of better answers.
You rarely write this mask by hand. In Hugging Face TRL's SFTTrainer, prompt-completion datasets use completion-only loss by default. Conversational datasets can set assistant_only_loss=True, provided their chat template exposes assistant spans through generation markers.[1] Always inspect a tokenized batch before launching training: confirm that only answer and stop tokens carry non--100 labels.
Inspect masks after truncation and causal shifting. Keeping only the start can remove the completion. A mean-reduced PyTorch cross-entropy batch with no active shifted targets returns NaN; an individual empty row in a mixed batch instead adds context and cost without response supervision. Count dropped or empty rows so dataset loss isn't hidden.
As checked September 22, 2026, TRL's pinned preparation path already filters fully masked rows after nonpacked truncation. Its packing path needs a separate audit.[11] Custom collators or skipped preparation can change this behavior. Check active labels in the actual collated tensors rather than assuming every empty row reaches the loss.
When sequence packing is active, this truncation step is skipped and overflow is handled by the packing strategy instead; bfd discards overflow tokens.[1] Predict what the guard below does when five prompt tokens consume the length budget:
1IGNORE_INDEX = -100
2
3prompt_tokens = ["<system>", "policy", "<user>", "stale_key", "<assistant>"]
4answer_tokens = ["open_ticket", "<end_of_turn>"]
5max_length = 5
6
7kept_tokens = (prompt_tokens + answer_tokens)[:max_length]
8labels = [
9 token if position >= len(prompt_tokens) else IGNORE_INDEX
10 for position, token in enumerate(kept_tokens)
11]
12has_answer_label = any(label != IGNORE_INDEX for label in labels)
13
14# A label at position zero has no preceding logit in this isolated sequence.
15first_only_labels = [7, IGNORE_INDEX]
16has_shifted_target = any(label != IGNORE_INDEX for label in first_only_labels[1:])
17assert any(label != IGNORE_INDEX for label in first_only_labels)
18assert not has_shifted_target
19
20print("kept_tokens=", kept_tokens)
21print("has_answer_label=", has_answer_label)
22print("decision=", "keep" if has_answer_label else "reject_or_increase_max_length")
23print("first_position_label_has_shifted_target=", has_shifted_target)1kept_tokens= ['<system>', 'policy', '<user>', 'stale_key', '<assistant>']
2has_answer_label= False
3decision= reject_or_increase_max_length
4first_position_label_has_shifted_target= FalsePack short examples while preserving their boundaries
Short demonstrations can leave many unused token slots when padded to a large fixed length. The fraction of padded slots isn't a measured compute-waste percentage: dynamic batch padding, length grouping, and attention kernels change the work performed. Packing is another efficiency option. For this run's independent-example objective, preserve each example's attention and loss boundaries.
For independent demonstrations, a reply for case R550 shouldn't condition on T102 merely because the loader put them in one block. Reset position IDs alone don't block attention. Enforce a causal block-diagonal pattern using an explicit mask or a kernel's segment metadata. The predicate below describes the intended edges; it doesn't execute an attention kernel:
1segments = [
2 ["T102 prompt", "T102 response", "<eos>"],
3 ["R550 prompt", "R550 response", "<eos>"],
4]
5
6flat = [(sequence_id, token) for sequence_id, row in enumerate(segments) for token in row]
7position_ids = [position for row in segments for position, _ in enumerate(row)]
8
9def may_attend(query_index, key_index):
10 query_sequence, _ = flat[query_index]
11 key_sequence, _ = flat[key_index]
12 return key_index <= query_index and query_sequence == key_sequence
13
14r550_response = 4
15assert position_ids == [0, 1, 2, 0, 1, 2]
16assert not may_attend(r550_response, 1)
17print("position_ids=", position_ids)
18print("R550_response_sees_T102_response=", may_attend(r550_response, 1))1position_ids= [0, 1, 2, 0, 1, 2]
2R550_response_sees_T102_response= FalseReset position IDs provide positional coordinates, not an attention barrier. The underlying attention kernel must enforce those boundaries during computation. Resetting positions without segment masking causes cross-example contamination:
| Index | Token | position_id | Intended attend-to | Broken path (no segment mask) |
|---|---|---|---|---|
| 0 | T102 prompt | 0 | T102 only | T102 only |
| 1 | T102 response | 1 | T102 only | T102 only |
| 2 | T102 eos | 2 | T102 only | T102 only |
| 3 | R550 prompt | 0 | R550 only | T102 + R550 (cross-example leak) |
| 4 | R550 response | 1 | R550 only | T102 + R550 (R550 answer conditions on T102) |
| 5 | R550 eos | 2 | R550 only | T102 + R550 |
Packing can also fail when an example gets truncated before being packed. If T102 survives truncation with every response token stripped away, packing still incorporates the prefix to fill the block, but that segment contributes zero learning signal. Inspect segment boundaries and verify non-zero answer labels before starting a training run.
TRL's bfd packs examples after truncating overlong ones; bfd_split splits overlong sequences into chunks; wrapped concatenates a stream and cuts token sequences.[1] Preserving stored tokens through splitting doesn't preserve every completion's original prompt context or every trainable target. Choose the strategy from the learning objective, then inspect the resulting segments.
The pinned TRL trainer enables padding-free mode for both bfd and bfd_split, and warns when its attention implementation isn't among the supported FlashAttention variants.[11] An ordinary global causal kernel can read across examples. Disable packing while debugging if necessary; verify boundary reachability before enabling it. A warning or successful forward pass isn't a boundary test.
Attention isolation alone also isn't enough. A globally shifted loss must not ask the last token of one independent segment to predict the first token of the next. The pinned padding-free collator masks labels wherever segment position IDs reset to zero.[11] Audit that loss boundary as well as the attention boundary.
Packing changes loader-item counts and can remove or split targets. A fixed max_steps can therefore imply a different number of data passes, without proving overfitting. Log processed input tokens, active shifted targets, and retained original examples. Throughput and learning exposure use different denominators.
A packed batch resets position IDs for R550, but R550's response can still attend to T102. Is the packing path correct?
Answer
No. Reset positions don't by themselves enforce attention isolation. The attention implementation must honor segment boundaries, and shifted labels must retain the intended answer targets without creating a target across an independent segment boundary.
Hyperparameters are experiments, not promises
A learning-rate sweep asks how much change this objective and parameterization can tolerate. Check task improvement and regressions on general capabilities together. Neither a frozen base nor a small learning rate guarantees preserved behavior.
The fragment below uses 1e-4 as an adapter candidate. That isn't a universal safe range or a fixed multiplier over full fine-tuning. Compare nearby rates under the same data exposure and evaluation rules. Dataset quality, optimizer, adapter scaling, model size, and loss normalization across microbatches all affect the result.
| Knob | Baseline to test | What decides whether it survives |
|---|---|---|
| Learning rate | a short sweep for the actual update surface | held-out task gains and regressions at matched exposure |
| Schedule | explicit warmup duration, decay shape, and final LR | observed early updates and behavior at the chosen checkpoint |
| Passes over data | a bounded pilot with measured retained targets | validation trends and task behavior, not a fixed epoch ceiling |
| Regularization & clipping | record decay groups and clipping threshold | compare changes in gradients, parameter updates, and held-out behavior |
| Precision | BF16 where supported; otherwise choose an explicit path | finite numerics and task parity under the actual hardware |
Warmup reduces early step sizes; it doesn't make AdamW's bias correction unnecessary or guarantee stability. A useful historical counterexample is InstructGPT's standalone SFT baseline: 16 epochs, no warmup, and cosine decay to 10% of its initial LR, with learning rates and epochs selected by searches.[3] That recipe isn't a recommendation to copy 16 epochs. It disproves a mandatory 3%-10% warmup or 1-3 epoch rule.
A cosine schedule's floor is configurable; standard Transformers cosine warmup decays to zero, while a minimum-LR variant can use a nonzero floor.[12] Resolve the actual scheduler instead of assuming 10%. AdamW applies decoupled decay separately from moment estimation. Decay doesn't bound distance from the pretrained weights, and clipping raw gradients doesn't directly bound AdamW's parameter update.[13] Inspect both gradients and update norms.
Effective batch size is also a token budget
The batch size configured on a single device isn't the batch size the optimizer sees during an update step. For synchronous training across complete accumulation windows, the number of items per update follows:
1loader_items_per_update = per_device_batch * gradient_accumulation_steps * world_sizeWithout packing, a loader item is one demonstration. With packing, each item is a composite block containing multiple short demonstrations, so the formula counts packed blocks rather than raw examples. Track both raw examples and actual supervised tokens per update. The CPU check below models an unpacked dataset across a four-GPU setup:
1per_device_batch = 4
2grad_accum = 8
3world_size = 4
4examples_per_update = per_device_batch * grad_accum * world_size
5
6supervised_tokens_per_microbatch_on_one_rank = [180, 212, 164, 220, 198, 204, 175, 207]
7supervised_tokens_per_update = sum(supervised_tokens_per_microbatch_on_one_rank) * world_size
8
9print(f"examples_per_update={examples_per_update}")
10print(f"supervised_tokens_per_update={supervised_tokens_per_update}")1examples_per_update=128
2supervised_tokens_per_update=6240Always calculate the combined impact of batch changes rather than reasoning about single knobs. Notice how doubling world size or accumulation doubles the update volume:
1def examples_per_update(per_device_batch, accumulation, world_size):
2 return per_device_batch * accumulation * world_size
3
4run_a = examples_per_update(4, 8, 4)
5run_b = examples_per_update(2, 16, 8)
6
7print("run_a_examples_per_update=", run_a)
8print("run_b_examples_per_update=", run_b)
9print("ratio_b_over_a=", run_b / run_a)
10assert run_b == 2 * run_a1run_a_examples_per_update= 128
2run_b_examples_per_update= 256
3ratio_b_over_a= 2.0When sharing batch configurations, specify accumulation steps, world size, sequence packing mode, and supervised token counts.
Loss weighting across accumulation steps matters as well. Suppose one microbatch has one target loss of 4.0 and another has nine target losses of 1.0 each. Averaging batch means gives 2.5; the token mean is . Equal example weights and equal token weights are different objectives. For the token mean, normalize loss sums by active shifted targets across the complete update window.
1microbatches = [[4.0], [1.0] * 9]
2mean_of_means = sum(sum(batch) / len(batch) for batch in microbatches) / len(microbatches)
3token_mean = sum(map(sum, microbatches)) / sum(map(len, microbatches))
4assert mean_of_means == 2.5 and token_mean == 1.3
5print(f"mean_of_means={mean_of_means:.1f}; token_mean={token_mean:.1f}")1mean_of_means=2.5; token_mean=1.3In ordinary distributed data parallel training, rank gradients are averaged. If each rank divides its local loss sum by a globally summed token count, that averaging introduces another factor of world_size. Compensate for it, or use a trainer that performs the normalization correctly. Don't add the factor twice when the trainer already handles it. Unequal target counts, partial accumulation windows, and actual loss gradients need checks; the scalar-loss fixture above doesn't execute distributed training.
What a resumable run must save
A model.safetensors file doesn't encode the optimizer history or next data batch. It may be part of an inference artifact, but serving also needs compatible architecture and tokenization. An adapter export additionally needs its base model and adapter configuration.
To reconstruct the continuation, retain the state your recipe actually uses:
- model weights or LoRA adapter weights
- optimizer state (AdamW first moment , second moment , and step counters)
- scheduler state (step counter, warmup phase, decay multiplier)
- global step and completed epoch indices
- sampler progress and random generator states, not only original seeds
- tokenizer revision and chat-template hash
- resolved training configuration and dataset manifests
- best validation metric achieved so far
- loss-scaler state if using FP16 dynamic scaling, plus any other stateful training machinery
Optimizer history determines direction and scaling. Scheduler progress determines the next LR, which needn't be the peak if reset incorrectly. Sampler and generator states determine upcoming data and randomness. Preserve the exact tokenizer/template artifacts and the adapter's base revision. Metadata labels or filenames alone don't guarantee identical contents or complete worker state.
For ordinary AdamW without AMSGrad, the update uses newly computed, bias-corrected moments:
Here is the current gradient and squaring is elementwise. A fresh optimizer with nonzero first computes and . Its update doesn't divide the current gradient by epsilon alone. Dropping state changes the trajectory, but doesn't imply a mandatory spike.[13]
Predict the direction after gradients 2.0, then -1.0. A fresh AdamW step follows the negative gradient upward, while retained positive momentum can still move the weight downward:
1import copy
2import math
3import torch
4
5weight = torch.nn.Parameter(torch.tensor(1.0, dtype=torch.float64))
6optimizer = torch.optim.AdamW([weight], lr=0.1, weight_decay=0.0)
7
8def update(parameter, opt, gradient):
9 opt.zero_grad(set_to_none=True)
10 parameter.grad = torch.tensor(gradient, dtype=torch.float64)
11 opt.step()
12
13update(weight, optimizer, 2.0)
14saved_weight = weight.detach().clone()
15saved_optimizer = copy.deepcopy(optimizer.state_dict())
16update(weight, optimizer, -1.0)
17
18resumed = torch.nn.Parameter(saved_weight.clone())
19resumed_optimizer = torch.optim.AdamW([resumed], lr=0.1, weight_decay=0.0)
20resumed_optimizer.load_state_dict(saved_optimizer)
21update(resumed, resumed_optimizer, -1.0)
22
23fresh = torch.nn.Parameter(saved_weight.clone())
24fresh_optimizer = torch.optim.AdamW([fresh], lr=0.1, weight_decay=0.0)
25update(fresh, fresh_optimizer, -1.0)
26assert torch.equal(resumed, weight)
27assert resumed.item() < saved_weight.item() < fresh.item()
28assert math.isclose(fresh.item() - saved_weight.item(), 0.1, abs_tol=1e-8)
29print(f"uninterrupted={weight.item():.6f}; resumed={resumed.item():.6f}")
30print(f"fresh_optimizer={fresh.item():.6f}; all_finite={bool(torch.isfinite(fresh))}")1uninterrupted=0.873366; resumed=0.873366
2fresh_optimizer=1.000000; all_finite=TrueThis controlled CPU experiment restores optimizer history while supplying the same next gradient; it doesn't replay a stochastic LLM job. Prefer saving at completed optimizer-step boundaries. Mid-accumulation continuation also needs partial gradients and microbatch progress. Compare the next batch, LR, loss, and update against an uninterrupted reference. Bitwise parity depends on deterministic operations and a compatible environment.[14][15]

A manifest can list required continuation artifacts. This schema fixture checks key presence only; it doesn't open files, validate hashes, or demonstrate a successful resume:
1required = {
2 "model_state",
3 "optimizer_state",
4 "scheduler_state",
5 "global_step",
6 "rng_state",
7 "sampler_state",
8 "tokenizer_version",
9 "chat_template_version",
10 "data_manifest",
11 "training_config",
12 "eval_manifest",
13 "best_metric",
14}
15
16resume_bundle = {
17 "model_state": "weights/step_600.safetensors",
18 "optimizer_state": "optimizer/step_600.pt",
19 "scheduler_state": "scheduler/step_600.pt",
20 "global_step": 600,
21 "rng_state": "rng/step_600.pt",
22 "sampler_state": {"epoch": 1, "batches_consumed": 120},
23 "tokenizer_version": "policy-sft-tokenizer-v3",
24 "chat_template_version": "llama3-access-policy-v2",
25 "data_manifest": "data/sft_manifest_2026-05-20.json",
26 "training_config": "config/resolved_training.json",
27 "eval_manifest": "eval/access_policy_behavior_v4.json",
28 "best_metric": {"name": "policy_pass_rate", "value": 0.97},
29}
30
31missing = sorted(required - resume_bundle.keys())
32print("schema_complete=", not missing)
33print("best_metric=", resume_bundle["best_metric"])1schema_complete= True
2best_metric= {'name': 'policy_pass_rate', 'value': 0.97}A resumed run reloads model weights and global step but not optimizer, scheduler, sampler, or tokenizer-template versions. Can its next metrics be compared as uninterrupted continuation?
Answer
No. The next update may use different momentum, learning rate, data position, randomness, or token formatting. Treat it as a new run unless the complete resume manifest is restored and validated.
The scalar SGD-with-momentum fixture below makes the dependency on velocity explicit. Its fixed next gradient isolates optimizer state; no training data or model loss is evaluated:
1import json
2
3def step(state, gradient):
4 velocity = 0.9 * state["velocity"] + gradient
5 return {"weight": state["weight"] - 0.1 * velocity,
6 "velocity": velocity, "step": state["step"] + 1}
7
8saved = step({"weight": 1.0, "velocity": 0.0, "step": 0}, 2.0)
9reference = step(saved, 1.0)
10restored = json.loads(json.dumps(saved))
11resumed = step(restored, 1.0)
12weights_only = step({**restored, "velocity": 0.0}, 1.0)
13assert resumed == reference and weights_only != reference
14print(f"uninterrupted={reference['weight']:.2f}; resumed={resumed['weight']:.2f}")
15print(f"reset_momentum={weights_only['weight']:.2f}")1uninterrupted=0.52; resumed=0.52
2reset_momentum=0.70Choosing the best checkpoint
Choose checkpoints against the deployment objective on validation data. Teacher-forced loss measures likelihood of reference tokens given reference prefixes; it doesn't directly measure whether a freely generated reply opens the required ticket. Generate outputs for task evaluation too, and keep a final test separate from repeated selection.
Unreliable selection criteria:
- lowest training loss alone (doesn't measure unseen-case performance)
- latest training step (risks overfitting or late degradation)
- file size or raw token count
Rigorous selection criteria:
- strict syntax and format pass rate (JSON schema compliance, valid function call formatting)
- domain policy pass rate on held-out cases via deterministic rule verifiers or calibrated judges
- human evaluation on blinded, randomized completion pairs
- validation loss with a declared role in selection, rather than an assumed universal priority
For the authored selection fixture below, require format validity of at least 98%, then maximize policy compliance and break ties with validation loss, followed by the earlier step. These rates and this threshold are supplied data and policy choices, not trained-model results or a universal release standard. In a real evaluation, report denominators, uncertainty, important failure slices, and regressions on general behavior.

1checkpoints = [
2 {"step": 200, "val_loss": 1.91, "policy_pass_rate": 0.93, "format_pass_rate": 0.99},
3 {"step": 400, "val_loss": 1.77, "policy_pass_rate": 0.91, "format_pass_rate": 1.00},
4 {"step": 600, "val_loss": 1.79, "policy_pass_rate": 0.97, "format_pass_rate": 0.98},
5]
6
7def select_checkpoint(rows):
8 eligible = [row for row in rows if row["format_pass_rate"] >= 0.98]
9 return max(eligible, key=lambda row: (row["policy_pass_rate"], -row["val_loss"], -row["step"]), default=None)
10
11best = select_checkpoint(checkpoints)
12assert best is not None
13assert select_checkpoint([]) is None
14assert select_checkpoint([{**checkpoints[0], "format_pass_rate": 0.97}]) is None
15assert select_checkpoint([{**checkpoints[0], "step": 600}, checkpoints[0]])["step"] == 200
16assert select_checkpoint([{**checkpoints[0], "val_loss": 2.0}, {**checkpoints[0], "step": 400}])["step"] == 400
17
18print("best_step=", best["step"])
19print("best_policy_pass_rate=", best["policy_pass_rate"])
20print("lowest_loss_step=", min(checkpoints, key=lambda row: row["val_loss"])["step"])1best_step= 600
2best_policy_pass_rate= 0.97
3lowest_loss_step= 400Step 400 shows the lowest validation loss (), but step 600 achieves the highest policy compliance rate () while clearing the format threshold (). The selection gate picks step 600. If no checkpoint meets the format threshold, the gate yields None, signaling that the recipe needs fixing before anything ships.
Single GPU first, then scale out
Don't start debugging thread T102 across an eight-GPU cluster. If prompt templating or label masking contains subtle bugs, scaling out distributes those defects across collective communication calls rather than resolving them. Prove your pipeline on a single device first:
- single GPU execution
- small, representative evaluation split
- frequent checkpointing
- verified chat templating and response-only label masking
- clear task compliance metric
Validate the smallest feasible run before adding scale. Some required models can't fit one device at all; debug data contracts with small fixtures, then use the minimum necessary sharding. Single-device success also doesn't prove the distributed loss normalization or checkpoint path is correct.
Signs that one GPU is still fine
A single device is feasible when weights, gradients, optimizer state, activations, and temporary buffers fit at the chosen precision and the run meets its time budget. A model that fits for inference may still fail this training-memory test.
Signs that you need FSDP or ZeRO
Consider FSDP or ZeRO when persistent model weights, gradients, and AdamW optimizer states exceed available GPU memory. Measure activation memory independently: sharding model parameters doesn't automatically eliminate activation overhead. Longer contexts often require activation checkpointing, reduced micro-batch sizes, or context parallelism. Adding more GPUs introduces network communication latency and won't speed up an improperly configured small job.
FSDP and ZeRO offer different state-sharding strategies; the chosen stage determines what is sharded.[16][17] Specify precision explicitly. Current TRL defaults to BF16 when FP16 isn't set, but a default doesn't establish hardware support or numerical parity.[1]
Minimal operational skeleton
The overall architectural pipeline remains consistent across modern training libraries:
1dataset -> template -> tokenizer -> collator/mask -> trainer
2 -> periodic eval -> checkpoint save -> best-checkpoint exportTRL, torchtune, and the Alignment Handbook implement variations of this workflow.[1][18][19]
The configuration fragment below is an unexecuted integration blueprint for conversational prompt-completion rows. Its pilot disables packing until boundaries are verified. API fields were checked against TRL on September 22, 2026; pin a tested package version and model revision for a real run.
1from trl import SFTConfig, SFTTrainer
2
3args = SFTConfig(
4 output_dir="runs/access-policy-sft-lora",
5 learning_rate=1e-4, # adapter baseline to evaluate, not a universal rule
6 max_length=1024,
7 bf16=True, # only on supported hardware; verify numerics
8 per_device_train_batch_size=1,
9 gradient_accumulation_steps=8,
10 max_steps=200, # authored pilot budget, not a recommended duration
11 lr_scheduler_type="linear", # explicit schedule; reaches zero at the budget endpoint
12 warmup_steps=0, # authored candidate to compare, not a stability promise
13 seed=42,
14 data_seed=42,
15 packing=False, # enable only after attention and loss-boundary checks
16 eval_packing=False, # keep validation examples separate
17 completion_only_loss=True, # train on the completion field, not the prompt
18 loss_type="nll", # explicit standard loss for this adapter fragment
19 eval_strategy="steps",
20 eval_steps=100,
21 save_steps=100,
22)
23
24trainer = SFTTrainer(
25 model=model_with_lora_adapters,
26 args=args,
27 train_dataset=train_rows, # prompt/completion are role/content message lists
28 eval_dataset=held_out_rows,
29 processing_class=tokenizer,
30)This fragment omits model loading, adapter initialization, and dataset construction. Convert the earlier authoring rows to message lists so the trainer can apply the checkpoint's template; raw string fields don't automatically become a chat conversation. We explicitly choose loss_type="nll"; current TRL's default chunked NLL computes the same target objective with a different projection/memory path.[1] The evaluation loop computes teacher-forced loss, not the supplied policy/format rates. Generate and grade validation replies separately before applying the selection rule.
Common pitfalls
When a training run regresses, analyze the symptoms and identify which system boundary failed. This prevents data formatting or masking bugs from turning into futile hyperparameter sweeps.
Confusing the objective with the update surface
Saying "we tried LoRA instead of SFT" mixes two distinct decisions. LoRA and full fine-tuning describe which weights are updated; SFT defines the supervised next-token cross-entropy objective. Define the behavioral target and objective first. Compare adapter tuning against full-weight tuning only when parameter capacity or memory budget is the variable under test.
Evaluating examples that share a case with training
High validation scores followed by failures on new documents can reflect leakage, distribution shift, or a mismatched metric. Check whether related cases crossed splits and whether the grouping actually tests the deployment claim. A thread-disjoint split doesn't test unseen documents shared across those threads.
Reporting micro-batch as if it were the real batch
When two experiments claim "batch size 4" but demonstrate drastically different stability, check accumulation steps and GPU count before modifying the optimizer. Always publish per-device batch size, gradient accumulation steps, world size, packing mode, and total supervised tokens per update step.
Saving weights without resume state
Changed loss or batches after resume prompts a state comparison; it doesn't uniquely diagnose an omitted checkpoint field. Check optimizer counters, LR, upcoming data, RNG states, precision/scaler state, and artifact identities against the uninterrupted reference.
Using a packing path that crosses example boundaries
Incoherent answers don't uniquely diagnose packing leakage. Test attention reachability and shifted-label boundaries directly. For independent demonstrations, use a verified segment-aware path; for intentional continuous-text training, cross-document context may be part of the chosen objective. A reset position ID or EOS token isn't itself an enforced barrier.
Selecting the best checkpoint by train loss
Training loss measures reference-token likelihood on seen examples, not memorization by itself. Use validation behavior and declared constraints for selection, then assess the frozen choice on the final test. Checks need a task-appropriate grader; fluent output and low loss don't certify policy compliance.