[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:
Yuzhen Zhou
2025-12-19 15:38:46 +08:00
committed by GitHub
co-authored by Sabre Shao Yusheng Su Hubert Lu xsun
parent f2d64e6782
commit 4bf06635fc
34 changed files with 823 additions and 72 deletions
@@ -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"),