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

от автора

Если вы обучаете графные нейросети или 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 torchembedding = 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 torchfrom 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 osimport sysimport gcimport shutilimport timeimport subprocessimport torch# 1. Автоматическая установка из Kaggle Dataset или GitHubtry:    from dsa.optimizer import DiskSparseRiemannianAdamexcept 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 DiskSparseRiemannianAdamdevice = torch.device("cuda" if torch.cuda.is_available() else "cpu")# Путь к виртуальному NVMe диску KaggleKAGGLE_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_000embedding_dim = 128batch_size = 2048initial_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 0optimizer = 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 = 20start_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_timesamples_per_sec = (epochs * batch_size) / total_timeprint(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 optimizergc.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 MBEpoch 05/20 | Loss: 0.000078 | GPU VRAM Overhead: 0.00 MBEpoch 10/20 | Loss: 0.000054 | GPU VRAM Overhead: 0.00 MBEpoch 15/20 | Loss: 0.000039 | GPU VRAM Overhead: 0.00 MBEpoch 20/20 | Loss: 0.000028 | GPU VRAM Overhead: 0.00 MB📊 МЕТРИКИ БЕНЧМАРКА: 🔹 Пропускная способность : 134,212 образцов / сек 🔹 Прирост VRAM на GPU     : 0.00 MB (Состояния оптимизатора вынесены на диск) 🔹 Финальный Loss          : 0.000028

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

ссылка на оригинал статьи https://habr.com/ru/articles/1068202/