Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 23 additions & 17 deletions src/accelerate/accelerator.py
Original file line number Diff line number Diff line change
Expand Up @@ -1588,33 +1588,39 @@ def _prepare_tp(self, *args):

old_named_params = self._get_named_parameters(*tuple(result), drop_refs=True)

for arg in result:
if not isinstance(arg, torch.nn.Module):
continue
from torch.distributed.tensor import DTensor

from torch.distributed.tensor import DTensor, Replicate
from transformers.integrations.tensor_parallel import ReplicateParallel
if self.is_fsdp2:
for arg in result:
if not isinstance(arg, torch.nn.Module):
continue
Comment thread
SunMarc marked this conversation as resolved.

model: torch.nn.Module = arg
tp_plan = ReplicateParallel
from torch.distributed.tensor import Replicate
from transformers.integrations.tensor_parallel import ReplicateParallel

for name, param in model.named_parameters():
if isinstance(param, DTensor):
continue
model: torch.nn.Module = arg
tp_plan = ReplicateParallel

dp = DTensor.from_local(param, device_mesh=device_mesh["tp"], placements=[Replicate()])
param_name, param_type = name.rsplit(".", 1)
module_to_tp = model.get_submodule(param_name)
for name, param in model.named_parameters():
if isinstance(param, DTensor):
continue

dp = DTensor.from_local(param, device_mesh=device_mesh["tp"], placements=[Replicate()])
param_name, param_type = name.rsplit(".", 1)
module_to_tp = model.get_submodule(param_name)

tp_plan().prepare_module_tp(module_to_tp, device_mesh["tp"])
if not isinstance(dp, torch.nn.Parameter):
dp = torch.nn.Parameter(dp, requires_grad=param.requires_grad)
setattr(module_to_tp, param_type, dp)
tp_plan().prepare_module_tp(module_to_tp, device_mesh["tp"])
if not isinstance(dp, torch.nn.Parameter):
dp = torch.nn.Parameter(dp, requires_grad=param.requires_grad)
setattr(module_to_tp, param_type, dp)

new_named_params = self._get_named_parameters(*tuple(result), drop_refs=False)
# Build a map from old to new params
mapping = {p: new_named_params[n] for n, p in old_named_params.items()}

if not mapping:
return result

def _get_tensor_address(p):
if isinstance(p, DTensor):
return p._local_tensor.data_ptr()
Expand Down
Loading