Skip to content
Open
Show file tree
Hide file tree
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
2 changes: 2 additions & 0 deletions megatron/core/models/hybrid/hybrid_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -736,6 +736,7 @@ def forward(
mtp_input_mask=mtp_input_mask,
packed_seq_params=packed_seq_params,
cp_batch=cp_batch,
padding_mask=padding_mask,
)
if mtp_inputs.decoder_input is None:
assert mtp_inputs.input_ids is not None and mtp_inputs.position_ids is not None, (
Expand All @@ -754,6 +755,7 @@ def forward(
embedding=self.embedding,
decoder_input=mtp_inputs.decoder_input,
mtp_input_mask=mtp_inputs.mtp_input_mask,
padding_mask=mtp_inputs.padding_mask,
packed_seq_params_by_layout=packed_seq_params_by_layout,
cp_layout_plan=cp_layout_plan,
)
Expand Down
8 changes: 8 additions & 0 deletions megatron/core/transformer/moe/moe_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -131,6 +131,14 @@ def switch_load_balancing_loss_func(
mask_expanded = padding_mask.unsqueeze(-1)
probs = probs * mask_expanded

# An entirely padded MTP depth can have no valid tokens across the reduction
# group. Its masked probabilities contribute zero loss; keep normalization
# finite so backward does not turn that zero contribution into NaNs.
if isinstance(total_num_tokens, torch.Tensor):
total_num_tokens = total_num_tokens.clamp(min=1)
else:
total_num_tokens = max(total_num_tokens, 1)

if fused:
if not HAVE_TE or fused_moe_aux_loss is None:
raise ValueError("fused_moe_aux_loss is not available. Please install TE >= 2.7.0.")
Expand Down
62 changes: 59 additions & 3 deletions megatron/core/transformer/multi_token_prediction.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
from megatron.core.process_groups_config import ProcessGroupCollection
from megatron.core.tensor_observation import is_observing_tensor, observe_tensor
from megatron.core.tensor_parallel import (
gather_from_sequence_parallel_region,
gather_from_tensor_model_parallel_region,
scatter_to_sequence_parallel_region,
)
Expand Down Expand Up @@ -1574,6 +1575,8 @@ def _get_embeddings(
hidden_states (torch.Tensor): hidden states tensor of shape [s, b, h] where s is the
sequence length, b is the batch size, and h is the hidden size.
packed_seq_params (PackedSeqParams): Parameters for packed sequence processing.
padding_mask (torch.Tensor, optional): Padding flags of shape [b, s/tp] with
sequence parallelism, otherwise [b, s]. True marks padding.
mtp_input_mask (torch.Tensor, optional): Mask of conditioning tokens backed by
regular token embeddings. Shape: [b, s].
"""
Expand Down Expand Up @@ -1613,14 +1616,34 @@ def _get_embeddings(
return_sum=False,
)
if padding_mask is not None:
padding_mask, _ = roll_tensor(
padding_mask,
# GPT has already SP-sharded this mask. Reconstruct the CP-local
# sequence before rolling so TP boundaries are not mistaken for ends
# and packed/CP metadata still describes the tensor being shifted.
if self.config.sequence_parallel:
padding_mask = gather_from_sequence_parallel_region(
padding_mask.transpose(0, 1).contiguous(),
tensor_parallel_output_grad=False,
group=self.tp_group,
).transpose(0, 1)
# roll_tensor zero-fills sequence ends. Roll validity so these new
# positions remain padding (True), including packed/CP boundaries.
valid_mask, _ = roll_tensor(
Comment thread
cuichenx marked this conversation as resolved.
~padding_mask,
shifts=-1,
dims=-1,
cp_group=self.cp_group,
packed_seq_params=packed_seq_params,
return_sum=False,
)
padding_mask = ~valid_mask
Comment thread
cuichenx marked this conversation as resolved.
Comment thread
cuichenx marked this conversation as resolved.
if self.config.sequence_parallel:
padding_mask = (
scatter_to_sequence_parallel_region(
padding_mask.transpose(0, 1).contiguous(), group=self.tp_group
)
.transpose(0, 1)
.contiguous()
)
# embedding
decoder_input = embedding(input_ids=input_ids, position_ids=position_ids)

Expand Down Expand Up @@ -2218,6 +2241,7 @@ class MultiTokenPredictionInputs:
loss_mask: Optional[Tensor]
mtp_input_mask: Optional[Tensor]
packed_seq_params: Optional[PackedSeqParams]
padding_mask: Optional[Tensor] = None


def _get_mtp_block_submodules(
Expand Down Expand Up @@ -2381,8 +2405,9 @@ def prepare_cp_layout(
mtp_input_mask: Optional[Tensor],
packed_seq_params: Optional[PackedSeqParams],
cp_batch: Optional[ContextParallelBatch],
padding_mask: Optional[Tensor] = None,
) -> MultiTokenPredictionInputs:
"""Prepare activations and token-aligned inputs for the MTP block's CP layout."""
"""Prepare MTP inputs, including the batch-major, optionally SP-sharded padding mask."""
source_layout = (
cp_batch.boundary_layout if cp_batch is not None else self.config.linear_cp_layout
)
Expand Down Expand Up @@ -2426,6 +2451,23 @@ def prepare_cp_layout(
self.tp_cp_group,
cp_batch.thd_plan,
)
if padding_mask is not None:
# Convert validity so any new THD padding slots (zero-filled by
# layout conversion) remain excluded from routing.
padding_mask = (
~convert_cp_layout(
(~padding_mask).transpose(0, 1).contiguous(),
source_layout,
target_layout,
self.cp_group,
self.sequence_parallel,
self.tp_group,
self.tp_cp_group,
cp_batch.thd_plan,
)
.transpose(0, 1)
.contiguous()
)
packed_seq_params = cp_batch.get_packed_seq_params(target_layout)
layout_batch = cp_batch.get_batch(target_layout)
input_ids = layout_batch["tokens"]
Expand All @@ -2443,6 +2485,7 @@ def prepare_cp_layout(
loss_mask=loss_mask,
mtp_input_mask=mtp_input_mask,
packed_seq_params=packed_seq_params,
padding_mask=padding_mask,
)

def _build_layers(self, pg_collection):
Expand Down Expand Up @@ -2615,6 +2658,19 @@ def forward(
packed_seq_params=packed_seq_params,
return_sum=False,
)
if padding_mask is not None:
# Precomputed embeddings bypass the layer's _get_embeddings,
# so shift validity here alongside those embeddings.
valid_mask, _ = roll_tensor_precomputed_embeddings(
(~padding_mask).transpose(0, 1).contiguous(),
shifts=-1,
dims=0,
sp_group=self.tp_group if self.sequence_parallel else None,
cp_group=self.cp_group,
packed_seq_params=packed_seq_params,
return_sum=False,
)
padding_mask = ~valid_mask.transpose(0, 1).contiguous()

# Older HSM entries predict earlier targets than the newest entry. Roll
# them once per depth so all candidates correspond to the same target.
Expand Down
13 changes: 10 additions & 3 deletions tests/unit_tests/determinism/kernels/test_moe_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -237,18 +237,25 @@ def test_group_limited_topk_replays():
),
],
)
def test_switch_load_balancing_loss_replays(fused):
@pytest.mark.parametrize("all_padding", [False, True])
def test_switch_load_balancing_loss_replays(fused, all_padding):
seeded()
routing_map, probs = _routing(num_tokens=65536, num_experts=256, topk=8)
if all_padding:
routing_map.zero_()
probs.zero_()
probs = probs.detach().requires_grad_(True)
tokens_per_expert = routing_map.sum(dim=0)
total_num_tokens = tokens_per_expert.sum() // 8 if all_padding else 65536

