Personalize this lesson
Adapt explanations and teaching visuals to your background and preferred voice.
It's 3:15 AM. Your training cluster dies. You check the storage bucket: yesterday's weights are intact. You load them into your script, hit resume, and watch yesterday's training samples replay from the start while the learning rate restarts its warmup from zero. The model weights survived, but the training run died.
We'll use a concrete run named policy-sft-v4 to walk through every recovery boundary: four data-parallel replicas, two sequences per microbatch, eight microbatches per accumulation window, and a committed checkpoint at update 1840. The same operational rules govern full-weight pretraining and parameter-efficient fine-tuning (LoRA): which state survived, which data sample comes next, and did your replacement hardware alter the optimization trajectory?

Weights are only one part of a continuation
Training loops introduced the optimizer and scheduler. A restart must restore their internal state, not merely reconstruct the model architecture. You need to distinguish three operations before calling a checkpoint loader:
| Operation | State you need | Meaning of the next update |
|---|---|---|
| Continue training | Model weights, optimizer moments, scheduler counters, progress counters, data cursor, RNG states, scaler state | Continue from a recorded training boundary |
| Initialize a new run | Compatible weights or adapters and their base model | Start fresh optimization from those parameters |
| Export for inference | Weights or adapters, tokenizer, configuration, and serving assets | No optimizer update is implied |
For AdamW, optimizer state includes the first and second gradient moments ( and ) alongside parameter step counters [1]. For mixed precision, you must include the gradient scaler when the recipe uses one. Save per-rank random number generator (RNG) states and the data loader resumable cursor. Preserve the model, tokenizer, chat template, dataset slice, and training arguments needed to interpret them. Evaluation history helps select a checkpoint, but it isn't a substitute for optimization state.
Adapter checkpoints also require the correct frozen base weights and adapter configuration. A merged inference export isn't a resumable adapter checkpoint. Tooling frameworks distinguish intermediate training checkpoints from final exports; torchtune documents separate resume requirements for full-weight and LoRA recipes [2].
Continuation isn't a guarantee of bitwise replay. A different cluster topology can change all-reduce reduction trees, random streams, and data partition assignments. Restore the intended state, document any topology changes, and compare loss trajectories against an acceptable tolerance. Exact replay demands strict control over CUDA kernels, non-deterministic operations, and sample batching.
What lost optimizer moments break
What happens when you restore weights but let AdamW start with zeroed moments? AdamW normalizes gradient steps by running second moments: Here estimates a second moment, not the variance alone. If you reset both moments and the optimizer's step counter, the first bias-corrected moments are and . The adaptive gradient term moves a scalar coordinate by , whose magnitude is at most . A tiny gradient doesn't make that first update arbitrarily large. This bound excludes the separate weight-decay term.
Try three gradient magnitudes with the same learning rate:
1lr, epsilon = 1e-3, 1e-8
2for gradient in (1e-12, 1e-6, 1.0):
3 update = lr * gradient / (abs(gradient) + epsilon)
4 assert abs(update) <= lr
5 print(f"g={gradient:.0e}; |first update|={abs(update):.3g}")1g=1e-12; |first update|=1e-07
2g=1e-06; |first update|=0.00099
3g=1e+00; |first update|=0.001Resetting moments still discards the history that determines future direction and magnitude. Keeping an old step counter with fresh moments also changes their bias correction. Either restart can change the trajectory; neither necessarily causes a loss jump or divergence.
This scalar Adam example minimizes over a repeating target sequence. It tracks first and second moments, bias-correction steps, a step-based learning rate, and a consumed data cursor. JSON round-tripping models a serialized restart:
1import json
2import math
3
4TARGETS = [1.0, -2.0, 3.0, 0.5]
5
6def fresh(w=0.0):
7 return dict(w=w, m=0.0, v=0.0, step=0, cursor=0)
8
9def advance(state, updates):
10 if type(updates) is not int or updates < 0:
11 raise ValueError("updates must be a nonnegative integer")
12 state = dict(state)
13 for _ in range(updates):
14 g = state["w"] - TARGETS[state["cursor"] % len(TARGETS)]
15 state["step"] += 1
16 t = state["step"]
17 state["m"] = 0.9 * state["m"] + 0.1 * g
18 state["v"] = 0.999 * state["v"] + 0.001 * g * g
19 m_hat = state["m"] / (1 - 0.9 ** t)
20 v_hat = state["v"] / (1 - 0.999 ** t)
21 lr = 0.05 / (1 + 0.1 * (t - 1))
22 state["w"] -= lr * m_hat / (math.sqrt(v_hat) + 1e-8)
23 state["cursor"] += 1
24 return state
25
26saved = advance(fresh(), 7)
27restored = json.loads(json.dumps(saved))
28resumed = advance(restored, 5)
29baseline = advance(fresh(), 12)
30assert resumed == baseline
31
32# Isolate missing moments: keep the weights, step and data cursor.
33lost_moments = dict(restored, m=0.0, v=0.0)
34wrong = advance(lost_moments, 5)
35assert abs(wrong["w"] - baseline["w"]) > 1e-4
36assert advance(restored, 0) == restored
37print(f"continuous={baseline['w']:.6f}; resumed={resumed['w']:.6f}")
38print(f"reset moments={wrong['w']:.6f}; next cursor={resumed['cursor']}")1continuous=0.125436; resumed=0.125436
2reset moments=0.126672; next cursor=12Even keeping the step counter doesn't reconstruct the lost moments. In a real cluster restore drill, compare the next sample IDs, current learning rate, optimizer moments, and next gradient update against an uninterrupted reference run. A smooth-looking loss curve doesn't prove mathematical continuity.
Reliability and the two checkpoint completion points
Large clusters have more opportunities for interruption. Llama 3 405B used configurations with up to 16,384 H100 GPUs. Over one reported 54-day period, the team recorded 419 unexpected interruptions and 47 planned ones; these included several kinds of hardware and software faults [3]. Those observations don't establish a universal GPU failure rate.
For a simplified model of independent devices with identical, constant failure rates, where any device failure stops the job, rates add: Shared power, network faults, repairs, and software failures violate that model's assumptions. Measure your own run's interruption frequency and lost work rather than substituting a supposed server MTBF into a GPU-count formula.
An asynchronous checkpoint separates three places where state can live:
- Live training memory: Parameters, optimizer state, and other saved state can change as training proceeds.
- Stable host snapshot: Staging makes a copy independent of subsequent live mutations. Pinned host buffers can improve transfer performance, but consume reserved RAM.
- Storage: Background I/O persists the snapshot. Only successful storage completion and the required publication protocol establish a new recovery point.
PyTorch's standard async_save stages state before returning a background storage future. Its default staging can still block training. The optional DefaultStager, introduced in PyTorch 2.9, moves staging to a background CPU thread; the documented response separates staging_completion from upload_completion [4].
Wait for staging before modifying any live state whose copy is still pending. Parameters aren't the only concern: forward passes can change buffers, and the next batch can advance RNG or loader state. Capture progress metadata at the same boundary as the tensors. Wait for storage completion before publishing or deleting the previous recovery point. Bound concurrent saves, propagate future exceptions, and close the stager when finished [5].
Asynchronous I/O can overlap computation, but still competes for host memory, CPU time, and storage bandwidth. Measure staging latency, storage latency, and training throughput together; there is no fixed “100 GB in three seconds” guarantee.
Atomic manifest publishing and data cursor alignment
Use a new checkpoint identifier for each save. Keep the previous committed snapshot until the new shards, framework metadata, and application state have completed successfully. A manifest should identify one consistent boundary and the expected artifacts, with integrity information such as sizes and hashes. Directory existence or a .tmp suffix alone doesn't prove completeness.
Publication depends on the storage backend. On a suitable filesystem, an atomic rename can replace a pointer file after the checkpoint completes; crash durability also requires the relevant file and directory synchronization Linux fsync documentation. An object store doesn't provide a transaction over an arbitrary collection of shard uploads. For Google Cloud Storage, individual object writes are atomic; a generation-match precondition can protect a latest pointer update from competing writers [6] [7]. Publish that pointer only after its referenced checkpoint is complete, and validate the checkpoint when loading.
A clean weight restore still fails if the data position describes a different boundary. A recorded consumed batch count plus a deterministic sampler can suffice for a simple replayable dataset. A worker's prefetched producer position can be ahead of what training consumed. Stateful streaming, random transforms, and packing may need worker RNG, sampler, dataset, and buffered-example state as well.
TorchData's StatefulDataLoader supports state_dict and load_state_dict; its default strategy tracks yielded batches and fast-forwards, while stateful datasets and samplers can provide their own restore state. It aggregates state across workers, not across distributed ranks [8]. A topology change therefore needs an explicit data-partition policy, alongside tensor resharding.
Preemption is a strict deadline, not a save command
Suppose update 1840 is durably committed to storage, and a preemption warning signal (such as a cloud spot termination notice or a cluster scheduler SIGTERM) arrives after microbatch 5 of the next 8-microbatch accumulation window.
Your training policy saves only at completed optimizer boundaries. There are two valid operational outcomes:
- All nodes remain healthy and there's enough time to complete the remaining 3 microbatches, perform the optimizer step, stage the checkpoint to host RAM, and update the manifest for step 1841.
- Time runs out or a rank dies abruptly. The orchestrator terminates, discards partial in-flight state, and reboots from durable step 1840. In this illustrative run, every update consumed 64 examples and none was skipped, so the saved global consumed-example count is . That count isn't itself a dataset row ID or byte offset.
Never pair weights from step 1840 with a data cursor that advanced past microbatch 5. That mistake silently drops those five microbatches without applying their gradients. Mid-window recovery needs the in-flight accumulated gradients, accumulation counter, matching scaler, model buffers, RNG, and data state. With AMP, those gradients may still be scaled; preserve their representation and keep the scale constant throughout accumulation [9]. Unless your framework explicitly supports and tests that recovery, use completed update boundaries.
Signal handlers should only set an internal shutdown flag, not initiate distributed network I/O. The main loop coordinates the save at a clean boundary. If a member of an ordinary fixed process group dies, collectives can't complete normally; a fault-tolerant launcher must recover or reconfigure the job. Periodic durable checkpoints remain necessary because not every failure gives a warning.
Budget the entire remaining path
A 90-second checkpoint doesn't fit into a 90-second warning if you also need to finish active microbatches and publish manifests. The preemption budget equation is: Here includes time through durable publication, and absorbs network and disk jitter:
1import math
2
3def can_finish(grace, window, save, margin):
4 times = (grace, window, save, margin)
5 if any(isinstance(t, bool) or not isinstance(t, (int, float))
6 or not math.isfinite(t) or t < 0 for t in times):
7 raise ValueError("times must be finite nonnegative seconds")
8 return window + save + margin <= grace
9
10assert can_finish(120, 20, 90, 10)
11assert not can_finish(119, 20, 90, 10)
12assert not can_finish(90, 20, 90, 10)
13print("120s warning: feasible at the estimate; 119s: insufficient")
14for productive_seconds in (300, 1200):
15 fraction = 90 / (productive_seconds + 90)
16 print(f"{productive_seconds}s compute + 90s save: {fraction:.1%} save overhead")1120s warning: feasible at the estimate; 119s: insufficient
2300s compute + 90s save: 23.1% save overhead
31200s compute + 90s save: 7.0% save overheadA blocking save every five minutes of compute spends 23.1% of cluster wall-clock time writing files under these authored timings. Increasing the compute interval to twenty minutes drops that fraction to 7.0%, but increases work exposed to failure. With asynchronous saving, work performed while a snapshot uploads is also uncommitted until publication. Choose the interval using measured latency, interruption frequency, and acceptable replay cost. The warning calculation establishes feasibility at the estimates, not a deadline guarantee.
Sharded saves need a compatible restore path
FSDP and ZeRO-3 distribute persistent parameters and optimizer state rather than keeping a complete copy on every GPU [10]. FSDP can temporarily all-gather a module's weights for computation. For checkpointing, ask two questions: did the save complete reliably, and can the restore loader interpret its layout?
A full state dictionary can gather tensors onto one rank, possibly into host memory. For 70 billion parameters, BF16 weights occupy 140 decimal GB; each FP32 Adam moment occupies 280 GB. Two moments plus those weights total 700 GB, before any separately saved master weights or other state. Consolidation needs enough capacity at its destination and substantial transfer time; it isn't automatically an instant GPU OOM if the framework offloads to a sufficiently large host.
PyTorch Distributed Checkpoint (DCP) can save local shards without first constructing the entire tensor on one rank [11]. Its metadata describes parameter identities, global tensor properties, and shard coordinates. The storage path can stage tensors through CPU memory; “distributed save” doesn't imply a direct GPU-to-object-store transfer.
![Logical checkpoint resharding. Four source ranks save a 16-element tensor in slices [0:4] through [12:16]. Eight destination ranks use metadata to identify the source slices covering [0:2] through [14:16]. Distributed readers fill the new layout without requiring a full-tensor gather; their physical I/O and staging memory depend on the storage implementation.](/cdn/content-image/training/training-run-operations/illustrations/_generated/sharded_save_reshard_dark.png?v=29b1a2d1477c)
Resharding across topologies without single-node bottlenecks
When you intentionally change the partitioned layout, the loader must map saved coordinates into the new layout. Not every model-parallel mesh permits every world size: check divisibility, parameter identities, and the framework's supported model and optimizer mappings.
Under DCP, resharding happens without central gathering:
- Initialize compatible destination model and optimizer state dictionaries, with tensors allocated in the target distributed layout.
- Each destination rank consults the shared
.metadataindex file. - Each rank calculates which source shard files intersect its assigned tensor slice.
- Distributed readers load the relevant serialized data and copy the required logical slices into destination tensors.
DCP's PyTorch 2.14 filesystem reader loads a serialized tensor into CPU memory and then narrows it to the requested slice. Logical slicing therefore doesn't guarantee reading only the corresponding value bytes or allocating only the destination slice [12].
DCP explicitly doesn't guarantee checkpoint compatibility across PyTorch versions, including changes within a major release [5]. Test the actual target stack. DeepSpeed Universal Checkpointing uses a ZeRO save, explicit conversion, and Universal load; supported topology changes still require matching model parameter identities and compatible shapes [13].
This small coordinate exercise reconstructs eight destination slices from four source slices and rejects gaps and overlaps. It deliberately builds a complete Python list in one process for inspection. It is not a DCP implementation or a test of distributed I/O and memory:
1def reshard(shards, size, destinations):
2 if type(size) is not int or size <= 0:
3 raise ValueError("positive tensor size required")
4 if type(destinations) is not int or destinations <= 0 or size % destinations:
5 raise ValueError("this toy requires equal nonempty destination shards")
6 tensor = [None] * size
7 for start, values in shards:
8 if type(start) is not int or start < 0 or start + len(values) > size:
9 raise ValueError("invalid range")
10 for offset, value in enumerate(values, start):
11 if tensor[offset] is not None:
12 raise ValueError("overlapping ranges")
13 if value is None:
14 raise ValueError("None is reserved for a missing element")
15 tensor[offset] = value
16 if any(value is None for value in tensor):
17 raise ValueError("missing range")
18 width = size // destinations
19 return [(i, tensor[i:i + width]) for i in range(0, size, width)]
20
21saved = [(i, list(range(i, i + 4))) for i in range(0, 16, 4)]
22loaded = reshard(saved, 16, 8)
23assert [v for _, values in loaded for v in values] == list(range(16))
24assert reshard(loaded, 16, 4) == saved
25for broken in (saved[:-1], saved + [saved[0]]):
26 try:
27 reshard(broken, 16, 8)
28 except ValueError as error:
29 print(error)
30 else:
31 raise AssertionError("incomplete or overlapping state was accepted")
32print(loaded)1missing range
2overlapping ranges
3[(0, [0, 1]), (2, [2, 3]), (4, [4, 5]), (6, [6, 7]), (8, [8, 9]), (10, [10, 11]), (12, [12, 13]), (14, [14, 15])]Keep the update size separate from the GPU count
For complete accumulation windows with equal microbatch sizes: Here is the sequence count per replica's microbatch, is the number of accumulated microbatches per update, and counts independent data-parallel replicas. Workers collaborating on one replica don't contribute independent examples.
For a regular mesh whose distinct axes are tensor parallelism (), pipeline parallelism (), context parallelism (), and data parallelism: Context-parallel workers split one example's sequence rather than adding independent examples. This product formula isn't a universal description of overlapping expert-parallel or hybrid process groups. Eight GPUs with , , and form replicas. Halving accumulation while keeping cuts the nominal global batch in half.

