GPU Memory Optimization для LLM: техники и практики

opensourceaillmgpuperformanceit
← Back to Blog

Введение

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%

Ссылки