Если вы обучаете графные нейросети или 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, фидбеку в Issues и пулл-реквестам! Если у вас есть задачи с большими графами или эмбеддингами — попробуйте DSA и делитесь результатами в комментариях.
🤗 Интерактивный калькулятор памяти на Hugging Face:
🧪 Бенчмарк: Запуск на 1 000 000 сущностей в Kaggle Notebook
Задача: прогнать обучение на 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
Спасибо что дочитал)

