Skip to content

CUDA status ​

CUDA is supported for training and inference on Linux, on a single GPU. It is not available on macOS.

Tested on: Linux x86_64, RTX 4090 (compute capability 8.9), CUDA Toolkit 13.3, driver 610.43.

For the comparison with CPU, Metal and Vulkan, see SUPPORT.md.

Building ​

Enable the Cargo feature cuda:

sh
cargo build --features cuda
VariableDefaultEffect
RETRO_CUDA_ARCHITECTURESnativeGPU architectures to compile for, for example 89-real
RETRO_CUDA_GRAPHSoffEnable CUDA Graphs capture
GGML_CUDA_DEQUANT_BUDGET_MB64F32 scratch ceiling of the sliced OUT_PROD dequantization; 0 or less means no slicing

NCCL is always disabled.

What works on the GPU ​

The backward kernels below run on CUDA. Anything not listed falls back to the CPU.

KernelSupported
OUT_PRODF32, F16, BF16 and 23 quantized types
FUSED_SPARSE_CE and its backwardF32, F16, Q4_0, Q4_1, Q5_0, Q5_1, Q8_0, Q2_K to Q6_K
FLASH_ATTN_BACKhead_dim up to 512; K/V in F16 or F32; GQA; softcap; no attention sinks
SSM_CONV_BACK, SSM_SCAN_BACKcontiguous F32
SILU_BACK, RMS_NORM_BACK, SOFT_MAX_BACK, GET_ROWS_BACK, CROSS_ENTROPY_LOSS_BACKyes
CONV_RS_GATHERyes, used when RETRO_RECURRENT_ROLLBACK=auto
AdamW, SGDF16 and BF16 parameters, with stochastic rounding identical to the CPU
CPY F32 -> F16/BF16the store write of a master-copy step, bit for bit what the reference conversion writes

The kernel picks the smallest register bucket covering the head dimension - 128, 256 or 512 (FA_BACK_MAX_D); a wider head, or a model with attention sinks, takes the materialized F32 backward graph instead.

The runtime detects these capabilities from the device itself. backend_report shows them in gpu_device, cap_flash_attn_back and cap_device_sampling.

Base weights in F16 and BF16 ​

A GGUF whose matrices are stored as F16 or BF16 trains on CUDA in place, under AdamW or SGD, with no conversion of the file. The gradient and both AdamW moments stay F32; only the parameter is half, and the update is rounded stochastically from a stream seeded by the optimizer's iteration counter, so a resume lands on the same weights bit for bit. The numbers below are measured against each fixture's own F32 twin, on the hardware named above. Under AdamW:

CPU / F16CUDA / F16CPU / BF16CUDA / BF16
worst relative gradient gap, one step over 24576 elements5.5e-41.7e-35.5e-34.9e-3
elements whose update is more than one grid point from the F32 trajectory1.2e-32.1e-32.3e-32.2e-3
relative loss gap after 2000 steps2.1e-47.9e-45.9e-31.5e-3

And under SGD, whose step is the gradient itself rather than a normalized one:

CPU / F16CUDA / F16CPU / BF16CUDA / BF16
worst relative gradient gap, one step over 24576 elements5.5e-41.7e-35.5e-34.9e-3
elements whose update is more than one grid point from the F32 trajectory5.7e-43.3e-38.5e-46.5e-4
relative loss gap after 2000 steps2.2e-32.3e-56.3e-35.0e-3

The gradient row is the same under both optimizers: it comes out of the forward, which does not know which step will read it. In F16, CUDA's gaps are about three times the CPU's because the forward feeds them differently; in BF16 the two backends land together, since the storage costs an order of magnitude more than that difference does.

Through an F32 master copy ​

With training.master_weights, the update step runs in F32 on a per-parameter copy and a single CPY F32 -> store writes the weight, so no half-precision update kernel is involved at all. The store cast is bit for bit the reference conversion, and every stored element is exactly the master value rounded once. Against the same F32 twin:

CUDA / F16CUDA / BF16
worst relative gradient gap, one step over 24576 elements1.7e-34.9e-3
elements whose update is more than one grid point from the F32 trajectory, AdamW1.9e-32.0e-3
the same under SGD2.6e-37.3e-4
relative loss gap after 2000 steps, AdamW1.9e-42.0e-3
the same under SGD5.8e-35.2e-3

The gradient row is the in-place path's, unchanged: the master copy touches only the update.

Limitations ​

  • One GPU only. The runtime uses CUDA0. No multi-GPU, sharding or peer-copy.
  • No CI. CUDA tests are run by hand with the commands below. Run them before a PR that touches flash-attn-back.cu or out-prod.cu.
  • checkpoint_dtype F16/BF16 saves memory but costs about 1e-3 relative error. Use it only when a run does not fit otherwise.
  • VRAM numbers. In the memory report, scratch_* is the memory used by this process. device_* is for the whole GPU, so only compare it between two points in time.
  • A quantized OUT_PROD does not always decode in place. Wide projections use a dequantize+SGEMM path whose F32 scratch is bounded by GGML_CUDA_DEQUANT_BUDGET_MB; small and strided outputs keep the scratch-free in-place decoder. So the budget moves a step's scratch peak by design, and it bounds the scratch without changing any gradient.

Running the tests ​

sh
# Device detection
RETRO_CUDA_ARCHITECTURES=89-real cargo test --features cuda --test backend_devices

# Kernel parity against the CPU
RETRO_CUDA_ARCHITECTURES=89-real cargo test --features cuda --test cuda_backend -- --test-threads=1

# Same, with CUDA Graphs
RETRO_CUDA_ARCHITECTURES=89-real RETRO_CUDA_GRAPHS=1 cargo test --features cuda --test cuda_backend -- --test-threads=1

# CPU non-regression
cargo test --no-default-features --features agent

# Half-precision base weights, CPU column and CUDA column in one run
RETRO_TINY_FIXTURE=tests/fixtures/retrograd-tiny-qwen2-f32.gguf \
RETRO_TINY_F16_FIXTURE=tests/fixtures/retrograd-tiny-qwen2-f16.gguf \
RETRO_TINY_BF16_FIXTURE=tests/fixtures/retrograd-tiny-qwen2-bf16.gguf \
RETRO_TINY_BF16_CONTROL_FIXTURE=tests/fixtures/retrograd-tiny-qwen2-bf16ctl-f32.gguf \
  cargo test --features cuda --test f16_base_training -- --test-threads=1 --nocapture

# The device-memory guard (discrete GPU only): on unified memory a "device"
# allocation is a host allocation.
RETRO_REQUIRE_GPU_RESIDENT=1 \
RETRO_CPU_FIXTURE=tests/fixtures/LFM2.5-230M-Q4_K_M.gguf \
RETRO_TINY_FIXTURE=tests/fixtures/retrograd-tiny-qwen2-f32.gguf \
  cargo test --release --features cuda --test device_memory --test gefen_ops -- --test-threads=1

Replace 89-real with the compute capability of your GPU.

Retrograd documentation