Contents
Map

04 · Pretraining at Scale

GPU Memory & Hardware

View as:

GPU Memory and Hardware Basics

Whether a model trains or serves on your hardware is mostly arithmetic: bytes per parameter for weights, gradients and optimizer state, plus activations and KV cache - and, once it spans several GPUs, how fast they can talk to each other. This chapter gives the memory formulas you should be able to do in your head, the effect of lower precision, ZeRO/FSDP sharding, and the GPU memory hierarchy and interconnects that explain FlashAttention and parallelism choices.

Learning objectives 45 min
By the end of this page you will be able to:
  • Estimate inference and training memory for a model from its parameter count, precision and optimizer (16 bytes/param for mixed-precision Adam)
  • Choose a precision (BF16, FP8, INT8, 4-bit) for a memory budget and name the trade-off
  • Compute per-GPU model-state memory under ZeRO stages 1-3 for N GPUs
  • Explain the GPU memory hierarchy and interconnect bandwidths, and why tensor parallelism stays inside NVLink domains

VRAM Estimation

Concept

Before deploying or training any LLM, you must know whether it fits in your GPU(s). VRAM estimation is a fundamental interview skill.

Model weights memory:

weights_GB = num_parameters × bytes_per_parameter / 10⁹   (divide by 1024³ for GiB - about 7% less)

Bytes per precision:
  FP32:  4 bytes
  BF16:  2 bytes  ← standard for training and most inference
  FP16:  2 bytes
  INT8:  1 byte
  INT4:  0.5 bytes
  NF4:   0.5 bytes

Quick mental math (BF16):

  • 7B model: 7 × 2 = 14 GB
  • 13B model: 13 × 2 = 26 GB
  • 70B model: 70 × 2 = 140 GB
  • 405B model: 405 × 2 = 810 GB → at least 11 × 80 GB GPUs for weights alone (in practice two 8-GPU nodes, or FP8 on one node)

Training overhead (Adam optimizer): Standard mixed-precision training with Adam keeps 16 bytes per parameter of model state (Rajbhandari et al., ZeRO, 2019):

BF16 weights            2 bytes   (used in forward/backward)
BF16 gradients          2 bytes
FP32 master weights     4 bytes   (the optimizer updates these, then re-casts to BF16)
FP32 Adam momentum (m)  4 bytes
FP32 Adam variance (v)  4 bytes
-------------------------------
Model state            16 bytes × num_parameters   (before activations)

  7B model:   7B × 16 = 112 GB  → does not fit one 80 GB GPU; shard with ZeRO/FSDP
  70B model: 70B × 16 = 1.12 TB → ~14+ GPUs of 80 GB just for model state, fully sharded

Activations come on top and scale with batch × sequence length × layers; activation checkpointing trades recompute for most of that memory. Memory-saving optimizers (8-bit Adam, Adafactor) and FP8 training reduce the 16-byte figure, but it remains the standard baseline for estimates.

Additional VRAM components:

ComponentSizeNotes
Model weightsparams × bytesDominant
FP32 master weights + Adam m, v12 bytes × paramsTraining only
Gradients2 bytes × params (BF16)Training only
ActivationsScales with batch × seq_len × d_model × n_layersReduced by activation (gradient) checkpointing
KV cache2 × n_layers × n_kv_heads × d_head × bytes × tokensInference; see KV Cache

Worked example - 70B inference on A100-80GB:

70B × 2 bytes (BF16) = 140 GB → needs 2× A100-80GB minimum
With 4-bit weights: 70B × 0.5 bytes = 35 GB → the weights fit on 1× A100-80GB, leaving ~40 GB for KV cache and activations

Lower-Precision Weights

Memory scales with bytes per parameter, so lower precision is the first lever for fitting a model:

FormatBytes/paramTypical useWhere it's covered
FP324Optimizer master weights and statesThis chapter
BF16 / FP162Training compute and standard inferenceThis chapter
FP8 (E4M3/E5M2)1Training on Hopper/Blackwell; the first serving choice on those GPUsAccelerators & Interconnects, Quantized Inference
INT81Weight(-and-activation) quantized servingQuantized Inference
INT4 (GPTQ, AWQ) / FP4 (MXFP4, NVFP4)~0.5Weight-only 4-bit serving; FP4 native on BlackwellQuantized Inference
NF4~0.5Frozen base weights in QLoRA fine-tuningLoRA & QLoRA Hands-On

4-bit formats store a scale per group of weights, so real sizes are a little above 0.5 bytes per parameter; quality loss depends on the method and the task, so always re-evaluate a quantized model on your own evaluation set.


Sharding Model State Across GPUs

When model state doesn't fit on one GPU, it is split across GPUs. The parallelism dimensions - data, tensor, pipeline, context and expert - are covered in Distributed Training at Scale; here is the memory arithmetic behind the most common one, ZeRO/FSDP sharding of data-parallel state.

