From c64681f162e4acce1c8b0116c728d31d7c1b6b4a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=90=D1=80=D1=82=D0=B5=D0=BC=20=D0=A1=D0=B0=D0=B2=D0=BA?= =?UTF-8?q?=D0=B8=D0=BD?= <58187114+OrangeRedeng@users.noreply.github.com> Date: Wed, 18 Mar 2026 09:05:12 +0300 Subject: [PATCH] [Bugfix] [diffusion] Fix cache-dit with sp-degree only (#19965) Co-authored-by: Mick Co-authored-by: ronnie_zheng --- .../runtime/cache/cache_dit_integration.py | 13 +++++++++---- 1 file changed, 9 insertions(+), 4 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py b/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py index e08124881..bf0172e6d 100644 --- a/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py +++ b/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py @@ -12,6 +12,11 @@ from typing import List, Optional import torch import torch.distributed as dist +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_ring_parallel_world_size, + get_tp_world_size, + get_ulysses_parallel_world_size, +) from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) @@ -107,15 +112,15 @@ def _build_parallelism_config( ulysses_size = None ring_size = None if sp_group is not None: - ulysses_size = getattr(sp_group, "ulysses_world_size", None) - ring_size = getattr(sp_group, "ring_world_size", None) + ulysses_size = get_ulysses_parallel_world_size() + ring_size = get_ring_parallel_world_size() tp_size = None if tp_group is not None: - tp_size = dist.get_world_size(tp_group) + tp_size = get_tp_world_size() return ParallelismConfig( - backend=ParallelismBackend.NATIVE_PYTORCH, + backend=ParallelismBackend.AUTO, ulysses_size=ulysses_size, ring_size=ring_size, tp_size=tp_size,