[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_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(
@@ -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():