[Bugfix] [diffusion] Fix cache-dit with sp-degree only (#19965)
Co-authored-by: Mick <mickjagger19@icloud.com> Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user