Распределённое обучение: от одной видеокарты до кластера
В статье про обучение сетей мы разобрали, как заставить модель сойтись, молчаливо предполагая, что всё происходит на одном ускорителе. В статье про инференс — как дёшево эксплуатировать готовый чекпоинт. Между ними зияет дыра, в которую проваливается любой, кто впервые берётся за модель крупнее пары сотен миллионов параметров: обучение больше не помещается в один GPU, и вся инженерия меняется.
Меняется не «немного». Появляется сеть между устройствами, а значит — латентность, пропускная способность и отказы, то есть вся физика распределённых систем. Появляется вопрос, какую именно ось разрезать: данные, слои, тензоры внутри слоя или контекст. Появляется экономика: кластер из 64 GPU стоит примерно как шестьдесят четыре GPU, а ускорение вы получите в 40 раз — и надо понимать, куда делись оставшиеся 24.
Эта статья — про инженерию масштаба. Она не про то, как придумать архитектуру, а про то, как физически разложить обучение по железу так, чтобы оно шло быстро, не падало каждые шесть часов и не сжигало бюджет впустую.
1. Три разные причины выйти за пределы одной карты
Прежде чем выбирать технику, надо честно назвать проблему. Их ровно три, и они требуют разных решений.
| Проблема | Симптом | Что реально помогает |
|---|---|---|
| Не влезает состояние обучения | CUDA out of memory уже при batch = 1 |
шардинг состояния (ZeRO/FSDP), тензорный и конвейерный параллелизм, checkpointing активаций |
| Слишком долго | эпоха идёт 40 часов, эксперимент — неделя | параллелизм по данным (DDP), больший эффективный батч |
| Не влезают данные | датасет в десятки терабайт, диск узкое место | шардированный стриминг датасета, префетч, отдельная тема |
Ошибка новичка — тащить FSDP туда, где хватило бы DDP и градиентного накопления. Ошибка ветерана — верить, что «сейчас докупим карт», когда модель просто не влезает и никакой DDP этого не изменит: DDP реплицирует модель, он не уменьшает её.
1.1 Бюджет памяти: считаем честно
Возьмём стандартный рецепт mixed precision из статьи про обучение: веса и градиенты в bf16, мастер-копия весов и моменты Adam в fp32. На один параметр:
| Что храним | Формат | Байт на параметр |
|---|---|---|
| Веса для вычислений | bf16 | 2 |
| Градиенты | bf16 | 2 |
| Мастер-копия весов | fp32 | 4 |
| Момент $m$ (Adam) | fp32 | 4 |
| Момент $v$ (Adam) | fp32 | 4 |
| Итого | 16 |
Для модели на 7 млрд параметров это 112 ГБ до единой активации. В H100 с 80 ГБ HBM такое не влезает физически — и никакая оптимизация батча не спасёт. Отдельно считаются активации: для трансформера порядок величины — $L \cdot b \cdot s \cdot h \cdot c$ байт, где $L$ — число слоёв, $b$ — размер батча, $s$ — длина последовательности, $h$ — скрытая размерность, $c$ — константа порядка 10–20, зависящая от того, что именно кэшируется для backward.
def training_memory_gb(params_b: float, layers: int, hidden: int,
batch: int, seq: int,
bytes_per_param: int = 16,
act_const: int = 14,
recompute: bool = False) -> dict:
"""Прикидка памяти на одну карту при обучении. params_b — миллиарды параметров."""
p = params_b * 1e9
state = p * bytes_per_param / 1e9 # веса, градиенты, состояния Adam
act = layers * batch * seq * hidden * act_const / 1e9 # активации для backward
if recompute:
# при полном checkpointing храним только границы слоёв: 2 байта на элемент
act = layers * batch * seq * hidden * 2 / 1e9
return {"состояние, ГБ": round(state, 1),
"активации, ГБ": round(act, 1),
"итого, ГБ": round(state + act, 1)}
print(training_memory_gb(7, layers=32, hidden=4096, batch=4, seq=4096))
# {'состояние, ГБ': 112.0, 'активации, ГБ': 30.1, 'итого, ГБ': 142.1}
print(training_memory_gb(7, layers=32, hidden=4096, batch=4, seq=4096, recompute=True))
# {'состояние, ГБ': 112.0, 'активации, ГБ': 4.3, 'итого, ГБ': 116.3}
Первое, что видно из этой таблички: активации сжимаются легко, состояние — нет. Отсюда и вырос весь шардинг.
2. Физика кластера: иерархия каналов
Всё распределённое обучение — это торговля между вычислениями и передачей данных. Чтобы принимать решения, надо помнить порядки величин (числа для типового узла DGX-класса 2024 года, у вас будут другие, но соотношения устойчивы).
| Канал | Пропускная способность | Во сколько раз медленнее HBM |
|---|---|---|
| HBM внутри GPU | ≈ 3300 ГБ/с | 1 |
| NVLink между GPU в узле | ≈ 900 ГБ/с | 3,7 |
| PCIe 5.0 x16 (GPU ↔ CPU) | ≈ 64 ГБ/с | 50 |
| InfiniBand NDR, 1 порт на GPU | ≈ 50 ГБ/с | 66 |
| Обычный Ethernet 25 Гбит/с | ≈ 3 ГБ/с | 1000 |
Правило, из которого следует вся раскладка: чем чаще операция требует синхронизации, тем толще должен быть канал под ней. Тензорный параллелизм синхронизируется несколько раз на каждый слой — ему нужен NVLink. Параллелизм по данным синхронизируется раз в шаг оптимизатора — ему хватит InfiniBand. Если перепутать, эффективность падает в разы, причём профайлер покажет «GPU простаивает», а не «сеть плохая».
Подробнее про то, что физически происходит в ускорителе и почему у него такая иерархия памяти, — в треке про железо: Ускорители и Подсистема памяти.
2.1 MFU — единственная честная метрика утилизации
Model FLOPs Utilization — доля пиковой производительности железа, которую вы реально используете:
$$\text{MFU} = \frac{6 N D / T}{P_{\text{peak}} \cdot G}$$
где $N$ — число параметров, $D$ — число обработанных токенов, $T$ — время в секундах, $G$ — число GPU, $P_{\text{peak}}$ — пиковые FLOPs одной карты. Коэффициент 6 — это классическая оценка «2 FLOP на параметр в forward, 4 в backward».
Ориентиры: 15–25 % MFU — типично для наивной реализации, 35–50 % — хорошо настроенное обучение большой модели (в отчёте по PaLM заявлено 46,2 %), выше 60 % бывает редко и обычно означает ошибку в подсчёте. Смысл метрики в том, что она не даёт себя обмануть: увеличили кластер вдвое, шаг ускорился в 1,3 раза — MFU честно упадёт, и станет видно, что вы платите за воздух.
def mfu(params_b: float, tokens_per_step: int, step_time_s: float,
n_gpu: int, peak_tflops: float = 989.0) -> float:
"""Model FLOPs Utilization в процентах. peak_tflops — пик BF16 одной карты."""
flops = 6 * params_b * 1e9 * tokens_per_step
achieved = flops / step_time_s
return 100 * achieved / (peak_tflops * 1e12 * n_gpu)
# 7B, глобальный батч 2 млн токенов, шаг 8,5 с, 64 карты H100
print(round(mfu(7, 2_000_000, 8.5, 64), 1)) # 15.7 — есть куда расти
3. Коллективные операции: словарь и цена
Всё общение при обучении сводится к нескольким коллективам, реализованным в NCCL. Понимать их надо не «по названию», а по объёму трафика.
| Операция | Что делает | Объём на ранг (буфер $S$ байт) |
|---|---|---|
all-reduce |
суммирует буферы всех рангов, результат у всех | $2S(N-1)/N \approx 2S$ |
reduce-scatter |
суммирует, но каждому — свой кусок | $S(N-1)/N \approx S$ |
all-gather |
собирает куски со всех, у всех целое | $S(N-1)/N \approx S$ |
broadcast |
один рассылает всем | $\approx S$ |
all-to-all |
каждый шлёт каждому свой кусок | $S(N-1)/N$, но $N^2$ соединений |
Ключевой факт: all-reduce = reduce-scatter + all-gather, и потому стоит ровно вдвое
дороже каждой из половинок. Отсюда почти вся арифметика дальше.
Кольцевая реализация all-reduce (классическое объяснение Baidu) устроена так: буфер режется на $N$ частей, дальше $N-1$ шагов «редукции по кольцу» и $N-1$ шагов «раздачи по кольцу». Каждый ранг на каждом шаге отправляет ровно $S/N$ байт.
Практический вывод, который экономит недели: время кольцевого all-reduce почти не зависит от числа участников, оно определяется пропускной способностью самого узкого канала. Поэтому DDP на 8 картах и на 512 картах масштабируется похоже — пока сеть однородна. Ломается всё, когда в кольце появляется медленное звено (одна карта на PCIe вместо NVLink, один узел на другом коммутаторе).
Проверять железо надо до первого обучения, утилитой nccl-tests:
# ожидаемая busbw на NVLink-узле — сотни ГБ/с; если видите десятки, ищите проблему
./build/all_reduce_perf -b 8M -e 2G -f 2 -g 8
Читать надо колонку busbw (bus bandwidth), а не algbw: она уже учитывает множитель
$2(N-1)/N$ и потому сравнима с паспортной скоростью канала.
4. Параллелизм по данным: DDP
Самая простая и самая полезная техника. Каждый ранг держит полную копию модели, получает свой кусок батча, считает свои градиенты, после чего градиенты усредняются через all-reduce — и все делают идентичный шаг оптимизатора.
Ключевая деталь реализации, описанная в статье про PyTorch DDP: градиенты собираются в бакеты и all-reduce для бакета запускается сразу, как только он заполнился, не дожидаясь конца backward. Коммуникация перекрывается вычислением, и на длинных сетях сетевая задержка почти исчезает из критического пути.
import os
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, DistributedSampler
def main():
# torchrun сам выставит RANK, LOCAL_RANK, WORLD_SIZE
dist.init_process_group(backend="nccl")
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
model = build_model().cuda(local_rank)
# bucket_cap_mb — размер бакета; крупнее = меньше вызовов, но хуже перекрытие
model = DDP(model, device_ids=[local_rank], bucket_cap_mb=25,
gradient_as_bucket_view=True)
sampler = DistributedSampler(train_ds, shuffle=True, drop_last=True)
loader = DataLoader(train_ds, batch_size=8, sampler=sampler,
num_workers=8, pin_memory=True, persistent_workers=True)
opt = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.1)
scaler_dtype = torch.bfloat16
for epoch in range(epochs):
sampler.set_epoch(epoch) # без этого перемешивание одинаково каждую эпоху
for step, (x, y) in enumerate(loader):
x, y = x.cuda(non_blocking=True), y.cuda(non_blocking=True)
with torch.autocast("cuda", dtype=scaler_dtype):
loss = model(x, y)
loss.backward() # здесь же идёт all-reduce по бакетам
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
opt.zero_grad(set_to_none=True)
dist.destroy_process_group()
if __name__ == "__main__":
main()
Запуск — через torchrun, а не самописные скрипты:
# один узел, 8 карт
torchrun --standalone --nproc_per_node=8 train.py
# четыре узла по 8 карт, с эластичностью: запуск продолжится при 3 живых узлах
torchrun --nnodes=3:4 --nproc_per_node=8 \
--rdzv-backend=c10d --rdzv-endpoint=head-node:29500 train.py
4.1 Градиентное накопление и эффективный батч
Если нужный глобальный батч не влезает в память, шаг оптимизатора делается раз в $k$ микробатчей:
$$B_{\text{эфф}} = b_{\text{микро}} \cdot k \cdot N_{\text{рангов}}$$
В DDP при этом обязательно оборачивать промежуточные backward в model.no_sync(),
иначе вы заплатите за all-reduce $k$ раз вместо одного:
for i, (x, y) in enumerate(loader):
is_last = (i + 1) % accum_steps == 0
ctx = nullcontext() if is_last else model.no_sync()
with ctx:
with torch.autocast("cuda", dtype=torch.bfloat16):
loss = model(x, y) / accum_steps # делим, иначе градиент в accum раз больше
loss.backward()
if is_last:
opt.step()
opt.zero_grad(set_to_none=True)
4.2 Большой батч ломает подобранный learning rate
Увеличив кластер, вы увеличили глобальный батч — и старый LR стал неправильным. Работающая эвристика из «Accurate, Large Minibatch SGD» (Goyal et al., 2017): при росте батча в $k$ раз увеличивайте LR в $k$ раз и добавьте разогрев в несколько сотен шагов, иначе первые шаги разнесут модель. Для Adam-подобных оптимизаторов чаще работает корневое масштабирование ($\sqrt{k}$), а на очень больших батчах — послойная адаптация в духе LAMB.
И помните про предел: за некоторым «критическим размером батча» дополнительные примеры почти не улучшают оценку градиента, а только сжигают вычисления — это измеримая величина, описанная в An Empirical Model of Large-Batch Training. Практически: если удвоение батча не уменьшило число шагов до целевой потери примерно вдвое, вы уже за пределом и просто платите деньги.
5. Шардинг состояния: ZeRO и FSDP
DDP хранит по полной копии всего на каждой карте — а это ровно те 16 байт на параметр, которые не влезают. Идея ZeRO (Rajbhandari et al., 2019): состояние не нужно дублировать, его достаточно собирать по требованию.
| Стадия | Что шардируется | Байт на параметр на карту при $N$ рангах | Трафик за шаг |
|---|---|---|---|
| DDP | ничего | 16 | $2\Psi$ |
| ZeRO-1 | состояния оптимизатора | $4 + 12/N$ | $2\Psi$ |
| ZeRO-2 | + градиенты | $2 + 14/N$ | $2\Psi$ |
| ZeRO-3 / FSDP | + сами параметры | $16/N$ | $3\Psi$ |
Здесь $\Psi$ — размер модели в байтах. Читать таблицу так: при 64 картах ZeRO-3 сводит 112 ГБ состояния 7B-модели к 1,75 ГБ на карту — ценой полуторакратного роста трафика. Механика ZeRO-3 (в PyTorch — FSDP):
перед его forward СобранСлой --> Шардировано: сразу после forward
освобождаем полные веса Шардировано --> СобранСлойBwd: all-gather тех же параметров
перед backward слоя СобранСлойBwd --> ГрадиентыСведены: reduce-scatter градиентов
каждому рангу свой шард ГрадиентыСведены --> Шардировано: локальный шаг оптимизатора
над своим шардом note right of Шардировано В покое карта держит 1/N весов, 1/N градиентов и 1/N состояний Adam end note
import functools
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP, ShardingStrategy
from torch.distributed.fsdp import MixedPrecision, CPUOffload
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
checkpoint_wrapper, apply_activation_checkpointing)
# оборачиваем поблочно: единица шардинга — один блок трансформера
wrap_policy = functools.partial(
transformer_auto_wrap_policy, transformer_layer_cls={TransformerBlock})
model = FSDP(
model,
auto_wrap_policy=wrap_policy,
sharding_strategy=ShardingStrategy.FULL_SHARD, # = ZeRO-3
mixed_precision=MixedPrecision(param_dtype=torch.bfloat16,
reduce_dtype=torch.float32, # редукция в fp32 — важно
buffer_dtype=torch.bfloat16),
cpu_offload=CPUOffload(offload_params=False), # включать только когда иначе OOM
device_id=torch.cuda.current_device(),
limit_all_gathers=True, # не даёт префетчу съесть память
)
# checkpointing активаций поверх FSDP: ~30 % лишних FLOPs за кратное падение памяти
apply_activation_checkpointing(
model,
checkpoint_wrapper_fn=checkpoint_wrapper,
check_fn=lambda m: isinstance(m, TransformerBlock))
Три практические заметки, каждая стоит потерянного дня:
HYBRID_SHARDчасто быстрееFULL_SHARDна нескольких узлах: шардим внутри узла по NVLink, реплицируем между узлами по InfiniBand. Трафик по медленному каналу сразу падает.reduce_dtype=float32. Редукция градиентов в bf16 на тысячах рангов копит ошибку округления и даёт тихую деградацию — потерь не видно, качество хуже.- Гранулярность обёртки решает всё. Обернуть модель целиком одним FSDP-юнитом =
собрать все веса разом = тот же OOM. Оборачивать каждый
nn.Linear= миллион мелких all-gather и мёртвая сеть. Единица — блок трансформера.
Если и этого мало — ZeRO-Offload и ZeRO-Infinity сгружают состояния в RAM и на NVMe. Это работает, но PCIe в 50 раз медленнее HBM: техника «дообучить 13B на одной карте за ночь», а не «обучать быстро».
6. Тензорный параллелизм: режем матрицу
Когда даже один слой не влезает в карту, режут сам слой. Подход Megatron-LM (Shoeybi et al., 2019) устроен геометрически красиво: для двух подряд идущих матриц MLP первую режут по столбцам, вторую по строкам — тогда между ними синхронизация не нужна вообще.
$$Y = \text{GELU}(X A), \quad Z = Y B$$
Разрежем $A = \lbrack A_1, A_2 \rbrack$ по столбцам: каждый ранг считает
$Y_i = \text{GELU}(X A_i)$ независимо (нелинейность поэлементная — это и есть причина
резать именно так). Разрежем $B$ по строкам: ранг считает $Z_i = Y_i B_i$, и результат
получается одним all-reduce: $Z = \sum_i Z_i$. Итого один коллектив на MLP-блок
в forward и один в backward. Внимание режется ещё естественнее — по головам,
каждая голова целиком живёт на своём ранге.
Цена: объём одного all-reduce равен $b \cdot s \cdot h \cdot 2$ байт, и таких операций 4 на слой (две в forward, две в backward). Для $b=4$, $s=4096$, $h=4096$ это 128 МБ, умноженные на 32 слоя — 4 ГБ трафика на каждый микробатч. По NVLink это единицы миллисекунд, по InfiniBand — уже десятки, по Ethernet — катастрофа. Отсюда железное правило: TP не выходит за пределы узла, степень TP ≤ числу карт с NVLink.
Родственная техника — sequence parallelism из Reducing Activation Recomputation: участки, где TP не помогает (LayerNorm, dropout), режутся по оси последовательности, что убирает дублирование активаций почти бесплатно. При длинном контексте её продолжение — Ring Attention, где по рангам режется сам контекст, а блоки внимания передаются по кольцу.
7. Конвейерный параллелизм и пузырь
Третья ось — разложить разные слои на разные устройства. Трафик минимален: между стадиями летят только активации на границе. Но появляется структурная беда: пока первая стадия считает, остальные ждут.
Лечение — микробатчи (GPipe, Huang et al., 2018). Глобальный батч режется на $m$ микробатчей, они текут по конвейеру внахлёст. Доля простоя («пузырь»):
$$\text{bubble} = \frac{P - 1}{m + P - 1}$$
где $P$ — число стадий. При $P = 4$ и $m = 4$ пузырь — 43 % (катастрофа), при $m = 32$ — 8,6 % (терпимо). Правило: $m \geq 4P$.
Расписание 1F1B (one-forward-one-backward) из Efficient Large-Scale Training on GPU Clusters (Narayanan et al., 2021) не уменьшает пузырь, но радикально уменьшает пиковую память: backward микробатча запускается сразу, как только он прошёл вперёд, и его активации освобождаются, вместо того чтобы копиться до конца. Чередующееся (interleaved) расписание, где каждая карта держит несколько несмежных участков модели, режет пузырь ещё в $v$ раз ценой дополнительного трафика.
Отдельная головная боль конвейера — балансировка стадий. Медленная стадия определяет темп всего конвейера; embedding и выходная проекция обычно тяжелее среднего блока, и наивное деление «поровну по слоям» даёт перекос в 20–30 %.
8. Как всё это комбинируют
Три оси ортогональны, их перемножают: $G = TP \times PP \times DP$. Плюс шардинг ZeRO
поверх DP, плюс экспертный параллелизм для MoE (эксперты раскладываются по картам,
маршрутизация превращается в all-to-all; см. Switch Transformer).
с batch >= 1?} B -->|Да| C{Достаточно быстро?} C -->|Да| D[Одна карта.
Не усложняйте] C -->|Нет| E[DDP + градиентное накопление
масштабируйте LR] B -->|Нет| F[Включите checkpointing активаций
и bf16] F --> G{Теперь влезает?} G -->|Да| E G -->|Нет| H[FSDP / ZeRO-3
шардим состояние] H --> I{Влезает?} I -->|Да| J[HYBRID_SHARD внутри узла,
реплика между узлами] I -->|Нет| K{Один слой больше карты?} K -->|Да| L[Тензорный параллелизм
ТОЛЬКО внутри узла по NVLink] K -->|Нет| M[Конвейерный параллелизм
между узлами, m >= 4P] L --> N[3D: TP внутри узла,
PP между узлами, DP сверху] M --> N J --> O[Измерьте MFU] E --> O N --> O O --> P{MFU выше 35 %?} P -->|Нет| Q[Профилируйте: dataloader,
дисбаланс, коллективы] P -->|Да| R[Запускайте длинный прогон]
Порядок выбора почти всегда такой:
- Уменьшить требования — checkpointing активаций, bf16, 8-битный оптимизатор (Dettmers et al., 2021, экономит 8 из 16 байт на параметр), flash-attention.
- Шардировать состояние — FSDP/ZeRO. Прозрачно для кода модели.
- Резать слой — TP, только внутри узла.
- Резать по слоям — PP, только когда TP исчерпан.
Обратный порядок (начать с 3D-параллелизма) — самый популярный способ потратить месяц на инфраструктуру ради модели, которую можно было обучить на четырёх картах.
9. Отказоустойчивость: обучение длиной в недели
Кластер из 512 GPU падает регулярно — это не аномалия, а статистика: если MTBF одного узла 30 дней, то у 64 узлов авария случается примерно раз в 11 часов. Обучение, которое не переживает падение, при таком масштабе не завершится никогда. Здесь трек смыкается с моделями отказов из распределённых систем — только вместо запросов у вас чекпоинты.
NCCL-коммуникатор создан Обучение --> Чекпоинт: каждые N шагов Чекпоинт --> Обучение: асинхронная запись,
обучение не ждёт диск Обучение --> ПадениеУзла: NCCL timeout / ECC / OOM ПадениеУзла --> Восстановление: планировщик даёт
новый узел Восстановление --> Инициализация: грузим последний
консистентный чекпоинт Обучение --> СпайкПотерь: loss взлетел СпайкПотерь --> Откат: rollback на чекпоинт до спайка,
пропустить проблемные данные Откат --> Обучение Обучение --> [*]: бюджет токенов израсходован note right of Чекпоинт Шардированный чекпоинт: каждый ранг пишет свой кусок. Иначе rank 0 становится узким местом на десятки минут end note
Что должно быть в чекпоинте, чтобы прогон продолжился бит в бит: веса, состояния оптимизатора, шаг планировщика LR, состояние RNG на каждом ранге и позиция в датасете. Пропустите последнее — после рестарта модель второй раз увидит те же данные, и это самая незаметная форма порчи прогона.
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict
def save(model, opt, step, data_pos, path):
"""Шардированный чекпоинт: каждый ранг пишет свою часть параллельно."""
msd, osd = get_state_dict(model, opt)
state = {"model": msd, "optim": osd, "step": step,
"data_position": data_pos,
"rng": torch.cuda.get_rng_state_all()}
dcp.save(state, checkpoint_id=f"{path}/step-{step}")
def load(model, opt, path, step):
state = {"model": {}, "optim": {}}
dcp.load(state, checkpoint_id=f"{path}/step-{step}")
set_state_dict(model, opt, model_state_dict=state["model"],
optim_state_dict=state["optim"])
return state
Частота чекпоинтов выбирается арифметикой, а не интуицией. Если авария происходит в среднем раз в $T_f$ часов, а чекпоинт стоит $c$ минут, то оптимальный интервал примерно $\sqrt{2 c T_f}$: при $T_f = 11$ ч и $c = 2$ мин — раз в ~50 минут.
Отдельная категория проблем — отстающие (stragglers). Одна карта с деградировавшим охлаждением тормозит весь кластер, потому что коллектив ждёт самого медленного. Симптом: MFU упал на 20 %, все GPU показывают высокую загрузку (они действительно заняты — ожиданием в NCCL). Лечится замером времени шага по рангам:
# раз в 100 шагов собираем длительность шага со всех рангов и ищем выброс
t = torch.tensor([step_time], device="cuda")
gathered = [torch.zeros_like(t) for _ in range(dist.get_world_size())]
dist.all_gather(gathered, t)
if dist.get_rank() == 0:
times = torch.cat(gathered)
slow = (times > times.median() * 1.15).nonzero().flatten().tolist()
if slow:
print(f"отстающие ранги: {slow}") # кандидаты на вывод из кластера
И третье — спайки потерь. На больших моделях loss иногда взлетает без видимой причины; стандартная практика, задокументированная в логбуке обучения OPT-175B — откатиться на чекпоинт до спайка, пропустить несколько сотен батчей и продолжить. Читать этот логбук стоит целиком: это лучший существующий документ о том, как на самом деле выглядит обучение большой модели.
10. Диагностика: симптом → причина
| Симптом | Вероятная причина | Что делать |
|---|---|---|
| MFU < 15 %, GPU периодически в нуле | dataloader не успевает | больше num_workers, префетч, предподготовленные шарды |
| Масштабирование хуже линейного с ростом узлов | all-reduce не перекрывается вычислением | крупнее бакеты, HYBRID_SHARD, проверить топологию NCCL |
| Обучение зависает на первом шаге | несогласованный порядок коллективов между рангами | ветвления по if rank == 0 вокруг forward — убрать |
NCCL timeout через 10 минут |
один ранг не дошёл до коллектива (упал/ушёл в OOM) | смотреть лог именно этого ранга, а не rank 0 |
| Потери отличаются при том же seed | RNG не синхронизирован, недетерминированные ядра | seed на ранг, torch.use_deterministic_algorithms |
| Память растёт со временем | активации не освобождаются, копится история графа | set_to_none=True, не хранить loss в списке |
| Загрузка GPU 100 %, но шаг медленный | отстающий ранг; коллектив ждёт | замер времени шага по рангам |
Инструменты: torch.profiler с трассировкой NCCL, nsys profile для полной картины,
NCCL_DEBUG=INFO для проверки, какие каналы реально выбраны. Общая методика измерений —
в Производительность: измерение, она
здесь работает без изменений: сначала измерьте, потом чините.
11. Типичные ошибки
- Начать с 3D-параллелизма. Почти всегда достаточно FSDP; TP и PP — это ответ на конкретное ограничение, а не «более продвинутый режим».
- Забыть
sampler.set_epoch(epoch). Каждая эпоха идёт в одном и том же порядке; качество тихо хуже, воспроизвести баг тяжело. - Не поделить loss на число шагов накопления. Эффективный LR оказывается в $k$ раз больше задуманного.
- Оставить LR от однокарточного прогона. Батч вырос в 32 раза — обучение либо расходится, либо сходится вдвое медленнее нужного.
- Логировать и считать метрики на всех рангах. Логи в 512 экземпляров, а метрики
без
all_reduceпоказывают статистику одного ранга. - Валидация только на rank 0. Остальные ранги ждут в коллективе и ловят timeout.
- Синхронизирующие вызовы в горячем цикле.
loss.item()каждый шаг заставляет CPU ждать GPU и разрушает перекрытие; накапливайте на GPU, синхронизируйтесь раз в N шагов. - Чекпоинт без позиции в датасете. После рестарта модель переучивается на тех же данных.
- BatchNorm в распределённом обучении. Статистики считаются по локальному микробатчу;
нужен
SyncBatchNorm— или, что чаще правильнее, LayerNorm.
12. Как это выглядит в проде
В командах, которые обучают модели регулярно, устройство примерно такое.
Планировщик. Slurm или Kubernetes с gang scheduling: задача либо получает все запрошенные узлы, либо не стартует вовсе. Частичный запуск — гарантированный дедлок в первом же коллективе.
Три уровня прогонов. Отладочный (1 GPU, крошечная модель, ловит ошибки кода за минуты), пилотный (1 узел, 1–2 часа, ловит ошибки конфигурации и даёт MFU), и только потом полный. Переход между уровнями — по чек-листу, а не по интуиции.
Наблюдаемость. В дашборде обязаны быть: loss и grad-norm, время шага по перцентилям, MFU, температуры и throttling карт, время в коллективах, скорость dataloader. Grad-norm особенно ценен — он взлетает раньше, чем расходится loss, и часто даёт пару минут форы. Подход к метрикам — тот же, что в Наблюдаемости распределённых систем.
Экономика. Час 8×H100 стоит ощутимых денег, и главный рычаг — не «ещё карт», а MFU. Поднять утилизацию с 18 % до 40 % — это 2,2 раза экономии, дешевле любой закупки. Spot-инстансы работают только при налаженных чекпоинтах и эластичном запуске.
Когда всего этого не нужно. Если задача — адаптировать готовую модель под домен, почти наверняка хватит LoRA на одной-двух картах: см. Дообучение и раздел про PEFT в статье про LLM. Полное обучение с нуля оправдано в единицах процентов реальных проектов.
Мини-итог
- Сначала назовите проблему: не влезает, долго или много данных — техники для них разные.
- Бюджет памяти обучения — 16 байт на параметр при mixed precision Adam плюс активации. Активации сжимаются checkpointing’ом, состояние — только шардингом.
- Вся раскладка следует из иерархии каналов: HBM ≫ NVLink ≫ PCIe ≈ InfiniBand ≫ Ethernet. Чем чаще синхронизация, тем толще нужен канал.
all-reduceстоит $2S$ и почти не зависит от числа рангов;reduce-scatterиall-gather— по $S$ каждый.- DDP реплицирует, FSDP/ZeRO шардирует ($16/N$ байт на параметр за +50 % трафика), TP режет матрицы внутри узла, PP режет слои между узлами и платит пузырём $(P-1)/(m+P-1)$.
- MFU — единственная метрика, которая не даёт обмануть себя ростом кластера.
- Длинный прогон — это распределённая система: шардированные чекпоинты, обнаружение отстающих, откат на спайках, эластичный рестарт.
Источники
- Rajbhandari et al. ZeRO: Memory Optimizations Toward Training Trillion Parameter Models (2019).
- Zhao et al. PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel (2023).
- Shoeybi et al. Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism (2019).
- Huang et al. GPipe: Efficient Training of Giant Neural Networks Using Pipeline Parallelism (2018).
- Narayanan et al. Efficient Large-Scale Language Model Training on GPU Clusters (2021).
- Chen et al. Training Deep Nets with Sublinear Memory Cost (2016).
- Korthikanti et al. Reducing Activation Recomputation in Large Transformer Models (2022).
- Goyal et al. Accurate, Large Minibatch SGD (2017).
- McCandlish et al. An Empirical Model of Large-Batch Training (2018).
- Li et al. PyTorch Distributed: Experiences on Accelerating Data Parallel Training (2020).
- Документация: PyTorch FSDP, torchrun и elastic, NCCL, DeepSpeed, nccl-tests.
Что дальше
Мы научились обучать модель любого разумного размера. Осталась область, где всё изученное в треке сходится вместе: модели, которые работают не с одним типом данных, а с несколькими сразу — звук, изображение и текст в одной сети.