def fn(probs):
return moe_utils.switch_load_balancing_loss_func(
probs, tokens_per_expert, 65536, 8, 256, 1e-2, fused=fused
probs, tokens_per_expert, total_num_tokens, 8, 256, 1e-2, fused=fused
)

assert_replays_bit_exact(fn, (probs,), replays=4, what=f"aux loss[fused={fused}]")
assert_replays_bit_exact(
fn, (probs,), replays=4, what=f"aux loss[fused={fused}, all_padding={all_padding}]"
)


@pytest.mark.parametrize("router_dtype", [torch.float32, torch.float64])
Expand Down
55 changes: 55 additions & 0 deletions tests/unit_tests/transformer/moe/test_routers.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@
from megatron.core.transformer.moe.moe_layer import MoELayer, MoESubmodules
from megatron.core.transformer.moe.moe_logging import MoEMetricsTracker
from megatron.core.transformer.moe.moe_utils import (
fused_compute_score_for_moe_aux_loss,
fused_moe_aux_loss,
get_default_pg_collection,
get_updated_expert_bias,
router_gating_linear,
Expand Down Expand Up @@ -296,6 +298,59 @@ def test_aux_loss(self):
out.sum().mul_(0).backward()
assert self.sequential_mlp.router.weight.grad.abs().sum() > 0

@pytest.mark.internal
@pytest.mark.parametrize("loss_type", ["aux_loss", "seq_aux_loss", "global_aux_loss"])
@pytest.mark.parametrize("fused", [False, True])
@pytest.mark.parametrize("with_history", [False, True])
def test_aux_loss_with_entirely_padded_input(self, monkeypatch, loss_type, fused, with_history):
"""Zero valid tokens contribute finite zero auxiliary gradients, even with history."""
if fused and (fused_moe_aux_loss is None or fused_compute_score_for_moe_aux_loss is None):
pytest.skip("TE fused auxiliary-loss operators are unavailable")
config = TransformerConfig(
num_layers=1,
hidden_size=16,
num_attention_heads=2,
num_moe_experts=4,
moe_router_topk=2,
moe_router_load_balancing_type=loss_type,
moe_aux_loss_coeff=0.001,
moe_router_aux_loss_fusion=fused,
moe_router_dtype="fp32",
params_dtype=torch.float32,
add_bias_linear=False,
)
losses = []
compute_aux_loss = router_mod.switch_load_balancing_loss_func

def record_aux_loss(*args, **kwargs):
loss = compute_aux_loss(*args, **kwargs)
losses.append(loss.detach())
return loss

monkeypatch.setattr(router_mod, "switch_load_balancing_loss_func", record_aux_loss)
router = TopKRouter(config, pg_collection=get_default_pg_collection()).cuda()
hidden_states = torch.randn(8, 2, 16, device="cuda", requires_grad=True)
if with_history:
probs, _ = router(hidden_states)
# Isolate the attached auxiliary gradient from the routing output loss.
(probs.sum() * 0).backward()
assert torch.isfinite(router.weight.grad).all()
assert torch.count_nonzero(router.weight.grad) > 0
router.zero_grad(set_to_none=True)
hidden_states.grad = None
previous_counts = (
router.global_tokens_per_expert.clone() if loss_type == "global_aux_loss" else None
)
padding_mask = torch.ones(8, 2, dtype=torch.bool, device="cuda")
probs, _ = router(hidden_states, padding_mask=padding_mask)
(probs.sum() * 0).backward()
assert len(losses) == (2 if with_history else 1)
torch.testing.assert_close(losses[-1], torch.zeros_like(losses[-1]))
torch.testing.assert_close(router.weight.grad, torch.zeros_like(router.weight.grad))
torch.testing.assert_close(hidden_states.grad, torch.zeros_like(hidden_states.grad))
if previous_counts is not None:
torch.testing.assert_close(router.global_tokens_per_expert, previous_counts)

@pytest.mark.internal
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_router_with_padding_mask(self):
Expand Down
Loading
Loading