Побеждаем OOM в PyTorch: как обучать гигантские графы на обычной видеокарте. Disk Sparse Adam.. Disk Sparse Adam. gnn.. Disk Sparse Adam. gnn. machine learning.. Disk Sparse Adam. gnn. machine learning. oom.. Disk Sparse Adam. gnn. machine learning. oom. Open source.. Disk Sparse Adam. gnn. machine learning. oom. Open source. python.. Disk Sparse Adam. gnn. machine learning. oom. Open source. python. PyTorch.. Disk Sparse Adam. gnn. machine learning. oom. Open source. python. PyTorch. SparseAdam.. Disk Sparse Adam. gnn. machine learning. oom. Open source. python. PyTorch. SparseAdam. глубокое обучение.. Disk Sparse Adam. gnn. machine learning. oom. Open source. python. PyTorch. SparseAdam. глубокое обучение. оптимизация памяти.

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

Я разработал маленький пакет Disk Sparse Adam (DSA) — 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 падает с классической ошибкой:

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

Конечно скорость обучения не сравнится с обучением на 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

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

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

👉 Kagge

Задача: прогнать обучение на 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

Источник