[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:
co-authored by
Mick
ronnie_zheng
parent
b6055e59cd
commit
c64681f162
@@ -12,6 +12,11 @@ from typing import List, Optional
|
|||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
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
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
@@ -107,15 +112,15 @@ def _build_parallelism_config(
|
|||||||
ulysses_size = None
|
ulysses_size = None
|
||||||
ring_size = None
|
ring_size = None
|
||||||
if sp_group is not None:
|
if sp_group is not None:
|
||||||
ulysses_size = getattr(sp_group, "ulysses_world_size", None)
|
ulysses_size = get_ulysses_parallel_world_size()
|
||||||
ring_size = getattr(sp_group, "ring_world_size", None)
|
ring_size = get_ring_parallel_world_size()
|
||||||
|
|
||||||
tp_size = None
|
tp_size = None
|
||||||
if tp_group is not None:
|
if tp_group is not None:
|
||||||
tp_size = dist.get_world_size(tp_group)
|
tp_size = get_tp_world_size()
|
||||||
|
|
||||||
return ParallelismConfig(
|
return ParallelismConfig(
|
||||||
backend=ParallelismBackend.NATIVE_PYTORCH,
|
backend=ParallelismBackend.AUTO,
|
||||||
ulysses_size=ulysses_size,
|
ulysses_size=ulysses_size,
|
||||||
ring_size=ring_size,
|
ring_size=ring_size,
|
||||||
tp_size=tp_size,
|
tp_size=tp_size,
|
||||||
|
|||||||
Reference in New Issue
Block a user