From 9401db3f290cb1bdffbd29c99114d3d68e2e0c0b Mon Sep 17 00:00:00 2001 From: Zhiyao Jiang Date: Tue, 18 Aug 2026 00:38:24 -0400 Subject: [PATCH] [AMD] Scope the EAGLE greedy-verify TP broadcast to ROCm only (#35195) --- python/sglang/srt/speculative/eagle_utils.py | 39 ++++++++++++++------ 1 file changed, 28 insertions(+), 11 deletions(-) diff --git a/python/sglang/srt/speculative/eagle_utils.py b/python/sglang/srt/speculative/eagle_utils.py index 18c9f0cbd..c4e23ec62 100644 --- a/python/sglang/srt/speculative/eagle_utils.py +++ b/python/sglang/srt/speculative/eagle_utils.py @@ -737,6 +737,21 @@ def eagle_sample( target_predict=target_predict, topk=verify_input.tree_topk, ) + + if _is_hip: + # On ROCm, the per-rank draft tokens can differ, so ranks accept a + # different number of drafts, desynchronize the committed seq_lens, and + # deadlock the next TP collective. Broadcast from rank 0 to ensure + # consistency, the same way the sampling branch below does. + tp_group = ( + get_parallel().attn_tp_group + if is_dp_attention_enabled() + else get_tp_group() + ) + if tp_group.world_size > 1: + tp_group.broadcast(predict, src=0) + tp_group.broadcast(accept_index, src=0) + tp_group.broadcast(num_correct_drafts, src=0) else: from sgl_kernel import ( top_k_renorm_prob, @@ -824,17 +839,19 @@ def eagle_sample( deterministic=True, ) - # Sync the verify decision across TP ranks: small per-rank differences in the - # tensors feeding this point can make one rank accept a different number of - # drafts, which desynchronizes the committed seq_lens and deadlocks the next - # TP collective. Broadcast from rank 0 to ensure consistency. - tp_group = ( - get_parallel().attn_tp_group if is_dp_attention_enabled() else get_tp_group() - ) - if tp_group.world_size > 1: - tp_group.broadcast(predict, src=0) - tp_group.broadcast(accept_index, src=0) - tp_group.broadcast(num_correct_drafts, src=0) + # Sync sampling results across TP ranks: different GPUs may + # produce slightly different target_probs due to floating-point + # non-determinism in softmax/top_k/top_p, causing different + # sampled tokens. Broadcast from rank 0 to ensure consistency. + tp_group = ( + get_parallel().attn_tp_group + if is_dp_attention_enabled() + else get_tp_group() + ) + if tp_group.world_size > 1: + tp_group.broadcast(predict, src=0) + tp_group.broadcast(accept_index, src=0) + tp_group.broadcast(num_correct_drafts, src=0) if SIMULATE_ACC_LEN > 0: # Do simulation. The helper builds (and returns) a replacement