| Microbatch | Accumulation | Data-Parallel | Global Sequences per Update |
|---|---|---|---|
| 2 | 8 | 4 | 64 |
| 2 | 8 | 8 | 128 |
| 2 | 4 | 8 | 64 |
Sequences versus supervised tokens and loss reduction
Sequence counts don't equal token throughput. In instruction tuning and conversational SFT, sequence packing, system prompt masking, and variable response lengths mean that two batches with identical sequence counts carry different numbers of loss-bearing tokens: If replica 0 processes 4,800 supervised tokens while replica 1 processes 9,600, averaging their per-device mean loss values equally computes . That naive average assigns twice as much weight to tokens on replica 0 as tokens on replica 1, distorting the true gradient.
If the intended objective is a mean over supervised tokens, use: This is a choice of objective: a deliberately equal-weighted per-example objective is different and can be valid. For the token mean, count supervised targets after shifting and masking across the whole accumulation window and all replicas. Ordinary DDP averages gradients across replicas [14]. With replicas and a global token count , backpropagating each replica's local loss sum multiplied by gives the desired gradient after that averaging. Check trainer normalization before adding any scale factor; don't also divide by when the global denominator already includes all microbatches.
Here the numbers are hypothetical scalar gradients of individual token losses. One replica has one token; another has three:
1tokens = [[10.0], [-2.0, -2.0, -2.0]]
2replicas = len(tokens)
3count = sum(len(rank) for rank in tokens)
4token_mean = sum(sum(rank) for rank in tokens) / count
5replica_mean = sum(sum(rank) / len(rank) for rank in tokens) / replicas
6local_scaled = [replicas * sum(rank) / count for rank in tokens]
7ddp_average = sum(local_scaled) / replicas
8assert token_mean == ddp_average == 1.0
9assert replica_mean == 4.0
10print(f"token mean gradient={token_mean:.1f}; replica mean={replica_mean:.1f}")
11print(f"DDP averaged corrected gradients={ddp_average:.1f}")1token mean gradient=1.0; replica mean=4.0
2DDP averaged corrected gradients=1.0This arithmetic checks the reduction, not a running process group. Inspect your distributed trainer's real gradient normalization with a small reference batch.
A 4-replica job with microbatch 2 and accumulation 8 restarts on eight GPUs configured as tensor parallel 2 and data parallel 4. Should you halve accumulation to preserve the batch?
Answer
No. Data parallelism is still four, so the batch remains 2 × 8 × 4 = 64. Halving accumulation would reduce it to 32. Count independent replicas, not devices.
Preserve the scheduler before tuning a new batch
When continuing a run at the same effective batch size, restore the scheduler phase directly. Don't restart warmup just because the process restarted.
Changing the global batch changes the optimization trajectory. Doubling tokens per update approximately halves the number of updates for a fixed token budget. Doubling sequences has that effect only if their average token contribution stays comparable.
A schedule with 100 warmup updates needs 6,400 examples at 64 examples/update and 12,800 at 128. Its milestones move to later data volumes, and a fixed data budget may end before the original step-based decay finishes. Decide whether to preserve update milestones or remap them to a data budget. Log whether the schedule uses updates, examples, or tokens.
A skipped AMP step isn't a successful update
In a standard FP16 GradScaler loop, nonfinite gradients can cause the optimizer step to be skipped. Keep separate counters for attempted accumulation windows, consumed data, and successful optimizer updates. A scheduler defined per successful update should stay fixed when no update occurred; a deliberately data-indexed scheduler follows its own policy.
Perform overflow checks and scale updates at effective-batch boundaries, keeping the scale constant during accumulation [9]. A skipped step can still consume data and change scaler state. Preserve both in the checkpoint. Don't infer success from scaler.step(...) returning None: many optimizers return None even after a real update. Use the framework's supported overflow handling.
Scaling rules are hypotheses to test
Goyal et al. popularized linear learning-rate scaling with gradual warmup for large-batch SGD, training ResNet-50 with batch 8,192 on 256 GPUs [15]. The approximation compares one large step with several smaller steps while assuming their gradients change little. It can break early in training or at very large batches; it isn't an AdamW SFT guarantee.
Malladi et al. derive an Adam scaling rule from stochastic differential equation approximations under specified noise and moment assumptions [16]. For a batch multiplier , their rule changes more than the learning rate:
The new must remain valid optimizer coefficients. This conditional result isn't a theorem that changing only AdamW's LR by preserves a fine-tuning run, nor that linear scaling always diverges. Compare controlled runs, including a fixed-LR baseline, and evaluate quality at a comparable data budget.
McCandlish et al. study a gradient-noise scale: where is the mean gradient and is the covariance of individual-example gradients [17]. It predicts the scale of the largest useful batch in their studied settings, rather than an exact universal threshold. Larger batches can reduce update count with diminishing returns, while increasing examples needed to reach a target. That optimization tradeoff doesn't imply linear wall-clock speedup: kernel efficiency and communication matter too.
Diagnose a symptom before choosing a recovery
A loss spike doesn't reveal its cause. Record the last known-good checkpoint, exact batch identity, token counts, learning rate, scaler state, and per-rank timings before changing the run.
| Symptom | Candidates to investigate | First check |
|---|---|---|
| Same weights, different first LR | Scheduler and optimizer parameter-group restore; warmup reset | Verify step counter and scheduler state dictionary |
| Old sample IDs reappear | Consumed cursor, sampler epoch seed, prefetch buffer state | Rewind cursor to exact completed update boundary |
| Straggler on one rank | Unequal input work, thermal or power limits, host stalls, PCIe issues | Compare phase timings and nvidia-smi -q -d PERFORMANCE clock event reasons |
| Spiking all-reduce latency | Rank arrival skew, network congestion, link errors | Compare compute completion times, NIC/switch error counters, and negotiated link rates |
| Loss NaN on a clean batch | Zero supervised tokens or nonfinite gradient under FP16 | Check loss scale factor, token count masks, and inputs |
| Sudden loss divergence | Data or schedule change, numerical instability, software bug, hardware corruption | Reproduce from saved state and inspect inputs before assigning a cause |
A slow rank can delay dependent collectives and extend the critical path. Communication may overlap backward computation, so this isn't a literal barrier where every GPU waits throughout the whole pass. Separate compute time, rank arrival skew, and collective completion time. Compare clock event reasons with a healthy device under matched work; there is no model-independent “bad clock” threshold [18].
Link errors and reduced bandwidth can leave a job alive but slower. Counter trends, link state, and measured throughput help distinguish fabric faults from compute imbalance. A particular rate change doesn't imply a fixed tenfold slowdown, and one counter alone isn't sufficient reason to evict a node.
The Llama 3 report recorded six silent data corruption incidents in its interruption breakdown [3]. Absence of an ECC alert doesn't rule out erroneous computation. Controlled canary computations can help identify faults, but require matched inputs, precision, kernels, and justified tolerances; ordinary floating-point differences aren't proof of bad silicon. A canary doesn't guarantee detecting every error before it affects training.
What a rollback policy must decide
Reject nonfinite training state as a new operational recovery point, and preserve diagnostic evidence. For an unexpected finite loss spike, use thresholds calibrated to the run's normal variability. A median plus three standard deviations over 50 steps is a possible authored heuristic, not an established universal training policy.
Reproduce the suspect batch from a known-good checkpoint, preferably on verified healthy hardware. Check masks, zero-token batches, learning rate, scaler behavior, and gradient norms. Use bounded retries so an automatic restart cannot loop indefinitely.
Choose how far to rewind from evidence about the first affected state. “Two checkpoints back” doesn't guarantee safety. Quarantine data only when the diagnosis supports it, with a recorded decision. Skipping examples or changing the shuffle seed changes the data trajectory; it is a deliberate new continuation policy, not exact recovery.
After a NaN loss, a canary passes on every GPU, but the batch has zero supervised targets. Should the recovery harness automatically quarantine a GPU?
Answer
No. The empty target set can make a mean loss undefined. Check masking and batch construction, fix the data or loss handling, and replay from consistent state. A passing canary doesn't prove every hardware operation is healthy, but this evidence doesn't establish a GPU fault.
Choose the objective, trainable weights, and precision separately
Training recipes aren't a single dropdown choice. SFT and continued pretraining define what the model learns. LoRA defines which parameters update. QLoRA adds 4-bit base-model quantization to adapter tuning. Distillation learns from a teacher, using soft targets [19] or generated sequences [20]; logits aren't required for every distillation method.
| Decision | Examples | What it changes |
|---|---|---|
| Learning signal | Labeled responses; domain text; teacher targets or responses | Optimization objective and data distribution |
| Trainable parameters | Full model weights; LoRA adapters | Parameter memory, optimizer footprint, and capacity |
| Frozen-base precision | BF16; FP8; 4-bit NormalFloat (NF4) | Base model memory, CUDA kernels, numerical behavior |
These are conceptually separate decisions, but their implementations aren't freely interchangeable. Check that the architecture, quantization backend, adapter implementation, and hardware support the combination you chose.
Domain-adaptive pretraining can improve specialized comprehension, but unfamiliar vocabulary alone doesn't prove it's necessary [21]. Small labeled datasets make LoRA SFT an attractive starting point, though adapters aren't guaranteed to beat full fine-tuning [22].

