From c972b4a80f3b9a84900d695252776501ea3d9ca3 Mon Sep 17 00:00:00 2001 From: xinqi Date: Tue, 16 Jun 2026 15:19:56 +0800 Subject: [PATCH] Add TV (total-variation) loss option for the MTP draft head MTP draft heads are trained with cross-entropy, but for speculative decoding the rejection-sampling acceptance rate is accept_rate = 1 - d_TV(p_main, q_draft), so total variation is the *direct* training target (CE only bounds it via Pinsker). This adds an opt-in TV loss for the MTP head: - TransformerConfig.mtp_loss_type ('ce' default | 'tv') + mtp_tv_chunk (int), auto-exposed as --mtp-loss-type / --mtp-tv-chunk via ArgumentGroupFactory. - GPTModel MTP loss: when mtp_loss_type=='tv', compute per-token 1 - sum_v min(softmax(p_main).detach(), softmax(q_draft)); p_main is the main model's next-token distribution (detached; hidden states rolled to align with the MTP target). Drop-in for the CE branch -> reuses MTPLossAutoScaler path. - mtp_tv_chunk>0 chunks the vocab-sum over the sequence dim and gradient- checkpoints each chunk, so backward memory is ~1 chunk instead of O(T*V), letting TV fit alongside full-backbone joint training. Limitation: TV needs the full vocab per rank for sum_v min(p,q), so it requires tensor-model-parallel size 1 (raises NotImplementedError for TP>1). Ref: "Breaking Entropy Bounds: Accelerating RL Training via MTP with Rejection Sampling" (Bebop). Co-Authored-By: Claude Opus 4.8 (1M context) --- megatron/core/models/gpt/gpt_model.py | 106 +++++++++++++++++- .../core/transformer/transformer_config.py | 13 +++ 2 files changed, 118 insertions(+), 1 deletion(-) diff --git a/megatron/core/models/gpt/gpt_model.py b/megatron/core/models/gpt/gpt_model.py index a83caf998c5..fb102975797 100644 --- a/megatron/core/models/gpt/gpt_model.py +++ b/megatron/core/models/gpt/gpt_model.py @@ -6,6 +6,7 @@ from typing import Dict, Literal, Optional import torch +import torch.utils.checkpoint from torch import Tensor from megatron.core import parallel_state, tensor_parallel @@ -46,6 +47,52 @@ ) +# --- MTP total-variation (TV) distillation loss for the draft head ----------- +# Paper: "Breaking Entropy Bounds: Accelerating RL Training via MTP with +# Rejection Sampling" (Bebop). Rejection-sampling accept_rate = 1 - d_TV(p, q), +# so TV distance is the *direct* training target for the draft head (vs CE, +# which only bounds it via Pinsker). +# +# Selected via TransformerConfig.mtp_loss_type ('ce' default | 'tv') and +# TransformerConfig.mtp_tv_chunk (sequence-dim chunk for the vocab-sum, 0=off). +# +# TP=1 ONLY: with tensor-model-parallel size 1 the full vocab lives on each rank, +# so softmax + sum_v min(p,q) over the last dim is exact (no vocab-parallel +# all-reduce). A guard raises if TP>1. + + +def _mtp_tv_per_token(draft_logits: Tensor, target_logits: Tensor, chunk: int = 0) -> Tensor: + """Per-position TV distance 1 - sum_v min(softmax(target), softmax(draft)). + + Gradients flow through ``draft_logits`` (the MTP head); ``target_logits`` + (the main model) must already be detached. Inputs are [s, b, V]; returns + [s, b]. Computed in fp32 for numerical stability. + """ + def _overlap(d: Tensor, t: Tensor) -> Tensor: + q = torch.softmax(d.float(), dim=-1) + p = torch.softmax(t.float(), dim=-1) + return torch.minimum(p, q).sum(dim=-1) + + if chunk and chunk > 0: + outs = [] + for i in range(0, draft_logits.shape[0], chunk): + d_i = draft_logits[i : i + chunk] + t_i = target_logits[i : i + chunk] + # Gradient-checkpoint each chunk: keep only the per-token overlap [chunk] in the + # forward graph and recompute softmax(draft) in backward. Without this, every chunk's + # [chunk,V] softmax graph is retained until backward (~O(T*V) ≈ 12GiB @ T~12k), which + # OOMs joint training (full backbone) — chunking the forward alone does NOT fix backward + # retention. _overlap is deterministic (no RNG), so recompute is numerically exact. + if torch.is_grad_enabled() and d_i.requires_grad: + outs.append(torch.utils.checkpoint.checkpoint(_overlap, d_i, t_i, use_reentrant=False)) + else: + outs.append(_overlap(d_i, t_i)) + overlap = torch.cat(outs, dim=0) + else: + overlap = _overlap(draft_logits, target_logits) + return 1.0 - overlap + + def apply_true_on_policy_logits_contract( logits: Optional[Tensor], *, vocab_size: Optional[int] = None ) -> Optional[Tensor]: @@ -734,7 +781,64 @@ def _postprocess( # Compute mtp loss without storing logits to save memory. # Detach output layer params to prevent MTP gradient flowing # to output layer (for RL training). - if self.fuse_linear_cross_entropy: + if self.config.mtp_loss_type == "tv": + # --- TV distillation loss (Bebop) ------------------------- + # accept_rate = 1 - d_TV(p_main, q_draft); minimize d_TV. + if parallel_state.get_tensor_model_parallel_world_size() != 1: + raise NotImplementedError( + "MTP_LOSS_TYPE=tv supports tensor-model-parallel size 1 " + "only (full vocab per rank); got TP>1." + ) + output_layer_params = { + k: v.detach() for k, v in self.output_layer.named_parameters() + } + output_layer_buffers = dict(self.output_layer.named_buffers()) + # q: draft-head logits (grads flow through the MTP hidden states) + mtp_logits, _ = torch.func.functional_call( + self.output_layer, + {**output_layer_params, **output_layer_buffers}, + (hidden_states_list[mtp_layer_number + 1],), + { + 'weight': output_weight.detach() if output_weight is not None else None, + 'runtime_gather_output': runtime_gather_output, + }, + ) + # p: main-model logits, fully detached. Align with mtp_labels + # by left-rolling the main hidden states (k+1) steps along + # seq (main predicts t+1; the k-th MTP target is t+k+2), + # mirroring the roll_tensor label shifts above. Roll hidden + # states (H) rather than logits (V) to save memory; the + # zero-padded boundary position is masked by loss_mask. + with torch.no_grad(): + # hidden states are [s, b, H] (seq-first). roll_tensor's + # packed-sequence path requires the seq dim to be LAST + # (it indexes tensor[..., cu_seqlens[i]:...]); the label + # roll above works because labels are [b, s]. So move seq + # to the last dim, roll (k+1) times to match mtp_labels' + # (k+1) in-loop shifts, then move it back. + main_hidden = hidden_states_list[0].movedim(0, -1) # [b, H, s] + for _ in range(mtp_layer_number + 1): + main_hidden, _ = roll_tensor( + main_hidden, + shifts=-1, + dims=-1, + cp_group=self.cp_group, + packed_seq_params=packed_seq_params, + ) + main_hidden = main_hidden.movedim(-1, 0).contiguous() # [s, b, H] + main_logits, _ = torch.func.functional_call( + self.output_layer, + {**output_layer_params, **output_layer_buffers}, + (main_hidden,), + { + 'weight': output_weight.detach() if output_weight is not None else None, + 'runtime_gather_output': runtime_gather_output, + }, + ) + # per-token TV [s, b] -> [b, s] to match loss_mask / CE layout + mtp_loss = _mtp_tv_per_token(mtp_logits, main_logits, chunk=self.config.mtp_tv_chunk) + mtp_loss = mtp_loss.transpose(0, 1).contiguous() + elif self.fuse_linear_cross_entropy: mtp_loss = self.output_layer( output_cross_entropy_loss=self.fuse_linear_cross_entropy, labels=mtp_labels, diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index 5c495ea8b2c..7e181e35934 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -54,6 +54,19 @@ class TransformerConfig(ModelParallelConfig): by using D sequential modules to predict D additional tokens. """ + mtp_loss_type: str = 'ce' + """Loss used to train the MTP draft head: 'ce' (cross-entropy, default) or 'tv' + (total-variation distillation to the main model's next-token distribution, + 1 - sum_v min(softmax(p_main), softmax(q_draft))). TV is the direct objective for + rejection-sampling acceptance rate (accept_rate = 1 - d_TV); CE only bounds it via + Pinsker. 'tv' requires tensor-model-parallel size 1 (full vocab per rank). + """ + + mtp_tv_chunk: int = 0 + """For mtp_loss_type='tv': chunk size along the sequence dim for the vocab-sum, + to cap peak/backward memory (each chunk is gradient-checkpointed). 0 = no chunking. + """ + mtp_loss_scaling_factor: Optional[float] = 0.1 """Weighting factor of Multi-Token Prediction (MTP) loss. We compute the average of the MTP losses across all depths,