ZeRO - Zero Redundancy Optimizer

Concept

ZeRO (Rajbhandari et al., 2019) eliminates redundant copies of optimizer states, gradients, and parameters across GPUs in data-parallel training. Three stages:

Stage 1 - Optimizer State Partitioning:
  Each GPU stores only 1/N of optimizer states (momentum, variance)
  Parameters and gradients: still replicated on all GPUs
  Memory reduction: up to ~4× as N grows (4 + 12/N bytes/param; optimizer state is 12 of the 16)

Stage 2 - Gradient Partitioning:
  Each GPU stores only 1/N of gradients during backward pass
  Parameters: still replicated
  Memory reduction: approaches 8× vs DDP as N grows (2 + 14/N bytes/param)

Stage 3 - Parameter Partitioning:
  Each GPU stores only 1/N of parameters
  Parameters are gathered via all-gather as needed during forward/backward
  Memory reduction: N× vs DDP (16/N bytes/param) - grows with the number of GPUs, enabling models far larger than one GPU's memory

ZeRO-Infinity: Extends ZeRO-3 to offload to CPU RAM and NVMe SSDs - can train trillion-parameter models on limited GPU clusters.

ZeRO StageWhat's partitionedMemory reductionCommunication overhead
0 (DDP)Nothing1×1 all-reduce
1Optimizer states~4×1 all-reduce + scatter
2+ Gradientsup to ~8×Same as 1 (reduce-scatter instead of all-reduce)
3+ ParametersN× (number of GPUs)~1.5× DDP (all-gather in forward and backward)

DeepSpeed implements ZeRO; PyTorch's FSDP2 (fully_shard) provides the same ZeRO-3-style sharding natively (see PyTorch for LLMs).


Flash Attention as Hardware Optimization

Concept

Flash Attention is primarily a hardware (GPU memory hierarchy) optimization - see Attention Mechanisms for the algorithm. Here is the hardware context:

GPU memory hierarchy:

Registers:            256 KB per SM (A100)
L1 / shared memory:   192 KB per SM on A100 (228 KB on H100) - the on-chip SRAM, ~19 TB/s aggregate
HBM:                  80 GB on A100 at ~2 TB/s (H100: 3.35 TB/s) - roughly 10× less bandwidth than SRAM

Standard attention writes the n×n attention matrix to HBM (slow), reads it back for softmax (slow), writes softmax output (slow), reads for ×V (slow). Flash Attention keeps all intermediate results in SRAM by tiling - eliminates the HBM round-trips.

Why this matters at scale:

  • For 128K tokens: one head's attention matrix = 128K × 128K × 2 bytes ≈ 34 GB - impossible to hold anywhere on-chip; FlashAttention never materializes it
  • Flash Attention makes long-context models practical, not just mathematically possible

Concept

GPU-to-GPU communication speed determines how well parallelism strategies scale.

InterconnectBandwidthLatencyUse
PCIe 4.0 ×16~32 GB/s per directionMediumConsumer GPUs, budget clusters
PCIe 5.0 ×16~64 GB/s per directionMediumPCIe H100/L40S servers
NVLink 3.0~600 GB/s (total per GPU)LowA100 SXM
NVLink 4.0~900 GB/s (total per GPU)LowH100/H200 SXM
NVLink 5.0~1.8 TB/s (total per GPU)LowB200/GB200 (NVL72 racks)
NVSwitchFull NVLink bandwidth between every pair of GPUs in a node (or a whole NVL72 rack)Very lowHGX/DGX nodes, NVL72

Practical impact:

  • Tensor parallelism requires all-reduce after every layer: high bandwidth demand. PCIe is the bottleneck - tensor parallelism scales poorly without NVLink.
  • Pipeline parallelism only passes activations at layer boundaries: lower bandwidth requirement - viable with PCIe.
  • A100 SXM4 (NVLink) vs A100 PCIe: tensor-parallel scaling efficiency drops from ~95% to ~50% at 8 GPUs.

Code

# VRAM estimation utility
def estimate_vram(
    num_params_billions,
    precision="bf16",
    mode="inference",
    kv_cache_tokens=4096,
    n_layers=32,
    n_kv_heads=8,
    d_head=128,
    batch_size=1
):
    """Estimate VRAM requirements in GB."""
    bytes_per_param = {"fp32": 4, "bf16": 2, "fp16": 2, "int8": 1, "int4": 0.5, "nf4": 0.5}
    bpp = bytes_per_param[precision]
    
    weight_gb = num_params_billions * 1e9 * bpp / 1e9
    
    if mode == "training":
        # Mixed-precision Adam: BF16 weights (2) + BF16 grads (2)
        # + FP32 master weights (4) + FP32 m (4) + FP32 v (4) = 16 bytes/param
        state_gb = num_params_billions * 16
        return {"weights": weight_gb, "model_state_excl_activations": state_gb}
    else:
        kv_gb = (n_layers * kv_cache_tokens * 2 * n_kv_heads * d_head * 2 * batch_size) / 1e9
        return {"weights": weight_gb, "kv_cache": kv_gb, "total_approx": weight_gb + kv_gb}

