Главная/Блог/Гайд/Host Offloading в JAX: Как обойти…
Гайд10 мин чтения · 11 июля 2026 г.

Host Offloading в JAX: Как обойти лимиты HBM при обучении LLM

Разбираем, как техника host offloading в JAX позволяет тренировать гигантские модели (DeepSeek-V3, Llama 3.1) на GPU, выходя за пределы физической памяти видеокарты за счет использования быстрой связи NVLink-C2C.

Host Offloading в JAX: Как обойти лимиты HBM при обучении LLM

Обучение современных больших языковых моделей (LLM) — это не просто задача вычислений, это постоянная борьба с ограничениями памяти. По мере того как модели становятся больше, а контекстное окно длиннее, видеопамять (HBM) на графических процессорах превращается в узкое горлышко. Весы модели, градиенты, состояния оптимизатора, буферы коммуникации и, что особенно критично, промежуточные активации (activations) конкурируют за каждый гигабайт доступной памяти. Когда объем этих данных превышает емкость HBM, обучение останавливается с ошибкой Out-of-Memory (OOM), даже если вычислительная мощность GPU простаивает.

Решение, которое долгое время считалось компромиссом, сегодня превращается в мощный инструмент масштабирования. Речь идет о host offloading (выгрузке на хост) — технике, при которой часть промежуточных данных перемещается из видеопамяти GPU в оперативную память сервера (host memory) во время прямого прохода (forward pass) и возвращается обратно для обратного прохода (backward pass). В сочетании с современными архитектурами NVIDIA, такими как Grace Blackwell и Vera Rubin, эта методика перестает быть «костылем» и становится стратегическим преимуществом, позволяющим запускать модели, которые ранее были невозможны для обучения на одном узле.

В этой статье мы подробно разберем, как работает host offloading в фреймворке JAX, почему он так эффективен на новых платформах NVIDIA, и какие конкретные приросты производительности (throughput) и емкости памяти он дает на примерах моделей DeepSeek-V3 и Llama 3.1. Мы также рассмотрим технические нюансы настройки, чтобы вы могли применить эти знания на практике.

01Почему активации — это главная проблема памяти?

Чтобы понять ценность host offloading, нужно сначала осознать масштаб проблемы. При обучении нейронной сети, особенно трансформеров, GPU должен хранить не только веса модели, но и все промежуточные результаты вычислений, необходимые для вычисления градиентов. Эти данные называются активациями.

В отличие от весов, которые можно легко квантовать или сжать, активации часто требуют высокой точности (например, bfloat16 или float32) для стабильности обучения. Кроме того, их объем линейно зависит от размера батча (batch size) и длины последовательности (sequence length). Если вы удвоите размер батча, объем требуемой памяти для активаций также удвоится. Для больших моделей, таких как Llama 3.1 405B или DeepSeek-V3 671B, активации могут занимать до 50-70% всей доступной видеопамяти, оставляя мало места для других необходимых структур данных.

Традиционное решение этой проблемы — activation rematerialization (пересоздание активаций). В этом подходе GPU не хранит активации в памяти, а пересчитывает их заново во время обратного прохода. Это экономит память, но увеличивает вычислительную нагрузку, так как одни и те же вычисления выполняются дважды. Host offloading предлагает альтернативу: вместо пересчета мы сохраняем активации в памяти сервера (host memory) и загружаем их обратно по мере необходимости. Это переносит часть нагрузки с вычислений на передачу данных, что может быть выгоднее, если каналы передачи данных достаточно быстры.

02Архитектурный прорыв: Роль NVLink-C2C

Долгое время host offloading был неэффективен из-за медленных каналов передачи данных между CPU и GPU. Стандартный интерфейс PCIe имеет пропускную способность, которая часто становится узким горлышком. Если GPU ждет данные с CPU, он простаивает, и общая производительность падает.

Однако архитектура NVIDIA Grace Blackwell кардинально меняет эту ситуацию. Процессоры Grace CPU и графические процессоры Blackwell соединены через технологию NVLink-C2C (Chip-to-Chip). Эта связь обеспечивает двунаправленную пропускную способность в 900 ГБ/с. Для сравнения, это на порядки быстрее, чем стандартный PCIe. Будущая платформа NVIDIA Vera Rubin еще больше увеличит эту скорость до 1,8 ТБ/с.

Высокая пропускная способность NVLink-C2C делает оперативную память сервера (host memory) практически продолжением видеопамяти GPU. Это позволяет использовать host memory в качестве эффективного «буфера» для активаций. Ключевой момент здесь — использование pinned host memory (зафиксированной памяти). Зафиксированная память не может быть перемещена в файл подкачки (swap) операционной системой, что позволяет GPU обращаться к ней напрямую через DMA (Direct Memory Access) без участия CPU, минимизируя задержки.

Но одной пропускной способности недостаточно. Чтобы host offloading работал быстро, передача данных должна происходить асинхронно и перекрываться с вычислениями. GPU должен успевать загружать данные с хоста, пока он выполняет другие задачи, и выгружать их, пока вычисляет следующий слой. Именно здесь в игру вступает оптимизация на уровне компилятора XLA.

