GPU Memory Optimization для LLM: техники и практики
opensourceaillmgpuperformanceit
Введение
GPU memory — главный bottleneck при запуске LLM. 70B модель требует ~40GB VRAM в INT4. Как уместить больше в меньше?
GPU Memory для LLM: основы
Сколько памяти нужно?
Модель | FP16 | INT8 | INT4 | FP8
Llama 3 8B | 16GB | 9GB | 5GB | 9GB
Llama 3 70B | 140GB | 75GB | 40GB | 75GB
Mixtral 8x7B | 260GB | 135GB | 70GB | 135GB
Формула:
Parameters × bytes_per_weight + KV_cache + overhead
FP16: 2 bytes/parameter
INT8: 1 byte/parameter
INT4: 0.5 byte/parameter
FP8: 1 byte/parameter
Из чего состоит GPU memory?
1. Model weights:
Llama 3 8B → 16GB (FP16)
Llama 3 70B → 140GB (FP16)
2. KV cache:
Зависит от context length и batch size
Llama 3 8B, ctx=8K, batch=1 → ~2GB
Llama 3 8B, ctx=8K, batch=8 → ~16GB
3. Activation memory:
Временные данные при forward pass
Зависит от batch size и context length
4. Overhead:
CuBLAS workspace, temporary buffers
~5-10% от общего объёма
Техника 1: Quantization
INT8 Quantization
import torch
from transformers import AutoModelForCausalLM
# Загрузка и квантование
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3-8B",
torch_dtype=torch.float16
)
# INT8 quantization
from bitsandbytes import quantization
model = quantization.quantize_model(model, bits=8)
# Экономия: 16GB → 9GB
INT4 Quantization (GGUF)
# Конвертация в GGUF Q4_K_M
python convert-hf-to-gguf.py model.bin --outtype q4_k_m
# Результат:
# Original FP16: 16GB
# GGUF Q4_K_M: 5.5GB
# Качество: ~98% от оригинала
NF4 Quantization (bitsandbytes)
from transformers import BitsAndBytesConfig
quantization_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4", # Normal Float 4-bit
bnb_4bit_compute_dtype=torch.float16,
bnb_4bit_use_double_quant=True, # Double quantization
)
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3-8B",
quantization_config=quantization_config
)
# 8B модель: 16GB → 4.7GB
FP8 Quantization
# FP8 — новый стандарт (Hopper, Ada Lovelace)
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3-8B",
torch_dtype=torch.float8_e4m3fn
)
# FP8_e4m3: 4 exponent bits, 3 mantissa bits
# FP8_e5m2: 5 exponent bits, 2 mantissa bits
# 8B модель: 16GB → 8GB
Сравнение quantization
Метод | Размер | Speed | Quality
FP16 | 100% | 1x | 100%
INT8 | 50% | 1.5x | 99%
FP8 | 50% | 1.8x | 98%
INT4 (NF4) | 25% | 2x | 95%
GGUF Q4_K_M | 30% | 2x | 97%
GGUF Q3 | 22% | 2.2x | 92%
Техника 2: KV Cache Optimization
PagedAttention (vLLM)
# vLLM использует PagedAttention
# Memory management как в OS (pages)
from vllm import LLM
llm = LLM(
model="meta-llama/Llama-3-8B",
gpu_memory_utilization=0.9, # 90% GPU memory для KV cache
max_num_batched_tokens=8192,
max_num_seqs=256,
)
# PagedAttention убирает fragmentation
# 2-4x больше batch size чем стандартный подход
KV Cache Quantization
# Квантование KV cache в INT8
from vllm import LLM
llm = LLM(
model="meta-llama/Llama-3-70B",
cache_dtype="int8", # Квантование KV cache в INT8
)
# Экономия: KV cache в 2x меньше
# 70B модель: 140GB + 30GB(KV) → 75GB + 15GB(KV)
KV Cache Eviction
# Удаление старых токенов из KV cache
# Для длинных контекстов
class KVCacheEviction:
def __init__(self, window_size=4096):
self.window_size = window_size
def forward(self, kv_cache, new_tokens):
if kv_cache.size > self.window_size:
# Удаляем старые блоки
return kv_cache[-self.window_size:]
return kv_cache
# Позволяет работать с context > 32K
Sliding Window Attention
# Только последние N токенов в KV cache
from transformers import LlamaForCausalLM
model = LlamaForCausalLM.from_pretrained(
"meta-llama/Llama-3-8B",
attn_implementation="flash_attention_2",
)
# При context=128K, sliding_window=4096:
# KV cache только для последних 4096 токенов
# Экономия: 128K/4096 = 31x меньше памяти
Техника 3: Gradient Checkpointing
# Trade memory для compute
# Сохраняем activations только для некоторых слоёв
from transformers import AutoModel
model = AutoModel.from_pretrained(
"meta-llama/Llama-3-8B",
gradient_checkpointing=True # Экономия 50-70% memory
)
# Memory: 80GB → 30GB
# Speed: 1x → 0.7x
Техника 4: Mixed Precision Training
# FP16 + FP8 комбинация
import torch.distributed as dist
class MixedPrecisionTrainer:
def __init__(self, model):
self.model = model
self.scaler = torch.amp.GradScaler('cuda')
def forward(self, x):
# Активные слои: FP16
# Остальные: FP8
with torch.amp.autocast('cuda', dtype=torch.float16):
return self.model(x)
# Экономия: 30-50% memory vs FP32
Техника 5: Model Parallelism
Tensor Parallelism
# Split one layer across multiple GPUs
from vllm import LLM
# 70B модель на 4 GPU
llm = LLM(
model="meta-llama/Llama-3-70B",
tensor_parallel_size=4, # 4 GPU
)
# Каждый слой split на 4 части
# 140GB → 4 × 35GB
Pipeline Parallelism
# Split layers across GPUs
# GPU 0: layers 0-11
# GPU 1: layers 12-23
# GPU 2: layers 24-31
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3-70B",
device_map="auto", # Автоматическое распределение
)
# 140GB на 4 GPU → 4 × 35GB + overhead
DeepSpeed ZeRO
# DeepSpeed Zero: split optimizer states, gradients, parameters
import deepspeed
deepspeed_config = {
"zero_optimization": {
"stage": 3, # Zero-3: всё split
"offload_param": {
"device": "cpu" # Offload на CPU
},
"contiguous_gradients": True,
"overlap_comm": True,
"reduce_scatter": True,
"round_robin_gradients": True,
},
"fp16": {
"enabled": True,
"loss_scale_window": 1000
}
}
# 70B на 1 GPU → 70B на 2 GPU (с offload)
Техника 6: Offloading
CPU Offloading
# parts of model on CPU, parts on GPU
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3-8B",
device_map="auto", # Автоматическое распределение
max_memory={0: "10GB", "cpu": "20GB"},
)
# GPU: 10GB (weight parts)
# CPU: 20GB (остальные parts)
# Медленнее, но работает на слабом GPU
Ollama CPU offload
# Ollama автоматически offloads на CPU
# Когда GPU memory full
OLLAMA_NUM_GPU=35 ollama run llama3.2
# 35 слоёв на GPU, остальные на CPU
HuggingFace accelerate
from accelerate import Accelerator
accelerator = Accelerator(
mixed_precision="fp16",
cpu_offload=True,
)
model, optimizer = accelerator.prepare(model, optimizer)
# Автоматическое управление memory
Техника 7: Batch Size Optimization
# Меньше batch = меньше memory
# Large batch (OOM risk)
batch_size = 64 # Требует 32GB VRAM
# Small batch (memory efficient)
batch_size = 8 # Требует 8GB VRAM
# Gradient accumulation для компенсации
accumulation_steps = 8 # 8 × 8 = 64 effective batch
for i, batch in enumerate(dataloader):
with torch.amp.autocast('cuda'):
outputs = model(**batch)
loss = criterion(outputs, labels) / accumulation_steps
loss.backward()
if (i + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()
Техника 8: Flash Attention
# Flash Attention: O(sqrt(N)) memory vs O(N)
# Обычный attention
# Attention(Q, K, V) = softmax(QK^T / sqrt(d))V
# Memory: O(batch × seq_len^2 × hidden)
# Flash Attention
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3-8B",
attn_implementation="flash_attention_2",
)
# Memory: O(batch × seq_len × sqrt(hidden))
# Экономия: 5-10x для длинных контекстов
Техника 9: Activation Checkpointing
# Не сохраняем все activations, только checkpoint'ы
model.gradient_checkpointing_enable(
gradient_checkpointing_kwargs={"use_reentrant": False}
)
# Memory saving: 40-60%
# Speed penalty: 10-20%
Техника 10: Memory-Efficient Architectures
Mamba / SSM
# State Space Models: O(1) memory vs O(n) для attention
from mamba_ssm import Mamba
model = Mamba(
d_model=2048,
d_state=16,
d_conv=4,
expand=2,
)
# Memory: constant (не зависит от seq_len)
# Llama 8B: KV cache = O(seq_len)
# Mamba: KV cache = O(1)
Grouped Query Attention (GQA)
# Fewer key-value heads = smaller KV cache
# Multi-Query (1 KV head):
# KV cache: 1/64 от обычного
# Grouped-Query (8 KV heads):
# KV cache: 8/64 от обычного
from transformers import AutoConfig
config = AutoConfig.from_pretrained("meta-llama/Llama-3-8B")
config.num_key_value_heads = 8 # GQA вместо MHA
# KV cache: 64 → 8 heads = 8x меньше
Практические рекомендации
Для 24GB GPU (RTX 4090)
Llama 3 8B:
- GGUF Q4_K_M: ✅ Отлично (5.5GB weights + KV)
- GGUF Q5_K_M: ✅ Хорошо (6.5GB)
- GGUF Q6_K: ✅ Хорошо (7.5GB)
- GGUF Q8_0: ⚠️ На пределе (9GB)
Llama 3 70B:
- GGUF Q4_K_M: ❌ Не влезает (40GB)
- Требуется: 2× GPU или CPU offload
Для 48GB GPU (RTX 6000)
Llama 3 8B:
- GGUF Q8_0: ✅ Отлично (9GB)
- Large batch: ✅ (batch=32+)
Llama 3 70B:
- GGUF Q4_K_M: ✅ Хорошо (40GB)
- Batch=1-2: ✅
Llama 3 70B Q8:
- GGUF Q8_0: ⚠️ На пределе (75GB)
Для 80GB GPU (A100)
Llama 3 70B:
- GGUF Q4_K_M: ✅ Отлично (40GB)
- GGUF Q8_0: ✅ Хорошо (75GB)
- FP16: ⚠️ На пределе (140GB - нет)
Llama 3 70B FP16:
- Tensor parallelism (2× A100): ✅
Для 80GB×4 (4× A100)
Llama 3 70B:
- FP16: ✅ Отлично (140GB / 4 = 35GB/GPU)
- INT8: ✅ (75GB / 4 = 19GB/GPU)
- INT4: ✅ (40GB / 4 = 10GB/GPU)
Mixtral 8x7B:
- FP16: ✅ (260GB / 4 = 65GB/GPU)
Мониторинг GPU Memory
# PyTorch memory tracking
import torch
# Начало
torch.cuda.reset_peak_memory_stats()
# ... training/inference ...
print(f"Allocated: {torch.cuda.memory_allocated() / 1e9:.2f} GB")
print(f"Cached: {torch.cuda.memory_reserved() / 1e9:.2f} GB")
print(f"Peak: {torch.cuda.max_memory_allocated() / 1e9:.2f} GB")
# nvidia-smi
nvidia-smi --query-gpu=memory.used,memory.total --format=csv
# Динамический мониторинг
watch -n 1 nvidia-smi
# PyNVML
python -c "
import pynvml
pynvml.nvmlInit()
handle = pynvml.nvmlDeviceGetHandleByIndex(0)
mem = pynvml.nvmlDeviceGetMemoryInfo(handle)
print(f'Used: {mem.used / 1e9:.2f} GB / {mem.total / 1e9:.2f} GB')
"
Чек-лист оптимизации
1. [ ] Quantize weights (INT4/INT8/FP8)
2. [ ] Use Flash Attention
3. [ ] Enable gradient checkpointing
4. [ ] Reduce batch size + accumulation
5. [ ] Use PagedAttention (vLLM)
6. [ ] KV cache quantization
7. [ ] Tensor parallelism (multi-GPU)
8. [ ] CPU offloading (если нужно)
9. [ ] Use GQA/MQA модели
10. [ ] Monitor memory usage
Итоги
| Техника | Экономия | Сложность | Влияние на speed |
|---|---|---|---|
| INT4 Quantization | 75% | ⭐ | +20% |
| INT8 Quantization | 50% | ⭐ | +50% |
| FP8 Quantization | 50% | ⭐⭐ | +80% |
| KV Cache Quant | 50% KV | ⭐⭐ | +10% |
| PagedAttention | 40-60% KV | ⭐ | +20% |
| Flash Attention | 5-10x KV | ⭐ | +30% |
| Gradient Checkpoint | 50% | ⭐ | -10% |
| Tensor Parallelism | N/GPU | ⭐⭐⭐ | -5% |
| CPU Offload | ∞ | ⭐ | -50% |
| Batch Size ↓ | Прямая | ⭐ | -30% |