Skip to content
Merged
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
1 change: 1 addition & 0 deletions megatron/core/models/gpt/gpt_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -682,6 +682,7 @@ def _postprocess(
compute_language_model_loss=self.compute_language_model_loss,
config=self.config,
cp_group=self.pg_collection.cp,
tp_group=self.tp_group,
packed_seq_params=packed_seq_params,
scale_logits_fn=self._scale_logits if self.config.use_mup else None,
input_ids=input_ids,
Expand Down
1 change: 1 addition & 0 deletions megatron/core/models/hybrid/hybrid_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -559,6 +559,7 @@ def forward(
compute_language_model_loss=self.compute_language_model_loss,
config=self.config,
cp_group=self.pg_collection.cp,
tp_group=self.tp_group,
packed_seq_params=packed_seq_params,
scale_logits_fn=self._scale_logits if self.config.use_mup else None,
input_ids=input_ids,
Expand Down
212 changes: 179 additions & 33 deletions megatron/core/transformer/multi_token_prediction.py
Original file line number Diff line number Diff line change
Expand Up @@ -340,80 +340,216 @@ def _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group=No


class MTPLossLoggingHelper:
"""Helper class for logging MTP losses."""
"""Helper class for logging MTP losses and acceptance rates."""

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This comment is not necessary to address for this PR itself, but eventually we will want to unify the MTP acceptance rate tracking for training and inference - for that effort, maybe there is some value in keeping the loss tracking separate from the acceptance rate tracking so that the latter can be shared while the former is used only for training.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Makes sense. We have this function here to calculate acceptance metrics that can be reused when things separate out.


tracker = {}

@staticmethod
def save_loss_to_tracker(
def save_metrics_to_tracker(
loss: torch.Tensor,
correct: torch.Tensor,
total: torch.Tensor,
layer_number: int,
num_layers: int,
reduce_group: Optional[torch.distributed.ProcessGroup] = None,
avg_group: Optional[torch.distributed.ProcessGroup] = None,
reduce_group: torch.distributed.ProcessGroup = None,
avg_group: torch.distributed.ProcessGroup = None,
):
"""Save the mtp loss for logging.
"""Save the mtp metrics (loss, correct, total) for logging.

Args:
loss (torch.Tensor): The loss tensor.
loss (torch.Tensor): The normalized loss value for this MTP layer.
correct (torch.Tensor): Number of correct predictions.
total (torch.Tensor): Total number of predictions.
layer_number (int): Layer index of the loss.
num_layers (int): The number of total layers.
reduce_group (torch.distributed.ProcessGroup): The group for reducing the loss.
mean_group (torch.distributed.ProcessGroup): The group for averaging the loss.
avg_group (torch.distributed.ProcessGroup): The group for averaging the loss.
"""
# Skip mtp loss logging if layer_number is None.
if layer_number is None:
return

tracker = MTPLossLoggingHelper.tracker
if "values" not in tracker:
tracker["values"] = torch.zeros(num_layers, device=torch.cuda.current_device())
tracker["values"][layer_number] += loss.detach()
if "loss_values" not in tracker:
tracker["loss_values"] = torch.zeros(num_layers, device=torch.cuda.current_device())
if "correct_values" not in tracker:
tracker["correct_values"] = torch.zeros(num_layers, device=torch.cuda.current_device())
if "total_values" not in tracker:
tracker["total_values"] = torch.zeros(num_layers, device=torch.cuda.current_device())

tracker["loss_values"][layer_number] += loss.detach()
tracker["correct_values"][layer_number] += correct.detach()
tracker["total_values"][layer_number] += total.detach()
tracker["reduce_group"] = reduce_group
tracker["avg_group"] = avg_group

def clean_loss_in_tracker():
"""Clear the mtp losses."""
@staticmethod
def clean_metrics_in_tracker():
"""Clear the mtp metrics."""
tracker = MTPLossLoggingHelper.tracker
tracker["values"].zero_()
if "loss_values" in tracker:
tracker["loss_values"].zero_()
if "correct_values" in tracker:
tracker["correct_values"].zero_()
if "total_values" in tracker:
tracker["total_values"].zero_()
tracker["reduce_group"] = None
tracker["avg_group"] = None

def reduce_loss_in_tracker():
"""Collect and reduce the mtp losses across ranks."""
@staticmethod
def reduce_metrics_in_tracker():
"""Collect and reduce the mtp metrics across ranks."""
tracker = MTPLossLoggingHelper.tracker
if "values" not in tracker:
if "loss_values" not in tracker:
return
values = tracker["values"]
# Reduce mtp losses across ranks.

loss_values = tracker["loss_values"]
if tracker.get('reduce_group') is not None:
torch.distributed.all_reduce(values, group=tracker.get('reduce_group'))
torch.distributed.all_reduce(loss_values, group=tracker.get('reduce_group'))
if tracker.get('avg_group') is not None:
torch.distributed.all_reduce(
values, group=tracker['avg_group'], op=torch.distributed.ReduceOp.AVG
loss_values, group=tracker['avg_group'], op=torch.distributed.ReduceOp.AVG
)

for key in ["correct_values", "total_values"]:
if key not in tracker:
continue
values = tracker[key]
if tracker.get('reduce_group') is not None:
torch.distributed.all_reduce(values, group=tracker.get('reduce_group'))
if tracker.get('avg_group') is not None:
torch.distributed.all_reduce(
values, group=tracker['avg_group'], op=torch.distributed.ReduceOp.SUM
)

@staticmethod
def track_mtp_metrics(loss_scale, iteration, writer, wandb_writer=None, total_loss_dict=None):
"""Track the Multi-Token Prediction (MTP) metrics for logging."""
MTPLossLoggingHelper.reduce_loss_in_tracker()
MTPLossLoggingHelper.reduce_metrics_in_tracker()
tracker = MTPLossLoggingHelper.tracker
if "values" not in tracker:
if "loss_values" not in tracker:
return
mtp_losses = tracker["values"] * loss_scale

mtp_losses = tracker["loss_values"] * loss_scale
mtp_corrects = tracker.get("correct_values", torch.zeros_like(mtp_losses))
mtp_totals = tracker.get("total_values", torch.ones_like(mtp_losses))

# Process-local logging state; cumulative rates intentionally reset after restart/resume.
if (
"cumulative_correct_values" not in tracker
or tracker["cumulative_correct_values"].shape != mtp_corrects.shape
):
tracker["cumulative_correct_values"] = torch.zeros_like(mtp_corrects)
if (
"cumulative_total_values" not in tracker
or tracker["cumulative_total_values"].shape != mtp_totals.shape
):
tracker["cumulative_total_values"] = torch.zeros_like(mtp_totals)

tracker["cumulative_correct_values"] += mtp_corrects
tracker["cumulative_total_values"] += mtp_totals
mtp_cumulative_corrects = tracker["cumulative_correct_values"]
mtp_cumulative_totals = tracker["cumulative_total_values"]

mtp_num_layers = mtp_losses.shape[0]
for i in range(mtp_num_layers):
name = f"mtp_{i + 1} loss"
loss_name = f"mtp_{i+1} loss"
step_acc_name = f"mtp_{i+1}_acceptance_rate"
cum_acc_name = f"mtp_{i+1}_cumulative_acceptance_rate"

loss = mtp_losses[i]
# Empty masks can leave no valid MTP positions, so clamp denominators to avoid NaNs.
step_rate = (mtp_corrects[i] / torch.clamp(mtp_totals[i], min=1)) * 100.0
cum_rate = (
mtp_cumulative_corrects[i] / torch.clamp(mtp_cumulative_totals[i], min=1)
) * 100.0

if total_loss_dict is not None:
if name in total_loss_dict:
total_loss_dict[name] += loss
else:
total_loss_dict[name] = loss
total_loss_dict[loss_name] = (
total_loss_dict.get(loss_name, torch.zeros_like(loss)) + loss
)

if writer is not None:
writer.add_scalar(name, loss, iteration)
writer.add_scalar(loss_name, loss, iteration)
writer.add_scalar(step_acc_name, step_rate, iteration)
writer.add_scalar(cum_acc_name, cum_rate, iteration)
if wandb_writer is not None:
wandb_writer.log({f"{name}": loss}, iteration)
wandb_writer.log({f"{loss_name}": loss}, iteration)
wandb_writer.log({f"{step_acc_name}": step_rate}, iteration)
wandb_writer.log({f"{cum_acc_name}": cum_rate}, iteration)

MTPLossLoggingHelper.clean_loss_in_tracker()
MTPLossLoggingHelper.clean_metrics_in_tracker()


def _mtp_logits_are_vocab_sharded(
output_layer: Callable, runtime_gather_output: Optional[bool]
) -> bool:
"""Return whether MTP logits are still vocab-sharded across tensor-parallel ranks."""
if runtime_gather_output is not None:
return not runtime_gather_output
return not getattr(output_layer, "gather_output", False)


def _vocab_parallel_argmax(
vocab_parallel_logits: Tensor, tp_group: torch.distributed.ProcessGroup, tp_size: int
) -> Tensor:
"""Return global argmax ids from logits sharded across the vocab dimension."""
vocab_shard_size = vocab_parallel_logits.size(-1)
local_max_vals, local_argmax = vocab_parallel_logits.max(dim=-1) # [s, b], [s, b]

gathered_max_vals = [torch.empty_like(local_max_vals) for _ in range(tp_size)]
gathered_argmax = [torch.empty_like(local_argmax) for _ in range(tp_size)]
torch.distributed.all_gather(gathered_max_vals, local_max_vals, group=tp_group)
torch.distributed.all_gather(gathered_argmax, local_argmax, group=tp_group)

stacked_max_vals = torch.stack(gathered_max_vals, dim=0)
stacked_argmax = torch.stack(gathered_argmax, dim=0)
winning_rank = stacked_max_vals.argmax(dim=0) # [s, b]
winning_local_argmax = torch.gather(stacked_argmax, 0, winning_rank.unsqueeze(0)).squeeze(
0
) # [s, b]
return winning_rank * vocab_shard_size + winning_local_argmax # [s, b]


def _compute_mtp_acceptance_counts(
mtp_logits: Tensor,
mtp_labels: Tensor,
loss_mask: Tensor,
output_layer: Callable,
runtime_gather_output: Optional[bool],
tp_group: Optional[torch.distributed.ProcessGroup] = None,
) -> tuple[Tensor, Tensor]:
"""Compute MTP acceptance correct/total counts."""
with torch.no_grad():
logits_are_vocab_sharded = _mtp_logits_are_vocab_sharded(
output_layer, runtime_gather_output
)
if (
tp_group is None
and logits_are_vocab_sharded
and parallel_state.is_initialized()
and parallel_state.get_tensor_model_parallel_world_size() > 1
):
raise ValueError(
"tp_group must be provided when computing MTP acceptance counts "
"from vocab-sharded logits under tensor model parallelism."
)
tp_size = torch.distributed.get_world_size(group=tp_group) if tp_group is not None else 1

# Apply TP rank offsets only when logits are vocab-sharded; gathered logits already
# contain global vocab ids in their last dimension.
if tp_group is not None and tp_size > 1 and logits_are_vocab_sharded:
preds = _vocab_parallel_argmax(mtp_logits, tp_group, tp_size)
else:
preds = torch.argmax(mtp_logits, dim=-1) # [s, b]

labels_match = mtp_labels.transpose(0, 1).contiguous() # [b, s] => [s, b]
mask_match = loss_mask.transpose(0, 1).contiguous() # [b, s] => [s, b]
valid_positions = mask_match.bool()
correct = ((preds == labels_match) & valid_positions).sum().float()
total = valid_positions.sum().float()

return correct, total


@dataclass
Expand Down Expand Up @@ -632,6 +768,7 @@ def process_mtp_loss(
compute_language_model_loss: Callable,
config: TransformerConfig,
cp_group: Optional[torch.distributed.ProcessGroup] = None,
tp_group: Optional[torch.distributed.ProcessGroup] = None,
packed_seq_params: Optional[PackedSeqParams] = None,
scale_logits_fn: Optional[Callable[[Tensor], Tensor]] = None,
input_ids: Optional[Tensor] = None,
Expand All @@ -652,6 +789,7 @@ def process_mtp_loss(
compute_language_model_loss (Callable): Method to compute language model loss.
config (TransformerConfig): Model configuration containing mtp_num_layers etc.
cp_group (Optional[ProcessGroup]): Context parallelism process group.
tp_group (Optional[ProcessGroup]): Tensor parallelism process group.
packed_seq_params (Optional[PackedSeqParams]): Packed sequence parameters.
scale_logits_fn (Optional[Callable[[Tensor], Tensor]]): Optional function to
scale logits before loss computation (e.g., MuP output scaling).
Expand Down Expand Up @@ -707,15 +845,23 @@ def process_mtp_loss(
loss_mask, num_tokens = roll_tensor(
loss_mask, shifts=-1, dims=-1, cp_group=cp_group, packed_seq_params=packed_seq_params
)

mtp_loss = compute_language_model_loss(mtp_labels, mtp_logits)

mtp_loss = loss_mask * mtp_loss

if is_training:
# Safe divide without sync: mask numerator when num_tokens==0, divide by clamp(min=1)
mtp_loss_for_log = (
torch.sum(mtp_loss) * (num_tokens > 0).to(mtp_loss.dtype)
) / num_tokens.clamp(min=1)
MTPLossLoggingHelper.save_loss_to_tracker(
correct, total = _compute_mtp_acceptance_counts(
mtp_logits, mtp_labels, loss_mask, output_layer, runtime_gather_output, tp_group
)

MTPLossLoggingHelper.save_metrics_to_tracker(
mtp_loss_for_log,
correct,
total,
mtp_layer_number,
config.mtp_num_layers,
avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True),
Expand Down
Loading
Loading