# Examples
print("LLaMA-3 8B inference (BF16, 4K context):")
print(estimate_vram(8, "bf16", "inference"))

print("\nLLaMA-3 70B inference (INT4, 4K context):")
print(estimate_vram(70, "int4", "inference"))

print("\nA 7B model, mixed-precision training:")
print(estimate_vram(7, "bf16", "training"))

# BitsAndBytes quantization
from transformers import AutoModelForCausalLM, BitsAndBytesConfig
import torch

config_8bit = BitsAndBytesConfig(load_in_8bit=True)
config_4bit = BitsAndBytesConfig(
    load_in_4bit=True,
    bnb_4bit_quant_type="nf4",
    bnb_4bit_compute_dtype=torch.bfloat16,
    bnb_4bit_use_double_quant=True
)

# Load in 8-bit (comment out when running, requires GPU + large model)
# model_8bit = AutoModelForCausalLM.from_pretrained(
#     "meta-llama/Llama-3.2-1B",
#     quantization_config=config_8bit
# )
# print(f"8-bit model memory: {model_8bit.get_memory_footprint() / 1e9:.2f} GB")

# GPU memory profiling
if torch.cuda.is_available():
    print(f"GPU: {torch.cuda.get_device_name()}")
    print(f"Total VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB")
    free, total = torch.cuda.mem_get_info()
    print(f"Free now: {free / 1e9:.1f} GB; reserved by PyTorch: {torch.cuda.memory_reserved(0) / 1e9:.2f} GB")

Hands-On Lab

Want to apply this VRAM/quantization math to an actual training run? See the Fine-Tuning Lab module - configuring BitsAndBytesConfig for real 4-bit QLoRA loading, and a benchmark script that reads measured peak VRAM (torch.cuda.max_memory_allocated()) rather than estimating it.


Study Notes

  • Weights = params × bytes (BF16 2, FP8/INT8 1, 4-bit ~0.5); 70B in BF16 = 140 GB
  • Mixed-precision Adam model state = 16 bytes/param (BF16 weights 2 + grads 2, FP32 master 4 + m 4 + v 4), before activations
  • ZeRO per-GPU bytes/param: stage 1 = 4 + 12/N, stage 2 = 2 + 14/N, stage 3 = 16/N; FSDP2 is PyTorch's native ZeRO-3
  • Activations scale with batch × sequence × width × depth; activation checkpointing trades recompute for most of them
  • On-chip SRAM is ~10× faster than HBM - FlashAttention's whole trick; decode is memory-bandwidth bound
  • NVLink (600-1,800 GB/s per GPU) vs PCIe (32-64 GB/s per direction): tensor parallelism stays inside NVLink domains

Check Yourself

Check yourself
0 / 4 answered
  1. How much memory does model state need to train a 13B model with mixed-precision Adam, before activations?
  2. With ZeRO stage 3 across 8 GPUs, how many bytes of model state per parameter does each GPU hold?
  3. Why is tensor parallelism usually kept within one NVLink-connected node?
  4. A 70B model is quantized to 4-bit weights for serving on one 80 GB GPU. What still limits how many users it can serve?

Exercises

Exercise - Size a fine-tuning job

You want to fully fine-tune an 8B model with AdamW in bf16 mixed precision on 8 × 80 GB GPUs, sequence length 4,096, micro-batch 1 per GPU. Estimate per-GPU model-state memory under DDP, ZeRO-2 and ZeRO-3/FSDP, and say which options fit once you leave ~20 GB for activations.

Solution

Model state 8B × 16 = 128 GB total. DDP: 128 GB per GPU - does not fit. ZeRO-2: (2 + 14/8) × 8B = 3.75 × 8 = 30 GB per GPU - fits with room for activations. ZeRO-3: 16/8 × 8B = 16 GB per GPU - fits, at the cost of extra all-gather traffic. Activation checkpointing keeps activations within the ~20 GB at 4K tokens.

Exercise - Measure the memory hierarchy

Benchmark a large element-wise operation (memory-bound) and a large matrix multiplication (compute-bound) on your GPU. Compute achieved bandwidth (GB/s) and throughput (TFLOPS) and compare with the data sheet.

Solution

The element-wise op reaches a large fraction of HBM bandwidth but a tiny fraction of peak FLOPS; the matmul reaches a high fraction of dense tensor-core FLOPS. This is the roofline view: operations with low arithmetic intensity (decode, attention without FlashAttention) are bandwidth-bound.

References

Last reviewed: 2026-09

⚡AI-assisted content - always verify, always explore multiple perspectives·