Skip to content

Commit 19b794b

Browse files
author
Tom Long
committed
Fix Mamba conv params under fine-grained FSDP gather
Mamba's fused path reads conv1d weights directly instead of calling Conv1d.forward(), so fine-grained Megatron-FSDP never gathered those child parameters before the second forward. Register the conv module as an extra forward-gather source and resolve context-parallel Mamba params from the live mixer object.
1 parent 83acf2d commit 19b794b

3 files changed

Lines changed: 40 additions & 13 deletions

File tree

‎megatron/core/distributed/fsdp/src/megatron_fsdp/megatron_fsdp.py‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -759,6 +759,17 @@ def _pre_forward_param_unshard(module: nn.Module, *unused):
759759
if self.enable_fine_grained_param_gather_hook:
760760
param_list = list(module.parameters(recurse=False))
761761

762+
extra_forward_param_modules = getattr(module, "_fsdp_extra_forward_param_modules", ())
763+
if isinstance(extra_forward_param_modules, nn.Module):
764+
extra_forward_param_modules = (extra_forward_param_modules,)
765+
if extra_forward_param_modules:
766+
seen_param_ids = {id(param) for param in param_list}
767+
for extra_module in extra_forward_param_modules:
768+
for extra_param in extra_module.parameters():
769+
if id(extra_param) not in seen_param_ids:
770+
param_list.append(extra_param)
771+
seen_param_ids.add(id(extra_param))
772+
762773
# All-gather the parameters before the forward pass.
763774
self.all_gather_and_wait_parameters_ready(
764775
params=param_list,

‎megatron/core/ssm/mamba_context_parallel.py‎

Lines changed: 12 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -70,6 +70,7 @@ def __init__(
7070
A_log_cp1: torch.Tensor,
7171
D_cp1: torch.Tensor,
7272
D_has_hdim: bool,
73+
mixer=None,
7374
) -> None:
7475
if not HAVE_EINOPS:
7576
raise ImportError("einops is required by the Mamba model but cannot be imported")
@@ -84,6 +85,7 @@ def __init__(
8485
self.A_log_cp1 = A_log_cp1
8586
self.D_cp1 = D_cp1
8687
self.D_has_hdim = D_has_hdim
88+
self._mixer = mixer
8789

8890
self.cp_size = self.cp_group.size()
8991

@@ -231,24 +233,29 @@ def conv1d_channels(self):
231233
def get_conv1d_weight(self) -> torch.Tensor:
232234
"""Returns a slice of the conv1d weight relevant to the current context parallel rank"""
233235
# weight shape: [conv_dim, 1, d_conv]
234-
return self._slice_conv_param(self.conv1d_cp1.weight)
236+
conv1d = self._mixer.conv1d if self._mixer is not None else self.conv1d_cp1
237+
return self._slice_conv_param(conv1d.weight)
235238

236239
def get_conv1d_bias(self) -> torch.Tensor:
237240
"""Returns a slice of the conv1d bias relevant to the current context parallel rank"""
238241
# bias shape: [conv_dim]
239-
return self._slice_conv_param(self.conv1d_cp1.bias)
242+
conv1d = self._mixer.conv1d if self._mixer is not None else self.conv1d_cp1
243+
return self._slice_conv_param(conv1d.bias)
240244

241245
def get_dt_bias(self) -> torch.Tensor:
242246
"""Returns a slice of dt_bias relevant to the current context parallel rank"""
243-
return self._slice_vector_param(self.dt_bias_cp1)
247+
param = self._mixer.dt_bias if self._mixer is not None else self.dt_bias_cp1
248+
return self._slice_vector_param(param)
244249

245250
def get_A_log(self) -> torch.Tensor:
246251
"""Returns a slice of A_log relevant to the current context parallel rank"""
247-
return self._slice_vector_param(self.A_log_cp1)
252+
param = self._mixer.A_log if self._mixer is not None else self.A_log_cp1
253+
return self._slice_vector_param(param)
248254

249255
def get_D(self) -> torch.Tensor:
250256
"""Returns a slice of D relevant to the current context parallel rank"""
251-
return self._slice_vector_param(self.D_cp1, has_hdim=self.D_has_hdim)
257+
param = self._mixer.D if self._mixer is not None else self.D_cp1
258+
return self._slice_vector_param(param, has_hdim=self.D_has_hdim)
252259

253260
def _slice_conv_param(self, param: torch.Tensor) -> torch.Tensor:
254261
"""

‎megatron/core/ssm/mamba_mixer.py‎

Lines changed: 17 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -316,6 +316,9 @@ def __init__(
316316
else:
317317
nn.init.kaiming_uniform_(self.conv1d.weight, a=math.sqrt(5))
318318

319+
# The fused Mamba path reads conv1d weights directly instead of calling conv1d.forward().
320+
self._fsdp_extra_forward_param_modules = (self.conv1d,)
321+
319322
self.activation = "silu"
320323
self.act = nn.SiLU()
321324

@@ -408,6 +411,7 @@ def __init__(
408411
A_log_cp1=self.A_log,
409412
D_cp1=self.D,
410413
D_has_hdim=self.D_has_hdim,
414+
mixer=self,
411415
)
412416
self.tp_group = pg_collection.tp
413417

@@ -700,17 +704,22 @@ def _ssm_training(
700704
assert sequence_packing_available, reason_for_no_sequence_packing
701705
seq_idx = packed_seq_params.seq_idx
702706

707+
conv1d_weight = rearrange(self.cp.get_conv1d_weight(), "d 1 w -> d w")
708+
conv1d_bias = self.cp.get_conv1d_bias()
709+
dt_bias = self.cp.get_dt_bias().float()
710+
D = (
711+
rearrange(self.cp.get_D().float(), "(h p) -> h p", p=self.headdim)
712+
if self.D_has_hdim
713+
else self.cp.get_D()
714+
)
715+
703716
y = mamba_split_conv1d_scan_combined(
704717
zxBCdt,
705-
rearrange(self.cp.get_conv1d_weight(), "d 1 w -> d w"),
706-
self.cp.get_conv1d_bias(),
707-
self.cp.get_dt_bias().float(),
718+
conv1d_weight,
719+
conv1d_bias,
720+
dt_bias,
708721
A,
709-
D=(
710-
rearrange(self.cp.get_D().float(), "(h p) -> h p", p=self.headdim)
711-
if self.D_has_hdim
712-
else self.cp.get_D()
713-
),
722+
D=D,
714723
chunk_size=self.chunk_size,
715724
activation=self.activation,
716725
headdim=None if self.D_has_hdim else self.headdim,

0 commit comments

Comments
 (0)