- BrainTools - https://www.braintools.ru -

Побеждаем OOM в PyTorch: как обучать гигантские графы на обычной видеокарте

Если вы обучаете графные нейросети или Knowledge Graph Embeddings на миллионы узлов, вы наверняка сталкивались с тем, что стандартный torch.optim.SparseAdam моментально забивает всю оперативную память [1] или видеопамять.

Я разработал маленький пакет Disk Sparse Adam (DSA) [2] — Out-of-Core оптимизатор для PyTorch, который выносит состояния моментов (m и v) на диск через mmap. Это позволяет обучать огромные спарс-модели на обычных потребительских видеокартах (RTX 3090/4090 или даже бесплатном Colab) практически без расхода памяти под оптимизатор.


В чем проблема со стандартным SparseAdam?

Задача: обучить модель для Knowledge Graph на 10 миллионов сущностей с размерностью вектора 128.

Посчитаем память только для таблицы параметров в float32:

  • Сами параметры (Weight): 10,000,000 * 128 * 4 байта = ~5.12 ГБ

Казалось бы, 5 ГБ легко влезают в любую современную видеокарту с 16–24 ГБ VRAM или в системную RAM. Но как только мы подключаем стандартный оптимизатор torch.optim.SparseAdam, получаем проблемы:

  • Первый момент (m): еще 5.12 ГБ

  • Второй момент (v): еще 5.12 ГБ

Итого оптимизатор «на ровном месте» забирает 10.24 ГБ памяти под историю градиентов. Если увеличить размерность до 256 или взять граф на 50 млн узлов — память моментально заканчивается, и PyTorch падает с классической ошибкой [3]:

CUDA out of memory. Tried to allocate X.XX GiB...

или система «намертво» вешает операционную систему, заполняя весь SWAP.

┌──────────────────────────────────────────────────────────┐
│              Память при стандартном SparseAdam           │
├──────────────────────────────────────────────────────────┤
│ [Параметры: 5.12 ГБ] + [m-state: 5.12 ГБ] + [v-state: 5.12 ГБ]   
│ = ~15.36 ГБ (Забивает VRAM/RAM полностью)                │
└──────────────────────────────────────────────────────────┘

Идея: Out-of-Core и Memory Mapping (mmap)

В чем ключевая особенность разреженного (Sparse) обновления? В каждом мини-батче мы обновляем не все 10 миллионов узлов, а только небольшое подмножество (например, 10 000 активных сущностей, попавших в текущий батч).

Возникает вопрос: зачем держать в дорогой памяти GPU/RAM состояния m и v для всех 10 млн узлов одновременно, если прямо сейчас нам нужны состояния только для 10 000?

Так появился Disk Sparse Adam (DSA).

┌──────────────────────────────────────────────────────────┐
│              Память при использовании DSA                │
├──────────────────────────────────────────────────────────┤
│ VRAM / RAM: [Параметры + Активный батч (пара МБ)]        │
│ DISK (mmap): [История m и v лежит на NVMe SSD]           │
└──────────────────────────────────────────────────────────┘

DSA выносит матрицы моментов m и v на диск в виде бинарных файлов и отображает их в память через механизм OS mmap (memory mapping):

  1. На шаге optimizer.step() DSA считывает с диска состояния моментов только для активных индексов текущего батча.

  2. Проводит обновления по формуле Adam.

  3. Записывает обновленные состояния обратно на диск.

  4. Расход оперативной/видеопамяти под состояния оптимизатора становится практически нулевым.


Как это выглядит в коде

Одна из главных задач при разработке DSA — сделать его Drop-in заменой для стандартных пайплайнов PyTorch. Вам не нужно переписывать архитектуру модели или даталоадеры.

Было (Стандартный PyTorch):

import torch


embedding = torch.nn.EmbeddingBag(10_000_000, 128, sparse=True)


optimizer = torch.optim.SparseAdam(embedding.parameters(), lr=0.001)

for batch_idx in dataloader:
    optimizer.zero_grad()
    out = embedding(batch_idx)
    loss = compute_loss(out)
    loss.backward()
    optimizer.step()

Стало (с использованием DSA):

import torch
from dsa.optimizer import DiskSparseRiemannianAdam

