From 582389cec557c613f87a054c6968bdbf52165f20 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Mon, 21 Sep 2026 17:23:15 -0700 Subject: [PATCH] [Fix] Keep diffusion encoder TP context bindings consistent (#40646) --- .../runtime/distributed/parallel_state.py | 3 + ...est_component_accuracy_parallel_runtime.py | 77 ++++++++++++------- 2 files changed, 54 insertions(+), 26 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py index 6dad38a9b..7f76a04fc 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py +++ b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py @@ -699,6 +699,9 @@ def use_tensor_parallel_group(tp_group: GroupCoordinator): tp_size=tp_group.world_size, tp_rank=tp_group.rank_in_group, tp_group=tp_group, + attn_tp_group=tp_group, + attn_tp_rank=tp_group.rank_in_group, + moe_tp_rank=tp_group.rank_in_group, # Only tensor parallelism folds here, so every other dimension is # one and the quotients come out of the shared derivation. **derive_parallel_widths( diff --git a/python/sglang/multimodal_gen/test/unit/test_component_accuracy_parallel_runtime.py b/python/sglang/multimodal_gen/test/unit/test_component_accuracy_parallel_runtime.py index 30502af55..81b7d3582 100644 --- a/python/sglang/multimodal_gen/test/unit/test_component_accuracy_parallel_runtime.py +++ b/python/sglang/multimodal_gen/test/unit/test_component_accuracy_parallel_runtime.py @@ -18,7 +18,7 @@ from sglang.multimodal_gen.test.single_test_file.component_accuracy.utils import initialize_parallel_runtime, ) from sglang.srt.distributed import parallel_state as srt_parallel_state -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import ParallelContext, get_parallel _UTILS = "sglang.multimodal_gen.test.single_test_file.component_accuracy.utils" @@ -178,62 +178,87 @@ def test_srt_owned_groups_are_not_overwritten_or_cleared(): assert srt_parallel_state._ATTN_TP is srt_attention_tp_group -def test_srt_tp_groups_follow_encoder_folding_context(): - original_diffusion_tp_group = object() - original_srt_tp_group = object() - original_srt_attention_tp_group = object() - folding_tp_group = _tp_group(world_size=2, rank_in_group=1) +@pytest.mark.parametrize("rank", [0, 1]) +def test_srt_tp_groups_follow_encoder_folding_context(rank): + original_tp_group = _tp_group() + folding_tp_group = _tp_group(world_size=2, rank_in_group=rank) with ( - patch.object(parallel_state, "_TP", original_diffusion_tp_group), - patch.object(srt_parallel_state, "_TP", original_srt_tp_group), - patch.object( - srt_parallel_state, - "_ATTN_TP", - original_srt_attention_tp_group, - ), + patch.object(parallel_state, "_TP", original_tp_group), + patch.object(srt_parallel_state, "_TP", None), + patch.object(srt_parallel_state, "_ATTN_TP", None), + patch("sglang.srt.runtime_context._PARALLEL", ParallelContext()), ): + # Diffusion initialization stamps the original TP=1 group before an + # encoder temporarily folds the sequence-parallel group into TP=2. + parallel_state._sync_srt_tp_group() with parallel_state.use_tensor_parallel_group(folding_tp_group): assert parallel_state._TP is folding_tp_group assert srt_parallel_state._TP is folding_tp_group assert srt_parallel_state._ATTN_TP is folding_tp_group assert get_parallel().tp_size == 2 - assert get_parallel().tp_rank == 1 + assert get_parallel().tp_rank == rank assert get_parallel().tp_group is folding_tp_group + assert get_parallel().attn_tp_size == 2 + assert get_parallel().attn_tp_rank == rank + assert get_parallel().attn_tp_group is folding_tp_group + assert get_parallel().moe_tp_rank == rank - assert parallel_state._TP is original_diffusion_tp_group - assert srt_parallel_state._TP is original_srt_tp_group - assert srt_parallel_state._ATTN_TP is original_srt_attention_tp_group + assert parallel_state._TP is original_tp_group + assert srt_parallel_state._TP is original_tp_group + assert srt_parallel_state._ATTN_TP is original_tp_group + assert get_parallel().tp_size == 1 + assert get_parallel().tp_rank == 0 + assert get_parallel().tp_group is original_tp_group + assert get_parallel().attn_tp_size == 1 + assert get_parallel().attn_tp_rank == 0 + assert get_parallel().attn_tp_group is original_tp_group + assert get_parallel().moe_tp_rank == 0 def test_encoder_folding_context_is_nested_and_restores_each_group(): - original_tp_group = object() + original_tp_group = _tp_group() outer_tp_group = _tp_group(world_size=4, rank_in_group=3) inner_tp_group = _tp_group(world_size=2, rank_in_group=1) with ( patch.object(parallel_state, "_TP", original_tp_group), - patch.object(srt_parallel_state, "_TP", original_tp_group), - patch.object(srt_parallel_state, "_ATTN_TP", original_tp_group), + patch.object(srt_parallel_state, "_TP", None), + patch.object(srt_parallel_state, "_ATTN_TP", None), + patch("sglang.srt.runtime_context._PARALLEL", ParallelContext()), ): + parallel_state._sync_srt_tp_group() with parallel_state.use_tensor_parallel_group(outer_tp_group): assert get_parallel().tp_size == 4 - with parallel_state.use_tensor_parallel_group(inner_tp_group): - assert parallel_state._TP is inner_tp_group - assert srt_parallel_state._TP is inner_tp_group - assert srt_parallel_state._ATTN_TP is inner_tp_group - assert get_parallel().tp_size == 2 - assert get_parallel().tp_rank == 1 + with pytest.raises(RuntimeError, match="encoder load failed"): + with parallel_state.use_tensor_parallel_group(inner_tp_group): + assert parallel_state._TP is inner_tp_group + assert srt_parallel_state._TP is inner_tp_group + assert srt_parallel_state._ATTN_TP is inner_tp_group + assert get_parallel().tp_size == 2 + assert get_parallel().tp_rank == 1 + assert get_parallel().attn_tp_group is inner_tp_group + assert get_parallel().attn_tp_rank == 1 + assert get_parallel().moe_tp_rank == 1 + raise RuntimeError("encoder load failed") assert parallel_state._TP is outer_tp_group assert srt_parallel_state._TP is outer_tp_group assert srt_parallel_state._ATTN_TP is outer_tp_group assert get_parallel().tp_size == 4 assert get_parallel().tp_rank == 3 + assert get_parallel().attn_tp_group is outer_tp_group + assert get_parallel().attn_tp_rank == 3 + assert get_parallel().moe_tp_rank == 3 assert parallel_state._TP is original_tp_group assert srt_parallel_state._TP is original_tp_group assert srt_parallel_state._ATTN_TP is original_tp_group + assert get_parallel().tp_group is original_tp_group + assert get_parallel().tp_rank == 0 + assert get_parallel().attn_tp_group is original_tp_group + assert get_parallel().attn_tp_rank == 0 + assert get_parallel().moe_tp_rank == 0 def test_weight_transfer_uses_loader_for_implicit_srt_shard():