From 3a94f755567aa945b42f2b3018a508b628c9c9eb Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Thu, 16 Jul 2026 19:06:52 -0700 Subject: [PATCH] [https://nvbugs/6460072][fix] Restore LM-head-TP group argmax for MTP-Eagle greedy draft under attention DP The draft_sampler consolidation dropped the ADP+LM-head-TP branch: under enable_lm_head_tp_in_adp the MTP shared head returns the LM-head-TP group's row-stacked, vocab-sharded logits, but the worker trimmed rows [:token_count] (group rank 0's requests, not this rank's) and greedy sampling fell through to a plain per-rank argmax over a local vocab shard, producing wrong draft tokens on every rank. Restore the group-wide distributed argmax in greedy_sample_draft_with_tp_gather: combine each rank's local (global index, max) over mapping_lm_head_tp, slice this rank's own row segment at tp_rank * max_num_requests, and only then trim the padding to token_count in the caller. Advanced (rejection) draft sampling asserts it never sees this layout, as rejection is config-gated off under attention DP. Signed-off-by: ZhaoyangWang --- tensorrt_llm/_torch/speculative/eagle3.py | 21 ++++----- tensorrt_llm/_torch/speculative/interface.py | 47 +++++++++++++++----- 2 files changed, 48 insertions(+), 20 deletions(-) diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index d9b5be1ecc1e..731d0236f049 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -894,18 +894,17 @@ def _forward_linear_draft_loop(self, inputs, attn_metadata, spec_metadata, self._d2t, draft_step=i) - # When ADP+LM-head-TP pads logits to max_num_requests, the - # padded rows are zero-filled placeholders only required so - # every TP rank produces logits of identical shape for the - # LM-head-TP all-gather. Drop them *before* sampling: the - # per-request sampling params (temperatures/top_k/top_p) are - # sized to token_count (== batch_size), so the padded logits - # would otherwise fail to broadcast in apply_temperature. This - # also keeps next_draft_tokens and the draft_probs buffer - # token_count-sized without a post-hoc trim. + # ADP+LM-head-TP logits are the LM-head-TP group's row-stacked + # batch (each rank's rows padded to max_num_requests, then + # all-gathered along dim 0) with the vocab sharded across the + # group. Rows [:token_count] would be group rank 0's requests, + # not this rank's, and a per-rank argmax would return a + # shard-local index -- so keep the full stacked logits and let + # greedy_sample_draft_with_tp_gather combine the group's vocab + # shards and slice this rank's own row segment; only then trim + # the max_num_requests padding down to token_count. mapping_lm_head_tp = None if use_lm_head_tp_in_adp: - logits = logits[:token_count] # The MTP head built this per-forward mapping when producing # the vocab-sharded logits; the sampler needs it to gather. mapping_lm_head_tp = getattr( @@ -917,6 +916,8 @@ def _forward_linear_draft_loop(self, inputs, attn_metadata, spec_metadata, batch_size, draft_step=i, mapping_lm_head_tp=mapping_lm_head_tp) + if use_lm_head_tp_in_adp: + new_draft_token = new_draft_token[:token_count] next_draft_tokens.append(new_draft_token) # Update hidden states for the next iteration. diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 308fc86ad6b9..48ae77f716e8 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -1378,12 +1378,15 @@ def maybe_gather_sharded_draft_logits(self, (see ``_draft_logits_are_sharded``); replicated full-vocab logits are returned unchanged. - Plain TP gathers vocab shards over ``self.mapping``. Under ADP + LM-head - TP the worker has already trimmed the LM-head-TP padding rows so each rank - holds ``[token_count, vocab_shard]`` for its own tokens; a vocab-dim - all-gather over ``mapping_lm_head_tp`` restores full vocab (no token - re-slice is needed after the trim). - """ + Plain TP gathers vocab shards over ``self.mapping``. ADP + LM-head TP + never reaches this path: rejection sampling (the only consumer of + advanced draft sampling) is config-gated off under attention DP, and + the group-stacked sharded logits it produces are handled by the greedy + path in ``greedy_sample_draft_with_tp_gather``. + """ + assert mapping_lm_head_tp is None, ( + "Advanced draft sampling is not supported under ADP + LM-head TP " + "(rejection sampling is config-gated off with attention DP)") if (spec_metadata is None or spec_metadata.is_all_greedy_sample or not self._draft_logits_are_sharded(logits, spec_metadata)): return logits @@ -1750,7 +1753,26 @@ def greedy_sample_draft_with_tp_gather(self, vocab-sharded (see ``_draft_logits_are_sharded``) -- e.g. a borrowed or gathered full-vocab draft head. Returns tokens in draft-vocab space (the caller applies d2t). Expects 2D ``[num_tokens, vocab_shard]`` logits. + + Under ADP + LM-head TP (``mapping_lm_head_tp`` given) the logits are the + LM-head-TP group's row-stacked batch (``tp_size`` segments of + ``max_num_requests`` padded rows, all-gathered along dim 0 by the MTP + shared head) with the vocab sharded across the group. The global argmax + must combine the group's vocab shards, and each rank must read its own + row segment at offset ``tp_rank * max_num_requests`` -- NOT rows + ``[:batch]``, which belong to group rank 0. """ + if (mapping_lm_head_tp is not None + and getattr(mapping_lm_head_tp, "tp_size", 1) > 1): + from ..distributed.ops import allgather + combined = self._get_local_max_and_combined(logits, + mapping_lm_head_tp) + gathered = allgather(combined, mapping_lm_head_tp, dim=-1) + group_size = mapping_lm_head_tp.tp_size + local_rows = logits.shape[0] // group_size + own_segment = gathered.view(group_size, local_rows, + -1)[mapping_lm_head_tp.tp_rank] + return self._get_draft_tokens_from_gathered(own_segment) mapping = self.mapping sharded = self._draft_logits_are_sharded(logits, spec_metadata) if (sharded and mapping is not None @@ -1760,9 +1782,9 @@ def greedy_sample_draft_with_tp_gather(self, combined = self._get_local_max_and_combined(logits) gathered = allgather(combined, mapping, dim=-1) return self._get_draft_tokens_from_gathered(gathered) - # No TP gather for attention-DP (incl. ADP + LM-head TP): each rank owns - # its own requests, so a per-rank argmax is the correct proposal and a - # cross-rank gather here would desync the ranks (see + # No cross-rank gather for plain attention-DP: each rank owns its own + # requests with replicated full-vocab logits, so a per-rank argmax is + # the correct proposal and a gather would desync the ranks (see # _draft_logits_are_sharded). Plain argmax; caller applies d2t. return torch.argmax(logits, dim=-1).type(torch.int32) @@ -1908,7 +1930,12 @@ def sample_draft_tokens(self, tokens = self.greedy_sample_draft_with_tp_gather( logits.reshape(-1, logits.shape[-1]), spec_metadata, mapping_lm_head_tp) - tokens = tokens.reshape(batch_shape) + if mapping_lm_head_tp is None: + tokens = tokens.reshape(batch_shape) + # else: ADP+LM-head-TP (2D step form only) -- the sampler returned + # this rank's own row segment, 1/tp_size of the stacked input rows, + # so the input batch shape no longer applies. Keep as-is; the + # caller trims the max_num_requests padding to token_count. else: # Advanced sampling gathers the vocab-sharded draft logits to full # vocab, then samples (scattering this step's proposal distribution