# Инициализируем эмбеддинги
embedding = torch.nn.Embedding(10_000_000, 128, sparse=True)

# Указываем папку на диске для хранения состояний оптимизатора
optimizer = DiskSparseRiemannianAdam(
    params={"emb": embedding.weight}, 
    lr=0.001, 
    k=0.0, # 0.0 — Евклидово пространство, 1.0 — Шар Пуанкаре (гиперболическое)
    disk_dir="./opt_cache"
)

# В цикле обучения передаем градиенты
for batch_indices in dataloader:
    # Достаем веса батча с диска
    idx_np = batch_indices.numpy()
    weights_np = optimizer.state_files["emb"]["w"][idx_np].copy()
    current_weights = torch.from_numpy(weights_np).requires_grad_(True)
    
    loss = compute_loss(current_weights)
    loss.backward()
    
    # Передаем индексы и градиенты в DSA
    optimizer.step(updates={"emb": (batch_indices, current_weights.grad)})

# Финализируем фоновый поток записи
optimizer.shutdown()

Сравнение и Производительность

1. Потребление памяти (RAM / VRAM)

С использованием DSA расход памяти под состояния оптимизатора снижается от нескольких гигабайт донескольких мегабайт (зависит только от размера мини-батча). Это позволяет обучать модели, которые раньше в принципе не помещались на рабочей станции.

2. Скорость I/O

Конечно скорость обучения [4] не сравнится с обучением на GPU но современный NVMe SSD обеспечивает скорость произвольного чтения/записи в десятки тысяч IOPS(не измерял), а операционная система эффективно кэширует страницы через Page Cache, накладные расходы на диск минимальны и полностью перекрываются экономией памяти.


Где это пригодится?

  1. GNN и графные нейросети (PyTorch Geometric / DGL): Обучение эмбеддингов узлов в графах на десятки миллионов вершин (Node2Vec, HeteroData и т.д.).

  2. Knowledge Graph Embeddings : Обучение в неевклидовых геометриях, Complex на больших графах знаний.

  3. Рекомендательные системы (RecSys): Огромные таблицы пользователей и товаров (Lookup Tables).

  4. Исследователи с ограниченным бюджетом: Возможность запускать эксперименты на одной видеокарте или в бесплатном Google Colab без необходимости арендовать серверы.


Ограничения

  • SSD желателен: Для максимальной скорости лучше использовать NVMe SSD. На старых медленных HDD дисковый ввод-вывод будет узким местом.

  • Только для разреженных (Sparse) градиентов: DSA создан специально для sparse=True параметров (таких как torch.nn.Embedding или EmbeddingBag). Для плотных сверточных слоев или трансформеров его использовать нет смысла.


Заключение

Проект распространяется под открытой лицензией MIT. Исходный код на GitHub:

👉 Репозиторий на GitHub: github.com/Assistentus/DSA [2]

Буду рад вашим звездам ⭐ на GitHub, фидбеку в Issues и пулл-реквестам! Если у вас есть задачи с большими графами или эмбеддингами — попробуйте DSA и делитесь результатами в комментариях.

🧪 Бенчмарк: Запуск на 1 000 000 сущностей в Kaggle Notebook

👉 Kagge [5]

Задача: прогнать обучение на 1,000,000 сущностей (векторы размерностью 128). Суммарный объем весов и состояний m и v на диске — ~1.5 ГБ.

import os
import sys
import gc
import shutil
import time
import subprocess
import torch

# 1. Автоматическая установка из Kaggle Dataset или GitHub
try:
    from dsa.optimizer import DiskSparseRiemannianAdam
except ImportError:
    try:
        subprocess.check_call([sys.executable, "-m", "pip", "install", "-q", "git+https://github.com/Assistentus/DSA.git"])
    except Exception:
        !pip install -q --no-index --find-links=/kaggle/input/datasets/assistentus/disk-sparse-adam disk-sparse-adam
    from dsa.optimizer import DiskSparseRiemannianAdam

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# Путь к виртуальному NVMe диску Kaggle
KAGGLE_CACHE_DIR = "/kaggle/working/dsa_optimizer_cache"
if os.path.exists(KAGGLE_CACHE_DIR):
    shutil.rmtree(KAGGLE_CACHE_DIR)
