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.
- 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
- Transformer Architecture - parameter counts
- Checkpointing & Mixed Precision
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:
| Component | Size | Notes |
|---|---|---|
| Model weights | params × bytes | Dominant |
| FP32 master weights + Adam m, v | 12 bytes × params | Training only |
| Gradients | 2 bytes × params (BF16) | Training only |
| Activations | Scales with batch × seq_len × d_model × n_layers | Reduced by activation (gradient) checkpointing |
| KV cache | 2 × n_layers × n_kv_heads × d_head × bytes × tokens | Inference; 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:
| Format | Bytes/param | Typical use | Where it's covered |
|---|---|---|---|
| FP32 | 4 | Optimizer master weights and states | This chapter |
| BF16 / FP16 | 2 | Training compute and standard inference | This chapter |
| FP8 (E4M3/E5M2) | 1 | Training on Hopper/Blackwell; the first serving choice on those GPUs | Accelerators & Interconnects, Quantized Inference |
| INT8 | 1 | Weight(-and-activation) quantized serving | Quantized Inference |
| INT4 (GPTQ, AWQ) / FP4 (MXFP4, NVFP4) | ~0.5 | Weight-only 4-bit serving; FP4 native on Blackwell | Quantized Inference |
| NF4 | ~0.5 | Frozen base weights in QLoRA fine-tuning | LoRA & 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 Stage | What's partitioned | Memory reduction | Communication overhead |
|---|---|---|---|
| 0 (DDP) | Nothing | 1× | 1 all-reduce |
| 1 | Optimizer states | ~4× | 1 all-reduce + scatter |
| 2 | + Gradients | up to ~8× | Same as 1 (reduce-scatter instead of all-reduce) |
| 3 | + Parameters | N× (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
PCIe vs NVLink
Concept
GPU-to-GPU communication speed determines how well parallelism strategies scale.
| Interconnect | Bandwidth | Latency | Use |
|---|---|---|---|
| PCIe 4.0 ×16 | ~32 GB/s per direction | Medium | Consumer GPUs, budget clusters |
| PCIe 5.0 ×16 | ~64 GB/s per direction | Medium | PCIe H100/L40S servers |
| NVLink 3.0 | ~600 GB/s (total per GPU) | Low | A100 SXM |
| NVLink 4.0 | ~900 GB/s (total per GPU) | Low | H100/H200 SXM |
| NVLink 5.0 | ~1.8 TB/s (total per GPU) | Low | B200/GB200 (NVL72 racks) |
| NVSwitch | Full NVLink bandwidth between every pair of GPUs in a node (or a whole NVL72 rack) | Very low | HGX/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
- How much memory does model state need to train a 13B model with mixed-precision Adam, before activations?
- With ZeRO stage 3 across 8 GPUs, how many bytes of model state per parameter does each GPU hold?
- Why is tensor parallelism usually kept within one NVLink-connected node?
- 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
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.
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
- Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models (2020)
- Micikevicius et al., Mixed Precision Training (2018)
- Dao et al., FlashAttention (2022)
- Shoeybi et al., Megatron-LM (2019)
- NVIDIA, A100 and H100 data sheets (2020-2023)
- PyTorch, FSDP2 (fully_shard) (2026)
Last reviewed: 2026-09