Настройка FSDP (Fully Sharded Data Parallel) для обучения

Модель LLaMA-2 70B не помещается в память A100 80GB при использовании DDP. FSDP решает эту проблему, шардируя параметры, градиенты и оптимизатор между GPU. Мы настраиваем FSDP под ключ — нативную реализацию fully sharded data parallelism в PyTorch, которая экономит до 70% памяти без потери скорости.

Направления AI-разработки

Часто задаваемые вопросы

Последние работы

  • image_website-b2b-advance_0.webp
    Разработка сайта компании B2B ADVANCE
    1441
  • image_web-applications_feedme_466_0.webp
    Разработка веб-приложения для компании FEEDME
    1302
  • image_websites_belfingroup_462_0.webp
    Разработка веб-сайта для компании БЕЛФИНГРУПП
    998
  • image_ecommerce_furnoro_435_0.webp
    Разработка интернет магазина для компании FURNORO
    1267
  • image_logo-advance_0.webp
    Разработка логотипа компании B2B Advance
    714
  • image_crm_enviok_479_0.webp
    Разработка веб-приложения для компании Enviok
    1006

Модель LLaMA-2 70B не помещается в память A100 80GB при использовании DDP. FSDP решает эту проблему, шардируя параметры, градиенты и оптимизатор между GPU. Мы настраиваем FSDP под ключ — нативную реализацию fully sharded data parallelism в PyTorch, которая экономит до 70% памяти без потери скорости. Сертифицированные инженеры с многолетним опытом в distributed training. За время работы мы выполнили более 50 проектов для моделей от 1B до 70B параметров. Наши клиенты экономят до 40% бюджета на облачные GPU за счет оптимальной конфигурации.

PyTorch FSDP documentation

Почему FSDP выгоднее DeepSpeed?

FSDP — часть PyTorch core и не требует внешних зависимостей. В отличие от DeepSpeed ZeRO-3, интеграция с Hugging Face Transformers и Accelerate происходит через нативные API. Мы используем FSDP в каждом втором проекте по fine-tuning больших моделей — от LLaMA до Mistral. PyTorch FSDP documentation

Как работает FSDP?

Принцип работы

При forward pass: параметры каждого sharded layer собираются (all-gather) со всех GPU перед вычислением. После forward — немедленно освобождаются, если включён reshard_after_forward. При backward pass: параметры снова собираются, градиенты вычисляются, затем reduce-scatter распределяет шарды градиентов по GPU. Это устраняет ситуацию, когда каждый GPU хранит полную копию модели, как в обычном DDP.

Базовая настройка

import torch import torch.distributed as dist from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp.fully_sharded_data_parallel import ( CPUOffload, BackwardPrefetch, ) from torch.distributed.fsdp.wrap import ( size_based_auto_wrap_policy, enable_wrap, wrap, ) import functools def setup_fsdp(rank, world_size): dist.init_process_group("nccl", rank=rank, world_size=world_size) torch.cuda.set_device(rank) def wrap_model_with_fsdp(model, rank): auto_wrap_policy = functools.partial( size_based_auto_wrap_policy, min_num_params=100_000_000 ) model = FSDP( model, auto_wrap_policy=auto_wrap_policy, cpu_offload=CPUOffload(offload_params=False), backward_prefetch=BackwardPrefetch.BACKWARD_PRE, device_id=torch.cuda.current_device(), sharding_strategy=ShardingStrategy.FULL_SHARD, mixed_precision=MixedPrecision( param_dtype=torch.bfloat16, reduce_dtype=torch.float32, buffer_dtype=torch.bfloat16, ), ) return model 

Как выбрать стратегию шардирования?

from torch.distributed.fsdp import ShardingStrategy # FULL_SHARD — полное шардирование (аналог ZeRO-3) strategy = ShardingStrategy.FULL_SHARD # SHARD_GRAD_OP — шардирование только градиентов и оптимизатора (ZeRO-2) strategy = ShardingStrategy.SHARD_GRAD_OP # NO_SHARD — обычный DDP strategy = ShardingStrategy.NO_SHARD # HYBRID_SHARD — FULL_SHARD внутри узла, репликация между узлами strategy = ShardingStrategy.HYBRID_SHARD 

Выбор стратегии зависит от размера модели, количества GPU и скорости межсоединений. Для 8 GPU с NVLink оптимален FULL_SHARD, для multi-node — HYBRID_SHARD.

Стратегии шардирования: сравнение памяти и скорости

Стратегия Экономия памяти Коммуникационный overhead Типичный сценарий
FULL_SHARD До 75% Высокий Одна нода с быстрым межсоединением
SHARD_GRAD_OP До 50% Средний Модели среднего размера
HYBRID_SHARD ~60% Низкий Multi-node кластеры
NO_SHARD 0% Низкий Базовая DDP

Как настроить FSDP: пошаговая инструкция

  1. Определите топологию кластера: количество GPU, узлов, тип межсоединения (NVLink, InfiniBand).
  2. Выберите стратегию шардирования: FULL_SHARD для одного узла с NVLink, HYBRID_SHARD для multi-node.
  3. Настройте mixed precision: используйте bfloat16 для параметров, float32 для reductions.
  4. Переопределите wrap policy: для трансформеров используйте transformer_auto_wrap_policy с указанием класса слоя.
  5. Оптимизируйте checkpointing: включите offload_to_cpu при сохранении full state dict.
  6. Профилируйте производительность: измерьте throughput, GPU utilization и latency p99.
Когда использовать HYBRID_SHARD? HYBRID_SHARD сочетает FULL_SHARD внутри ноды и репликацию между нодами. Это снижает межнодовый трафик, что критично при медленных межсоединениях (Ethernet). Подходит для кластеров из 2+ узлов с InfiniBand или RoCE.

Практические аспекты настройки

Wrap policy для трансформеров

Для трансформеров важно оборачивать каждый Transformer block отдельно:

from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy from transformers.models.llama.modeling_llama import LlamaDecoderLayer llama_auto_wrap_policy = functools.partial( transformer_auto_wrap_policy, transformer_layer_cls={LlamaDecoderLayer}, ) model = FSDP(model, auto_wrap_policy=llama_auto_wrap_policy) 

Сохранение и загрузка checkpoint

from torch.distributed.fsdp import FullStateDictConfig, StateDictType save_policy = FullStateDictConfig(offload_to_cpu=True, rank0_only=True) with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, save_policy): cpu_state = model.state_dict() if rank == 0: torch.save(cpu_state, "checkpoint.pt") if rank == 0: state_dict = torch.load("checkpoint.pt") else: state_dict = {} with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, save_policy): model.load_state_dict(state_dict) 

Интеграция с Hugging Face Accelerate

from accelerate import Accelerator from accelerate.utils import FullyShardedDataParallelPlugin from torch.distributed.fsdp.fully_sharded_data_parallel import FullOptimStateDictConfig, FullStateDictConfig fsdp_plugin = FullyShardedDataParallelPlugin( state_dict_config=FullStateDictConfig(offload_to_cpu=True, rank0_only=False), optim_state_dict_config=FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=False), ) accelerator = Accelerator(fsdp_plugin=fsdp_plugin) 

Как мы настроили FSDP для LLaMA-70B

В одном из проектов нам потребовалось дообучить LLaMA-2 70B на 8x A100 80GB. Исходно модель не влезала даже с DeepSpeed ZeRO-3. Мы выбрали FSDP с FULL_SHARD и гибридной точностью bfloat16, настроили transformer_auto_wrap_policy и backward prefetch. В результате throughput составил 850 tokens/s при batch size 4 на GPU. Экономия памяти — 68% по сравнению с DDP. Кроме того, мы сократили время на каждый epoch на 30% за счёт оптимизации коммуникации. Клиент сэкономил более 35% затрат на аренду GPU.

Что входит в настройку FSDP

  • Аудит модели и конфигурации GPU
  • Выбор оптимальной стратегии шардирования и mixed precision
  • Настройка wrap policy под архитектуру (трансформеры, CNN, GNN)
  • Интеграция с Accelerate и Hugging Face Trainer
  • Оптимизация checkpointing и загрузки
  • Профилирование производительности (throughput, memory, GPU utilization)
  • Документация и обучение вашей команды
  • Поддержка после деплоя

Типичные ошибки при настройке FSDP

  • OOM при сохранении checkpoint: используйте FullStateDictConfig с offload_to_cpu=True.
  • Медленная инициализация: попробуйте HYBRID_SHARD для multi-node.
  • Несовместимость с некоторыми слоями: проверьте auto_wrap_policy на все подмодули.

Сроки настройки — от 5 до 10 рабочих дней. Стоимость рассчитывается индивидуально после бесплатной консультации. Свяжитесь с нами, чтобы обсудить задачу. Закажите настройку FSDP под ключ — получите консультацию сертифицированного инженера.