[Fix] Keep diffusion encoder TP context bindings consistent (#40646)
This commit is contained in:
@@ -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(
|
||||||
|
|||||||
+51
-26
@@ -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():
|
||||||
|
|||||||
Reference in New Issue
Block a user