[diffusion] multi-platform: support diffusion on amd and fix encoder loading on MI325 (#13760)
Co-authored-by: Sabre Shao <sabre.shao@amd.com> Co-authored-by: Yusheng (Ethan) Su <yushengsu.thu@gmail.com> Co-authored-by: Hubert Lu <Hubert.Lu@amd.com> Co-authored-by: xsun <sunxiao04@gmail.com>
This commit is contained in:
co-authored by
Sabre Shao
Yusheng Su
Hubert Lu
xsun
parent
f2d64e6782
commit
4bf06635fc
@@ -404,15 +404,25 @@ def initialize_model_parallel(
|
||||
global _SP
|
||||
assert _SP is None, "sequence parallel group is already initialized"
|
||||
|
||||
from yunchang import set_seq_parallel_pg
|
||||
from yunchang.globals import PROCESS_GROUP
|
||||
try:
|
||||
from .yunchang import PROCESS_GROUP as _YC_PROCESS_GROUP
|
||||
from .yunchang import set_seq_parallel_pg as _set_seq_parallel_pg
|
||||
except ImportError:
|
||||
_set_seq_parallel_pg = None
|
||||
|
||||
set_seq_parallel_pg(
|
||||
sp_ulysses_degree=ulysses_degree,
|
||||
sp_ring_degree=ring_degree,
|
||||
rank=get_world_group().rank_in_group,
|
||||
world_size=dit_parallel_size,
|
||||
)
|
||||
class _DummyProcessGroup:
|
||||
ULYSSES_PG = torch.distributed.group.WORLD
|
||||
RING_PG = torch.distributed.group.WORLD
|
||||
|
||||
PROCESS_GROUP = _DummyProcessGroup()
|
||||
else:
|
||||
_set_seq_parallel_pg(
|
||||
sp_ulysses_degree=ulysses_degree,
|
||||
sp_ring_degree=ring_degree,
|
||||
rank=get_world_group().rank_in_group,
|
||||
world_size=dit_parallel_size,
|
||||
)
|
||||
PROCESS_GROUP = _YC_PROCESS_GROUP
|
||||
|
||||
_SP = init_parallel_group_coordinator(
|
||||
group_ranks=rank_generator.get_ranks("sp"),
|
||||
|
||||
Reference in New Issue
Block a user