Политика offloading активаций для повторяющегося слоя MoE в модели DeepSeek-V3 671B.
Политика offloading активаций для повторяющегося слоя MoE в модели DeepSeek-V3 671B.

03Глубокое погружение: Case Study DeepSeek-V3 671B

Давайте рассмотрим конкретный пример из исследований NVIDIA. Модель DeepSeek-V3 671B — это огромная модель с архитектурой Mixture-of-Experts (MoE). Она состоит из 61 декодирующего слоя. Первые три слоя используют плотные многослойные перцептроны (MLP), а остальные 58 слоев используют блоки MoE с многозаголовочным латентным вниманием (Multihead Latent Attention, MLA).

В таких моделях активации от проекций QKV (Query, Key, Value) и от блоков MoE (up/down projections) являются самыми «тяжелыми» по объему. Именно их выгрузка на хост дает наибольший эффект. В исследовании использовалась политика offloading, при которой эти крупные промежуточные активации сохранялись в памяти сервера.

Результаты производительности

Эксперименты проводились на кластере NVIDIA GB200 NVL72 с использованием 128 GPU и фреймворка MaxText (оптимизированная версия JAX для обучения LLM). Результаты были впечатляющими:

  • С базовым пересозданием активаций (rematerialization): Пропускная способность составила 578.3 TFLOPs/s/device.
  • С host offloading без оптимизаций: 541.6 TFLOPs/s/device (медленнее из-за накладных расходов на передачу данных без перекрытия).
  • С host offloading + Latency Hiding Scheduler (LHS) + pipelined transfers: 908.2 TFLOPs/s/device.

Использование оптимизированного host offloading дало прирост пропускной способности на 57% по сравнению с пересозданием активаций. Это означает, что обучение проходит значительно быстрее, несмотря на необходимость передачи данных по шине.

Сравнение пропускной способности (throughput) модели DeepSeek-V3 671B на NVIDIA GB200 в зависимости от политики размещения активаций.
Сравнение пропускной способности (throughput) модели DeepSeek-V3 671B на NVIDIA GB200 в зависимости от политики размещения активаций.

Почему LHS и пайплайнинг критичны?

В случае с DeepSeek-V3 простое включение offloading не дало бы такого эффекта. Ключом стала комбинация двух технологий:

  1. Latency Hiding Scheduler (LHS): Это компонент компилятора XLA, который автоматически планирует операции так, чтобы скрыть задержки передачи данных. Он позволяет GPU выполнять вычисления, пока данные передаются по шине.
  2. Pipelined Host Offloading: Эта функция позволяет конвейеризировать передачи данных. Пока GPU обрабатывает один блок, следующий блок активаций уже загружается с хоста, а предыдущий — выгружается.

Без этих оптимизаций передача данных блокировала бы GPU. С ними задержка передачи полностью «прячется» за полезной работой, что и дает такой огромный прирост скорости.

04Расширение возможностей: Увеличение размера батча

Помимо скорости, host offloading решает проблему емкости памяти. В таблице ниже приведены результаты для DeepSeek-V3 671B с разными конфигурациями.

💡
Важно понимать. Увеличение размера батча (batch size) не только ускоряет обучение за счет лучшей утилизации GPU, но и повышает стабильность градиентов. Однако, как мы видим, без offloading увеличить батч с 256 до 1024 просто невозможно из-за нехватки памяти.
Конфигурация Микро-батч Глобальный батч Throughput (TFLOPs/s) Пик GPU памяти (GiB)
Offloading + LHS + Pipeline 1024 1024 908.2 165.2
Rematerialization + LHS 1024 578.3 578.3 151.3
Offloading (без оптимизаций) 1024 541.6 541.6 145.6
Сохранение на GPU (без offload) 256 256 425.3 113.3
Сохранение на GPU (без offload) 1024 OOM 0.0 0.0

Как видно из таблицы, конфигурация с оптимизированным offloading позволяет использовать глобальный батч 1024, что невозможно при сохранении активаций на GPU (OOM). При этом использование GPU памяти возрастает с 145.6 GiB (базовый offload) до 165.2 GiB. Это кажется парадоксальным: как offloading может увеличивать использование GPU памяти?

Ответ кроется в буферах. Для обеспечения асинхронной передачи (overlap) GPU необходимо держать в памяти буферы для копирования данных и предварительно загруженные активации. Это небольшая плата за возможность работать с гораздо большим батчем и получать в 2 раза большую пропускную способность. Если бы мы не использовали offloading, мы бы вообще не смогли запустить этот батч.

Визуализация масштаба задачи обучения больших языковых моделей (LLM).
Визуализация масштаба задачи обучения больших языковых моделей (LLM).

05Сравнение с плотными моделями: Llama 3.1 405B

Для полноты картины рассмотрим плотную модель Llama 3.1 405B. В отличие от MoE-моделей, у нее нет «экспертов», и активации распределены более равномерно. Здесь offloading применялся выборочно, только для проекций QKV.

