Skip to main content
Automatic Mixed Precision (AMP) reduces memory consumption and increases throughput by running forward passes in a lower-precision floating-point format while keeping weights and critical operations in float32. Modern NVIDIA GPUs have dedicated tensor cores that execute half-precision matrix multiplications several times faster than full-precision equivalents. Clorch’s clorch.amp namespace provides the two building blocks you need: the autocast macro for precision narrowing and the grad-scaler record for dynamic loss scaling.

amp/autocast

autocast is a macro that enables LibTorch’s thread-local autocast state for the duration of its body. Eligible operations in the body run in the specified dtype; unsupported operations fall back to float32 automatically. Autocast restores the previous dtype, enabled flag, and cache setting when the body exits, even on error.
Options:

amp/grad-scaler

Float16 has a limited dynamic range; gradients near the minimum representable value underflow to zero. The grad scaler multiplies the loss by a large scale factor before backward, and then divides the gradients back before the optimizer step. If gradients overflow to Inf or NaN, the scaler discards that step and reduces the scale.
Constructor options:
For bfloat16 training, create the scaler with {:enabled? false} or skip it entirely. Bfloat16 has the same dynamic range as float32, so loss scaling is not required.

Scaling and Backward

amp/backward! multiplies the loss by the current scale and calls .backward. It does not unscale gradients; that happens inside amp/step!.

amp/step!

amp/step! performs the following atomically:
  1. Collects all gradients from the optimizer’s parameter list.
  2. Checks whether every gradient is finite (no Inf or NaN).
  3. In distributed training, performs an all-reduce of the finite flag so the decision is consistent across all ranks.
  4. If all gradients are finite, unscales them by dividing by the current scale and calls .step on the optimizer.
  5. Updates the dynamic scale: increases it after growth-interval clean steps, decreases it after an overflow.
  6. Returns true if the optimizer stepped, false if the step was skipped.

float16 vs bfloat16

  • Smaller dynamic range (1e-4 to 65504)
  • Requires dynamic loss scaling to avoid underflow
  • Faster on older Pascal/Volta hardware
  • Use amp/grad-scaler and amp/backward!

Full AMP Training Loop

The following example shows a complete float16 AMP loop including gradient accumulation.

Integration with DDP

When combining AMP and DDP, pass the scaler to ddp/optimizer-step! instead of calling amp/step! directly. This ensures the overflow flag is synchronized across all distributed ranks before any rank steps the optimizer.

Inspecting Scaler State

The scaler state is automatically saved and restored by dist/save-checkpoint! and dist/load-checkpoint!.