Reject impossible memory plans before launching
The original QLoRA paper demonstrated fine-tuning a 65B model on a single 48GB GPU [23]. That empirical result doesn't mean a 70B model fits within a 24 GiB usable device budget. Here 24 and 48 GiB are hypothetical capacities, not conversions of specific products' advertised GB.
At an ideal four bits per parameter, a 70B model requires: Base weights alone exceed 24 GiB by 8.6 GiB before quantization metadata, higher-precision tensors, LoRA adapters, activations, gradients, optimizer moments, or temporary buffers.
1def weight_gib(parameters, bits):
2 if type(parameters) is not int or parameters <= 0:
3 raise ValueError("positive integer parameter count required")
4 if type(bits) is not int or bits <= 0:
5 raise ValueError("positive integer bit width required")
6 return ((parameters * bits + 7) // 8) / (1024 ** 3)
7
8large = weight_gib(70_000_000_000, 4)
9small = weight_gib(8_000_000_000, 16)
10assert large > 24
11assert small < 24 # Only the base fits this lower-bound check.
12assert weight_gib(1, 4) == 1 / (1024 ** 3)
13print(f"70B at 4 bits: at least {large:.1f} GiB; exceeds 24 GiB")
14print(f"8B at 16 bits: at least {small:.1f} GiB; training needs more")170B at 4 bits: at least 32.6 GiB; exceeds 24 GiB
28B at 16 bits: at least 14.9 GiB; training needs moreThis all-resident four-bit plan cannot fit a 24 GiB budget. Change the plan: use a smaller model, a supported lower-bit representation, sharding, or offloading. Passing the base-only check doesn't establish training fit. For the hypothetical 48 GiB budget, the remaining 15.4 GiB must cover all omitted memory, not just adapters and activations.
Rehearse the restart before the long run
Don't wait for your first midnight preemption to discover that your checkpoint loader crashes. Run a restart drill on small test jobs:
- Local state continuity: Compare five uninterrupted steps against five steps resumed from an intermediate checkpoint. Verify that weights, Adam moments, learning rates, and data sample IDs match within floating-point tolerance.
- Distributed topology restoration: Save on four GPUs, then load on a supported target layout with eight GPUs (or another valid size). Verify model and optimizer resharding, data partitioning, and peak restore memory.
- Preemption deadline drill: Send
kill -TERM <training-pid>to your test job. Exercise both a warning that fits the measured path and one that doesn't. Force an interruption during saving and verify that the loader selects the previous complete checkpoint. - Diagnostics and recovery policy: Inject a controlled straggler or nonfinite test batch. Verify that monitoring identifies the symptom, unsafe state isn't published, retries are bounded, and any rollback or quarantine follows the configured policy.
Distinguish between your latest resumable checkpoint and your best evaluation checkpoint. Step 1840 might be your operational recovery point, while step 1600 holds your highest validation score. Store both with explicit metadata tags.