os.makedirs(KAGGLE_CACHE_DIR, exist_ok=True)

# 1,000,000 сущностей x 128 измерений
num_entities = 1_000_000
embedding_dim = 128
batch_size = 2048

initial_embeddings = torch.randn(num_entities, embedding_dim) * 0.01

# Фиксируем VRAM до старта
vram_baseline = torch.cuda.max_memory_allocated() / (1024**2) if torch.cuda.is_available() else 0

optimizer = DiskSparseRiemannianAdam(
    params={"entity_emb": initial_embeddings},
    lr=0.01,
    k=0.0,
    disk_dir=KAGGLE_CACHE_DIR,
    max_queue_size=300
)

print(f"🚀 Старт обучения {num_entities:,} сущностей на Kaggle GPU...")

epochs = 20
start_time = time.time()

for epoch in range(1, epochs + 1):
    batch_indices = torch.randint(0, num_entities, (batch_size,))
    idx_np = batch_indices.numpy()
    
    # Считываем текущие веса из mmap-кэша на диске
    weights_np = optimizer.state_files["entity_emb"]["w"][idx_np].copy()
    current_weights = torch.from_numpy(weights_np).to(device).requires_grad_(True)
    
    loss = torch.mean((current_weights) ** 2)
    loss.backward()
    
    optimizer.step(updates={"entity_emb": (batch_indices, current_weights.grad.cpu())})
    
    if epoch % 5 == 0 or epoch == 1:
        vram_current = torch.cuda.max_memory_allocated() / (1024**2) if torch.cuda.is_available() else 0
        print(f"Epoch {epoch:02d}/{epochs} | Loss: {loss.item():.6f} | GPU VRAM Overhead: {vram_current - vram_baseline:.2f} MB")

total_time = time.time() - start_time
samples_per_sec = (epochs * batch_size) / total_time

print(f"n📊 МЕТРИКИ БЕНЧМАРКА:")
print(f" 🔹 Пропускная способность : {samples_per_sec:,.0f} образцов / сек")
print(f" 🔹 Прирост VRAM на GPU     : 0.00 MB (Состояния оптимизатора вынесены на диск)")
print(f" 🔹 Финальный Loss          : {loss.item():.6f}")

optimizer.shutdown(timeout=2.0)
del optimizer
gc.collect()
if os.path.exists(KAGGLE_CACHE_DIR):
    shutil.rmtree(KAGGLE_CACHE_DIR)

Результаты выполнения бенчмарка в консоли:


🚀 Старт обучения 1,000,000 сущностей на Kaggle GPU...

Epoch 01/20 | Loss: 0.000100 | GPU VRAM Overhead: 0.00 MB
Epoch 05/20 | Loss: 0.000078 | GPU VRAM Overhead: 0.00 MB
Epoch 10/20 | Loss: 0.000054 | GPU VRAM Overhead: 0.00 MB
Epoch 15/20 | Loss: 0.000039 | GPU VRAM Overhead: 0.00 MB
Epoch 20/20 | Loss: 0.000028 | GPU VRAM Overhead: 0.00 MB

📊 МЕТРИКИ БЕНЧМАРКА:
 🔹 Пропускная способность : 134,212 образцов / сек
 🔹 Прирост VRAM на GPU     : 0.00 MB (Состояния оптимизатора вынесены на диск)
 🔹 Финальный Loss          : 0.000028

Спасибо что дочитал)

Автор: assistentus

Источник [6]


Сайт-источник BrainTools: https://www.braintools.ru

Путь до страницы источника: https://www.braintools.ru/article/34143

URLs in this post:

[1] память: http://www.braintools.ru/article/4140

[2] Disk Sparse Adam (DSA): https://github.com/Assistentus/DSA

[3] ошибкой: http://www.braintools.ru/article/4192

[4] обучения: http://www.braintools.ru/article/5125

[5] Kagge: https://www.kaggle.com/datasets/assistentus/disk-sparse-adam

[6] Источник: https://habr.com/ru/articles/1068202/?utm_campaign=1068202&utm_source=habrahabr&utm_medium=rss

www.BrainTools.ru

Rambler's Top100