Класс 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,
которая создаёт новую группу процессов для коммуникации