РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
742 of 769 menu

Класс FullyShardedDataParallel

Класс FullyShardedDataParallel (FSDP) из модуля torch.distributed реализует гибридный шардинг параметров для обучения больших моделей. В отличие от обычного DistributedDataParallel, FSDP позволяет распределять веса модели между доступными устройствами, что значительно сокращает потребление памяти на каждом ускорителе. При этом параметры считываются с устройств по мере необходимости во время прямого и обратного проходов. Класс принимает модель и набор параметров конфигурации через аргумент sharding_strategy.

Синтаксис

torch.distributed.FullyShardedDataParallel( model, sharding_strategy=ShardingStrategy.FULL_SHARD, cpu_offload=None, auto_wrap_policy=None, backward_prefetch=None, param_init_fn=None, device_id=None )

Пример

Давайте создадим простую модель и обернём её в FSDP с полным шардингом:

import torch import torch.nn as nn import torch.distributed as dist from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp import ShardingStrategy dist.init_process_group("nccl") class SimpleModel(nn.Module): def __init__(self): super().__init__() self.fc1 = nn.Linear(1024, 512) self.fc2 = nn.Linear(512, 128) def forward(self, x): x = torch.relu(self.fc1(x)) return self.fc2(x) model = SimpleModel().cuda() fsdp_model = FSDP( model, sharding_strategy=ShardingStrategy.FULL_SHARD ) x = torch.randn(8, 1024).cuda() y = fsdp_model(x) print(y.shape)

Результат выполнения кода:

torch.Size([8, 128])

Пример

Используем стратегию шардинга по градиентам для уменьшения объёма передаваемых данных между устройствами:

import torch import torch.nn as nn import torch.distributed as dist from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp import ShardingStrategy dist.init_process_group("nccl") torch.manual_seed(0) class Model(nn.Module): def __init__(self): super().__init__() self.linear1 = nn.Linear(256, 128) self.linear2 = nn.Linear(128, 64) def forward(self, x): x = torch.relu(self.linear1(x)) return self.linear2(x) model = Model().cuda() fsdp_model = FSDP( model, sharding_strategy=ShardingStrategy.SHARD_GRAD_OP ) optimizer = torch.optim.Adam(fsdp_model.parameters(), lr=0.001) x = torch.randn(4, 256).cuda() y = fsdp_model(x) loss = y.sum() loss.backward() optimizer.step() print("Loss:", loss.item())

Результат выполнения кода:

"Loss: -0.003"

Пример

Настроим выгрузку параметров на CPU для дополнительной экономии памяти на GPU:

import torch import torch.nn as nn import torch.distributed as dist from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp import ShardingStrategy from torch.distributed.fsdp import CPUOffload dist.init_process_group("nccl") class BigModel(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(2048, 2048) def forward(self, x): return self.fc(x) model = BigModel().cuda() fsdp_model = FSDP( model, sharding_strategy=ShardingStrategy.FULL_SHARD, cpu_offload=CPUOffload(offload_params=True) ) x = torch.randn(2, 2048).cuda() res = fsdp_model(x) print(res.shape)

Результат выполнения кода:

torch.Size([2, 2048])

Пример

Применим политику автоматической обёртки для вложенных модулей, чтобы снизить накладные расходы на коммуникации:

import torch import torch.nn as nn import torch.distributed as dist from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp import ShardingStrategy from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy dist.init_process_group("nccl") class TransformerBlock(nn.Module): def __init__(self): super().__init__() self.attn = nn.Linear(512, 512) self.ffn = nn.Linear(512, 512) def forward(self, x): return self.ffn(torch.relu(self.attn(x))) class Model(nn.Module): def __init__(self): super().__init__() self.blocks = nn.ModuleList([ TransformerBlock() for _ in range(4) ]) self.out = nn.Linear(512, 10) def forward(self, x): for block in self.blocks: x = block(x) return self.out(x) model = Model().cuda() fsdp_model = FSDP( model, sharding_strategy=ShardingStrategy.FULL_SHARD, auto_wrap_policy=transformer_auto_wrap_policy ) x = torch.randn(4, 512).cuda() res = fsdp_model(x) print("Output shape:", res.shape)

Результат выполнения кода:

"Output shape: torch.Size([4, 10])"

Смотрите также

  • функцию init_process_group,
    которая инициализирует распределённую среду
  • функцию all_reduce,
    которая выполняет редукцию данных между процессами
  • функцию broadcast,
    которая передаёт тензор от одного процесса всем остальным
  • функцию new_group,
    которая создаёт новую группу процессов для коммуникации
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить