Если вы обучаете графные нейросети или Knowledge Graph Embeddings на миллионы узлов, вы наверняка сталкивались с тем, что стандартный torch.optim.SparseAdam моментально забивает всю оперативную память или видеопамять.
Я разработал маленький пакет Disk Sparse Adam (DSA) — Out-of-Core оптимизатор для PyTorch, который выносит состояния моментов ( и
) на диск через
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, получаем проблемы:
-
Первый момент (
): еще 5.12 ГБ
-
Второй момент (
): еще 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 состояния и
для всех 10 млн узлов одновременно, если прямо сейчас нам нужны состояния только для 10 000?
Так появился Disk Sparse Adam (DSA).
┌──────────────────────────────────────────────────────────┐
│ Память при использовании DSA │
├──────────────────────────────────────────────────────────┤
│ VRAM / RAM: [Параметры + Активный батч (пара МБ)] │
│ DISK (mmap): [История m и v лежит на NVMe SSD] │
└──────────────────────────────────────────────────────────┘
DSA выносит матрицы моментов и
на диск в виде бинарных файлов и отображает их в память через механизм OS
mmap (memory mapping):
-
На шаге
optimizer.step()DSA считывает с диска состояния моментов только для активных индексов текущего батча. -
Проводит обновления по формуле Adam.
-
Записывает обновленные состояния обратно на диск.
-
Расход оперативной/видеопамяти под состояния оптимизатора становится практически нулевым.
Как это выглядит в коде
Одна из главных задач при разработке 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, накладные расходы на диск минимальны и полностью перекрываются экономией памяти.
Где это пригодится?
-
GNN и графные нейросети (PyTorch Geometric / DGL): Обучение эмбеддингов узлов в графах на десятки миллионов вершин (
Node2Vec,HeteroDataи т.д.). -
Knowledge Graph Embeddings : Обучение в неевклидовых геометриях, Complex на больших графах знаний.
-
Рекомендательные системы (RecSys): Огромные таблицы пользователей и товаров (Lookup Tables).
-
Исследователи с ограниченным бюджетом: Возможность запускать эксперименты на одной видеокарте или в бесплатном 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). Суммарный объем весов и состояний и
на диске — ~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


