Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 18 additions & 13 deletions diffsynth_engine/distributed/parallel_state.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -13,7 +13,6 @@

import torch
import torch.distributed
from torch.cuda import device_count, set_device

from diffsynth_engine.distributed.group_coordinator import (
GroupCoordinator,
Expand All@@ -22,7 +21,11 @@
)
from diffsynth_engine.utils import logging
from diffsynth_engine.utils.constants import IDLE_TIMEOUT_SEC
from diffsynth_engine.utils.platform import get_torch_distributed_backend
from diffsynth_engine.utils.platform import (
device_count,
get_torch_distributed_backend,
set_device,
)

logger = logging.get_logger(__name__)

Expand DownExpand Up@@ -425,28 +428,30 @@ def init_distributed_environment(
distributed_init_method,
backend,
)
# local_rank is not available in torch ProcessGroup,
# see https://github.com/pytorch/pytorch/issues/122816
if local_rank == -1:
# local rank not set, this usually happens in single-node
# setting, where we can use rank as local rank
if distributed_init_method == "env://":
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
else:
local_rank = rank if rank >= 0 else 0

if not torch.distributed.is_initialized():
assert distributed_init_method is not None, (
"distributed_init_method must be provided when initializing distributed environment"
)
# Bind device before init_process_group (required by HCCL on Ascend).
set_device(local_rank % max(device_count(), 1))
# this backend is used for WORLD
torch.distributed.init_process_group(
backend=backend,
init_method=distributed_init_method,
world_size=world_size,
rank=rank,
)
set_device(torch.distributed.get_rank() % device_count())
# set the local rank
# local_rank is not available in torch ProcessGroup,
# see https://github.com/pytorch/pytorch/issues/122816
if local_rank == -1:
# local rank not set, this usually happens in single-node
# setting, where we can use rank as local rank
if distributed_init_method == "env://":
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
else:
local_rank = rank
set_device(torch.distributed.get_rank() % max(device_count(), 1))
global _WORLD
if _WORLD is None:
ranks = list(range(torch.distributed.get_world_size()))
Expand Down
3 changes: 2 additions & 1 deletion diffsynth_engine/engine.py
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,14 @@
from typing import Any

import torch.multiprocessing as mp
from torch.cuda import set_device

from diffsynth_engine.configs import PipelineConfig
from diffsynth_engine.registry import (
get_pipeline_class,
get_pipeline_class_name,
)
from diffsynth_engine.utils import logging
from diffsynth_engine.utils.platform import align_config_device, set_device
from diffsynth_engine.utils.torch_profiler import TorchProfiler
from diffsynth_engine.worker import run_worker_loop

Expand All@@ -19,6 +19,7 @@ class DiffSynthEngine:
@classmethod
def from_pretrained(cls, model_path_or_config: str | PipelineConfig, **kwargs):
pipeline_config = _resolve_pipeline_config(model_path_or_config)
pipeline_config.device = align_config_device(pipeline_config.device)
num_workers = pipeline_config.parallelism
master_addr = kwargs.get("master_addr", "localhost")
master_port = kwargs.get("master_port", 29500)
Expand Down
3 changes: 2 additions & 1 deletion diffsynth_engine/layers/attention/__init__.py
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,10 @@
from .backends.abstract import AttentionMetadata, AttentionType
from .layer import LocalAttention, USPAttention
from .layer import LocalAttention, USPAttention, AscendLongContextAttention

__all__ = [
"AttentionType",
"AttentionMetadata",
"LocalAttention",
"USPAttention",
"AscendLongContextAttention",
]
1 change: 1 addition & 0 deletions diffsynth_engine/layers/attention/backends/abstract.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -26,6 +26,7 @@ class AttentionType(str, enum.Enum):
SAGE2 = "sage2"
SAGE3 = "sage3"
SPARGE = "sparge"
MINDIE = "mindie"

def __str__(self) -> str:
return self.value
Expand Down
106 changes: 106 additions & 0 deletions diffsynth_engine/layers/attention/backends/mindie_attn.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,106 @@
import torch

from diffsynth_engine.layers.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
AttentionType,
)
from diffsynth_engine.utils import logging

logger = logging.get_logger(__name__)


class MindieAttentionMetadataBuilder(AttentionMetadataBuilder):
def __init__(self) -> None:
pass

def build(self, **kwargs) -> AttentionMetadata:
return AttentionMetadata()


class MindieAttentionBackend(AttentionBackend):
@staticmethod
def check_availability() -> None:
from diffsynth_engine.platforms import AscendPlatform

if not AscendPlatform.supports("device"):
error_msg = "MindIE attention requires an available Ascend NPU device."
logger.error(error_msg)
raise RuntimeError(error_msg)
if not AscendPlatform.supports("mindie_attention"):
error_msg = (
"MindIE attention backend is not available. "
"Install MindIE-SD 3.x matching the current torch_npu and CANN versions, "
"and ensure mindiesd.layers.flash_attn.attention_forward works on NPU."
)
logger.error(error_msg)
raise RuntimeError(error_msg)

@staticmethod
def get_type() -> str:
return str(AttentionType.MINDIE)

@staticmethod
def get_impl_cls() -> type["AttentionImpl"]:
return MindieAttentionImpl

@staticmethod
def get_metadata_cls() -> type["AttentionMetadata"]:
return AttentionMetadata

@staticmethod
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
return MindieAttentionMetadataBuilder

@staticmethod
def get_supported_head_sizes() -> list[int]:
return []

@classmethod
def supports_ring_attention(cls) -> bool:
return False


class MindieAttentionImpl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
softmax_scale: float | None = None,
causal: bool = False,
num_kv_heads: int | None = None,
**extra_impl_args,
) -> None:
if num_kv_heads is None:
num_kv_heads = num_heads
self.num_kv_groups = num_heads // num_kv_heads
self.causal = causal
self.softmax_scale = softmax_scale
self.num_heads = num_heads
self.head_size = head_size

def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_mask: torch.Tensor | None = None,
attn_metadata: AttentionMetadata | None = None,
**kwargs,
) -> torch.Tensor:
from mindiesd.layers.flash_attn.attention_forward import attention_forward

return attention_forward(
query=query,
key=key,
value=value,
attn_mask=attn_mask,
scale=self.softmax_scale,
fused=True,
head_first=False,
opt_mode="manual",
op_type="fused_attn_score",
layout="BSND",
)
Loading