Результаты показали прирост пропускной способности с 2669 до 2746 TFLOPs/s/device (около 2.9%). Этот прирост меньше, чем у DeepSeek, но он демонстрирует универсальность метода. Важно отметить, что для Llama 3.1 основным фактором ускорения стал LHS, а не пайплайнинг. Это связано с тем, что в данной конфигурации LHS уже эффективно скрывал большую часть задержек, и добавление пайплайнинга давало diminishing returns.

Также интересно отметить использование памяти: 70.9 GiB хост-памяти было использовано для хранения QKV-активаций. При размере батча 2 и длине последовательности 8192, активации одного слоя занимают около 576 MiB. Поскольку в JAX часто используется опция scan_layers=True, активации обрабатываются послойно, и на GPU одновременно нужны активации только одного слоя. Это делает offloading QKV очень эффективным: мы не храним все активации сразу, а подгружаем их по мере необходимости.

06Как настроить Host Offloading в JAX/MaxText

Если вы хотите попробовать host offloading в своих проектах, вот практическое руководство. Основной инструмент — фреймворк MaxText, который является эталонной реализацией обучения LLM на JAX.

Шаг 1: Подготовка окружения

Используйте официальные контейнеры NVIDIA JAX-Toolbox. Для работы с большими моделями рекомендуется использовать образы, оптимизированные под конкретные архитектуры, например, ghcr.io/nvidia/jax:deepseek_v3_maxtext.

Шаг 2: Настройка флагов XLA

Ключевые флаги компилятора, которые необходимо включить для эффективного offloading:

terminalbash
--xla_gpu_enable_latency_hiding_scheduler=true
--xla_gpu_enable_pipelined_host_offloading=true
--xla_gpu_experimental_parallel_async_compute_limit=8

Последний флаг (parallel_async_compute_limit) определяет, сколько асинхронных операций может находиться в очереди одновременно. Увеличение этого значения дает планировщику LHS больше пространства для маневра, позволяя лучше перекрывать передачи данных с вычислениями и коммуникациями NCCL.

Шаг 3: Выбор политик сохранения (Checkpoint Policies)

В JAX вы можете контролировать, какие именно тензоры сохраняются. Для host offloading обычно используются политики, которые сохраняют крупные активации (например, outputs of attention projections) в памяти хоста. В MaxText это настраивается через параметры конфигурации, указывающие, какие слои или операции подвергать offloading.

⚠️
Важно. Не пытайтесь выгружать все активации. Это создаст избыточный трафик. Выбирайте только самые «тяжелые» тензоры, которые занимают значительную часть памяти и не требуются для каждого шага. Обычно это QKV-проекции в слоях внимания и выходы блоков MLP/MoE.

Шаг 4: Профилирование

Обязательно используйте NVIDIA Nsight Systems для анализа производительности. Вам нужно визуально убедиться, что операции копирования данных (D2H и H2D) перекрываются с операциями вычислений (compute) и коммуникациями (NCCL). Если вы видите «окна» простоя GPU, значит, offloading настроен неэффективно, и данные передаются синхронно, блокируя вычисления.

Архитектура Grace Hopper Superchip, демонстрирующая интеграцию CPU и GPU с высокой пропускной способностью.
Архитектура Grace Hopper Superchip, демонстрирующая интеграцию CPU и GPU с высокой пропускной способностью.

07Что это значит на практике

Host offloading в JAX — это не просто техническая деталь, а стратегический инструмент для инженеров по машинному обучению. Вот ключевые выводы, которые стоит запомнить:

  1. Преодоление ограничений памяти: Если ваша модель не помещается в GPU из-за активаций, host offloading позволяет увеличить размер батча или использовать более длинные контексты, не покупая новое оборудование.
  2. Ускорение обучения: На современных платформах NVIDIA (Grace Blackwell, Vera Rubin) offloading может ускорить обучение на 50% и более по сравнению с пересозданием активаций, за счет эффективного использования пропускной способности NVLink-C2C.
  3. Зависимость от оптимизаций: Простое включение offloading без LHS и пайплайнинга может замедлить обучение. Всегда используйте полный набор оптимизаций XLA.
  4. Профилирование обязательно: Теоретические оценки могут не учитывать накладные расходы на коммуникацию. Используйте Nsight Systems для проверки реального перекрытия операций.

С развитием архитектур, таких как Vera Rubin, где пропускная способность между CPU и GPU удваивается, host offloading станет еще более привлекательным. Это позволяет декуплировать производительность обучения от жестких физических ограничений памяти GPU, открывая путь к обучению моделей следующего поколения, которые ранее были невозможны.

Для тех, кто хочет начать, рекомендуется начать с небольших репрезентативных запусков, используя jax.remat с политикой сохранения на хост и jax.device_put с параметром memory_kind="pinned_host". Постепенно масштабируйте эксперименты, всегда измеряя реальное использование памяти и времени шага.

Host offloading — это мост между теоретической емкостью памяти сервера и практической производительностью GPU. В эпоху гигантских моделей умение эффективно использовать этот мост становится критически важным навыком для любого AI-инженера.

Источник: NVIDIA Developer ↗