[Fix] Keep diffusion encoder TP context bindings consistent (#40646)

This commit is contained in:
Cheng Wan
2026-09-21 17:23:15 -07:00
committed by GitHub
parent c53cc8e1eb
commit 582389cec5
2 changed files with 54 additions and 26 deletions
@@ -699,6 +699,9 @@ def use_tensor_parallel_group(tp_group: GroupCoordinator):
tp_size=tp_group.world_size, tp_size=tp_group.world_size,
tp_rank=tp_group.rank_in_group, tp_rank=tp_group.rank_in_group,
tp_group=tp_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 # Only tensor parallelism folds here, so every other dimension is
# one and the quotients come out of the shared derivation. # one and the quotients come out of the shared derivation.
**derive_parallel_widths( **derive_parallel_widths(
@@ -18,7 +18,7 @@ from sglang.multimodal_gen.test.single_test_file.component_accuracy.utils import
initialize_parallel_runtime, initialize_parallel_runtime,
) )
from sglang.srt.distributed import parallel_state as srt_parallel_state 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" _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 assert srt_parallel_state._ATTN_TP is srt_attention_tp_group
def test_srt_tp_groups_follow_encoder_folding_context(): @pytest.mark.parametrize("rank", [0, 1])
original_diffusion_tp_group = object() def test_srt_tp_groups_follow_encoder_folding_context(rank):
original_srt_tp_group = object() original_tp_group = _tp_group()
original_srt_attention_tp_group = object() folding_tp_group = _tp_group(world_size=2, rank_in_group=rank)
folding_tp_group = _tp_group(world_size=2, rank_in_group=1)
with ( with (
patch.object(parallel_state, "_TP", original_diffusion_tp_group), patch.object(parallel_state, "_TP", original_tp_group),
patch.object(srt_parallel_state, "_TP", original_srt_tp_group), patch.object(srt_parallel_state, "_TP", None),
patch.object( patch.object(srt_parallel_state, "_ATTN_TP", None),
srt_parallel_state, patch("sglang.srt.runtime_context._PARALLEL", ParallelContext()),
"_ATTN_TP",
original_srt_attention_tp_group,
),
): ):
# 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): with parallel_state.use_tensor_parallel_group(folding_tp_group):
assert parallel_state._TP is folding_tp_group assert parallel_state._TP is folding_tp_group
assert srt_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 srt_parallel_state._ATTN_TP is folding_tp_group
assert get_parallel().tp_size == 2 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().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 parallel_state._TP is original_tp_group
assert srt_parallel_state._TP is original_srt_tp_group assert srt_parallel_state._TP is original_tp_group
assert srt_parallel_state._ATTN_TP is original_srt_attention_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(): 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) outer_tp_group = _tp_group(world_size=4, rank_in_group=3)
inner_tp_group = _tp_group(world_size=2, rank_in_group=1) inner_tp_group = _tp_group(world_size=2, rank_in_group=1)
with ( with (
patch.object(parallel_state, "_TP", original_tp_group), patch.object(parallel_state, "_TP", original_tp_group),
patch.object(srt_parallel_state, "_TP", original_tp_group), patch.object(srt_parallel_state, "_TP", None),
patch.object(srt_parallel_state, "_ATTN_TP", original_tp_group), 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): with parallel_state.use_tensor_parallel_group(outer_tp_group):
assert get_parallel().tp_size == 4 assert get_parallel().tp_size == 4
with parallel_state.use_tensor_parallel_group(inner_tp_group): with pytest.raises(RuntimeError, match="encoder load failed"):
assert parallel_state._TP is inner_tp_group with parallel_state.use_tensor_parallel_group(inner_tp_group):
assert srt_parallel_state._TP is inner_tp_group assert parallel_state._TP is inner_tp_group
assert srt_parallel_state._ATTN_TP is inner_tp_group assert srt_parallel_state._TP is inner_tp_group
assert get_parallel().tp_size == 2 assert srt_parallel_state._ATTN_TP is inner_tp_group
assert get_parallel().tp_rank == 1 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 parallel_state._TP is outer_tp_group
assert srt_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 srt_parallel_state._ATTN_TP is outer_tp_group
assert get_parallel().tp_size == 4 assert get_parallel().tp_size == 4
assert get_parallel().tp_rank == 3 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 parallel_state._TP is original_tp_group
assert srt_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 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(): def test_weight_transfer_uses_loader_for_implicit_srt_shard():