Если вы обучаете графные нейросети или 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 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, накладные расходы на диск минимальны и полностью перекрываются экономией памяти.
Где это пригодится?
-
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 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/