[diffusion] platform: support WAN/FLUX/Qwen-Image/Qwen-Image-edit on Ascend (#13662)
Co-authored-by: dhx98 <haox.dai@gmail.com> Co-authored-by: DHX98 <haoxiand@andrew.cmu.edu> Co-authored-by: ronnie_zheng <zl19940307@163.com> Co-authored-by: DHX98 <DHX98@noreply.gitcode.com> Co-authored-by: Yuhao Yang <47235274+yhyang201@users.noreply.github.com>
This commit is contained in:
co-authored by
dhx98
DHX98
ronnie_zheng
DHX98
Yuhao Yang
parent
7b83659310
commit
00248d85c7
@@ -64,7 +64,7 @@ jobs:
|
|||||||
multimodal_gen:
|
multimodal_gen:
|
||||||
- "python/sglang/multimodal_gen/**"
|
- "python/sglang/multimodal_gen/**"
|
||||||
- "python/pyproject_npu.toml"
|
- "python/pyproject_npu.toml"
|
||||||
- "scripts/ci/npu_ci_install_dependency.sh"
|
- "scripts/ci/npu/npu_ci_install_dependency.sh"
|
||||||
- ".github/workflows/pr-test-npu.yml"
|
- ".github/workflows/pr-test-npu.yml"
|
||||||
|
|
||||||
# ==================== PR Gate ==================== #
|
# ==================== PR Gate ==================== #
|
||||||
@@ -241,3 +241,42 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
cd test/srt
|
cd test/srt
|
||||||
python3 run_suite.py --suite per-commit-16-npu-a3 --timeout-per-file 3600 --auto-partition-id ${{ matrix.part }} --auto-partition-size 2
|
python3 run_suite.py --suite per-commit-16-npu-a3 --timeout-per-file 3600 --auto-partition-id ${{ matrix.part }} --auto-partition-size 2
|
||||||
|
|
||||||
|
multimodal-gen-test-1-npu-a3:
|
||||||
|
needs: [check-changes, pr-gate]
|
||||||
|
if: needs.check-changes.outputs.multimodal_gen == 'true'
|
||||||
|
runs-on: linux-aarch64-a3-16
|
||||||
|
container:
|
||||||
|
image: swr.cn-southwest-2.myhuaweicloud.com/base_image/ascend-ci/cann:8.3.rc2-a3-ubuntu22.04-py3.11
|
||||||
|
steps:
|
||||||
|
- name: Checkout code
|
||||||
|
uses: actions/checkout@v4
|
||||||
|
|
||||||
|
- name: Install dependencies
|
||||||
|
run: |
|
||||||
|
# speed up by using infra cache services
|
||||||
|
CACHING_URL="cache-service.nginx-pypi-cache.svc.cluster.local"
|
||||||
|
sed -Ei "s@(ports|archive).ubuntu.com@${CACHING_URL}:8081@g" /etc/apt/sources.list
|
||||||
|
pip config set global.index-url http://${CACHING_URL}/pypi/simple
|
||||||
|
pip config set global.extra-index-url "https://pypi.tuna.tsinghua.edu.cn/simple"
|
||||||
|
pip config set global.trusted-host "${CACHING_URL} pypi.tuna.tsinghua.edu.cn"
|
||||||
|
|
||||||
|
bash scripts/ci/npu/npu_ci_install_dependency.sh a3
|
||||||
|
# copy required file from our daily cache
|
||||||
|
cp ~/.cache/modelscope/hub/datasets/otavia/ShareGPT_Vicuna_unfiltered/ShareGPT_V3_unfiltered_cleaned_split.json /tmp
|
||||||
|
# copy download through proxy
|
||||||
|
curl -o /tmp/test.jsonl -L https://gh-proxy.test.osinfra.cn/https://raw.githubusercontent.com/openai/grade-school-math/master/grade_school_math/data/test.jsonl
|
||||||
|
|
||||||
|
- name: Run test
|
||||||
|
timeout-minutes: 60
|
||||||
|
env:
|
||||||
|
SGLANG_USE_MODELSCOPE: true
|
||||||
|
SGLANG_IS_IN_CI: true
|
||||||
|
HF_ENDPOINT: https://hf-mirror.com
|
||||||
|
TORCH_EXTENSIONS_DIR: /tmp/torch_extensions
|
||||||
|
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
|
||||||
|
STREAMS_PER_DEVICE: 32
|
||||||
|
run: |
|
||||||
|
export PATH="/usr/local/Ascend/8.3.RC1/compiler/bishengir/bin:${PATH}"
|
||||||
|
cd python
|
||||||
|
python3 sglang/multimodal_gen/test/run_suite.py --suite 1-npu
|
||||||
|
|||||||
@@ -77,7 +77,8 @@ diffusion = [
|
|||||||
"moviepy>=2.0.0",
|
"moviepy>=2.0.0",
|
||||||
"opencv-python==4.10.0.84",
|
"opencv-python==4.10.0.84",
|
||||||
"remote-pdb",
|
"remote-pdb",
|
||||||
"cache-dit==1.1.8"
|
"cache-dit==1.2.1",
|
||||||
|
"addict"
|
||||||
]
|
]
|
||||||
|
|
||||||
tracing = [
|
tracing = [
|
||||||
|
|||||||
@@ -16,7 +16,6 @@ import torch.distributed
|
|||||||
from torch.cuda import synchronize
|
from torch.cuda import synchronize
|
||||||
from torch.distributed import Backend, ProcessGroup
|
from torch.distributed import Backend, ProcessGroup
|
||||||
|
|
||||||
from sglang.multimodal_gen import envs
|
|
||||||
from sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_communicator import (
|
from sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_communicator import (
|
||||||
DeviceCommunicatorBase,
|
DeviceCommunicatorBase,
|
||||||
)
|
)
|
||||||
@@ -46,11 +45,7 @@ _group_name_counter: dict[str, int] = {}
|
|||||||
def get_local_torch_device() -> torch.device:
|
def get_local_torch_device() -> torch.device:
|
||||||
"""Return the torch device for the current rank."""
|
"""Return the torch device for the current rank."""
|
||||||
|
|
||||||
return (
|
return current_platform.get_local_torch_device()
|
||||||
torch.device(f"cuda:{envs.LOCAL_RANK}")
|
|
||||||
if current_platform.is_cuda_alike()
|
|
||||||
else torch.device("mps")
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _get_unique_name(name: str) -> str:
|
def _get_unique_name(name: str) -> str:
|
||||||
@@ -190,8 +185,6 @@ class GroupCoordinator:
|
|||||||
# TODO: fix it for other platforms
|
# TODO: fix it for other platforms
|
||||||
self.device = get_local_torch_device()
|
self.device = get_local_torch_device()
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
|
||||||
|
|
||||||
self.use_device_communicator = use_device_communicator
|
self.use_device_communicator = use_device_communicator
|
||||||
|
|
||||||
self.device_communicator: DeviceCommunicatorBase = None # type: ignore
|
self.device_communicator: DeviceCommunicatorBase = None # type: ignore
|
||||||
@@ -287,9 +280,6 @@ class GroupCoordinator:
|
|||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def graph_capture(self, graph_capture_context: GraphCaptureContext | None = None):
|
def graph_capture(self, graph_capture_context: GraphCaptureContext | None = None):
|
||||||
# Platform-aware graph capture
|
|
||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
|
||||||
|
|
||||||
if current_platform.is_cuda_alike():
|
if current_platform.is_cuda_alike():
|
||||||
if graph_capture_context is None:
|
if graph_capture_context is None:
|
||||||
stream = torch.cuda.Stream()
|
stream = torch.cuda.Stream()
|
||||||
|
|||||||
@@ -248,7 +248,11 @@ def init_distributed_environment(
|
|||||||
# For MPS and MUSA, don't pass device_id as it doesn't support device indices
|
# For MPS and MUSA, don't pass device_id as it doesn't support device indices
|
||||||
extra_args = (
|
extra_args = (
|
||||||
{}
|
{}
|
||||||
if (current_platform.is_mps() or current_platform.is_musa())
|
if (
|
||||||
|
current_platform.is_mps()
|
||||||
|
or current_platform.is_musa()
|
||||||
|
or current_platform.is_npu()
|
||||||
|
)
|
||||||
else dict(device_id=device_id)
|
else dict(device_id=device_id)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -618,6 +622,7 @@ def maybe_init_distributed_environment_and_model_parallel(
|
|||||||
local_rank=local_rank,
|
local_rank=local_rank,
|
||||||
distributed_init_method=distributed_init_method,
|
distributed_init_method=distributed_init_method,
|
||||||
device_id=device,
|
device_id=device,
|
||||||
|
backend=current_platform.get_torch_distributed_backend_str(),
|
||||||
timeout=dist_timeout,
|
timeout=dist_timeout,
|
||||||
)
|
)
|
||||||
initialize_model_parallel(
|
initialize_model_parallel(
|
||||||
|
|||||||
@@ -14,8 +14,12 @@ from sglang.multimodal_gen.runtime.platforms import current_platform
|
|||||||
|
|
||||||
_is_cuda = current_platform.is_cuda()
|
_is_cuda = current_platform.is_cuda()
|
||||||
_is_hip = current_platform.is_hip()
|
_is_hip = current_platform.is_hip()
|
||||||
|
_is_npu = current_platform.is_npu()
|
||||||
if _is_cuda or _is_hip:
|
if _is_cuda or _is_hip:
|
||||||
from sgl_kernel import silu_and_mul
|
from sgl_kernel import silu_and_mul
|
||||||
|
|
||||||
|
if _is_npu:
|
||||||
|
import torch_npu
|
||||||
# TODO (will): remove this dependency
|
# TODO (will): remove this dependency
|
||||||
from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp
|
from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp
|
||||||
|
|
||||||
@@ -46,6 +50,10 @@ class SiluAndMul(CustomOp):
|
|||||||
d = x.shape[-1] // 2
|
d = x.shape[-1] // 2
|
||||||
return F.silu(x[..., :d]) * x[..., d:]
|
return F.silu(x[..., :d]) * x[..., d:]
|
||||||
|
|
||||||
|
def forward_npu(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
out = torch_npu.npu_swiglu(x)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
@CustomOp.register("gelu_and_mul")
|
@CustomOp.register("gelu_and_mul")
|
||||||
class GeluAndMul(CustomOp):
|
class GeluAndMul(CustomOp):
|
||||||
|
|||||||
@@ -64,6 +64,11 @@ class CustomOp(nn.Module):
|
|||||||
# PyTorch-native implementation.
|
# PyTorch-native implementation.
|
||||||
return self.forward_native(*args, **kwargs)
|
return self.forward_native(*args, **kwargs)
|
||||||
|
|
||||||
|
def forward_npu(self, *args, **kwargs) -> Any:
|
||||||
|
# By default, we assume that NPU ops are compatible with the
|
||||||
|
# PyTorch-native implementation.
|
||||||
|
return self.forward_native(*args, **kwargs)
|
||||||
|
|
||||||
def dispatch_forward(self) -> Callable:
|
def dispatch_forward(self) -> Callable:
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
return self.forward_cuda
|
return self.forward_cuda
|
||||||
|
|||||||
@@ -12,9 +12,13 @@ import torch.nn.functional as F
|
|||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
|
|
||||||
_is_cuda = current_platform.is_cuda()
|
_is_cuda = current_platform.is_cuda()
|
||||||
|
_is_npu = current_platform.is_npu()
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
from sgl_kernel import fused_add_rmsnorm, rmsnorm
|
from sgl_kernel import fused_add_rmsnorm, rmsnorm
|
||||||
|
|
||||||
|
if _is_npu:
|
||||||
|
import torch_npu
|
||||||
|
|
||||||
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm, fused_inplace_qknorm
|
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm, fused_inplace_qknorm
|
||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
get_tensor_model_parallel_rank,
|
get_tensor_model_parallel_rank,
|
||||||
@@ -28,11 +32,8 @@ from sglang.multimodal_gen.runtime.layers.triton_ops import (
|
|||||||
rms_norm_fn,
|
rms_norm_fn,
|
||||||
triton_one_pass_rms_norm,
|
triton_one_pass_rms_norm,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
|
||||||
from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var
|
from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var
|
||||||
|
|
||||||
_is_cuda = current_platform.is_cuda()
|
|
||||||
|
|
||||||
|
|
||||||
# Copied and adapted from sglang
|
# Copied and adapted from sglang
|
||||||
@CustomOp.register("rms_norm")
|
@CustomOp.register("rms_norm")
|
||||||
@@ -141,6 +142,18 @@ class RMSNorm(CustomOp):
|
|||||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
return self.forward_native(x, residual)
|
return self.forward_native(x, residual)
|
||||||
|
|
||||||
|
def forward_npu(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
residual: Optional[torch.Tensor] = None,
|
||||||
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
|
if residual is not None:
|
||||||
|
out, _, residual_out = torch_npu.npu_add_rms_norm(
|
||||||
|
residual, x, self.weight.data, self.variance_epsilon
|
||||||
|
)
|
||||||
|
return out, residual_out
|
||||||
|
return torch_npu.npu_rms_norm(x, self.weight.data, self.variance_epsilon)[0]
|
||||||
|
|
||||||
def forward_hip(
|
def forward_hip(
|
||||||
self,
|
self,
|
||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
@@ -214,7 +227,7 @@ class LayerNorm(CustomOp):
|
|||||||
x = x.view(-1, self.hidden_size)
|
x = x.view(-1, self.hidden_size)
|
||||||
return self.forward_triton(x).view(shape)
|
return self.forward_triton(x).view(shape)
|
||||||
|
|
||||||
@torch.compile(backend="inductor")
|
@torch.compile(backend="inductor", disable=current_platform.is_npu())
|
||||||
def forward_native(
|
def forward_native(
|
||||||
self,
|
self,
|
||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
|
|||||||
@@ -35,6 +35,7 @@ from sglang.multimodal_gen.runtime.models.parameter import (
|
|||||||
|
|
||||||
# yapf: enable
|
# yapf: enable
|
||||||
from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs
|
from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs
|
||||||
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
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__)
|
||||||
@@ -152,7 +153,7 @@ class UnquantizedLinearMethod(LinearMethodBase):
|
|||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
output = (
|
output = (
|
||||||
F.linear(x, layer.weight, bias)
|
F.linear(x, layer.weight, bias)
|
||||||
if torch.cuda.is_available() or bias is None
|
if current_platform.is_amp_supported() or bias is None
|
||||||
else F.linear(x, layer.weight, bias.to(x.dtype))
|
else F.linear(x, layer.weight, bias.to(x.dtype))
|
||||||
) # NOTE: this line assumes that we are using amp when using cuda and is needed to account for the fact that amp isn't supported in mps
|
) # NOTE: this line assumes that we are using amp when using cuda and is needed to account for the fact that amp isn't supported in mps
|
||||||
return output
|
return output
|
||||||
|
|||||||
@@ -8,6 +8,8 @@ import triton # type: ignore
|
|||||||
import triton.language as tl # type: ignore
|
import triton.language as tl # type: ignore
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
|
|
||||||
|
|
||||||
@triton.autotune(
|
@triton.autotune(
|
||||||
configs=[
|
configs=[
|
||||||
@@ -524,8 +526,14 @@ def triton_autotune_configs():
|
|||||||
max_threads_per_block = 1024
|
max_threads_per_block = 1024
|
||||||
# Default to warp size 32 if not defined by device
|
# Default to warp size 32 if not defined by device
|
||||||
warp_size = getattr(
|
warp_size = getattr(
|
||||||
torch.cuda.get_device_properties(torch.cuda.current_device()), "warp_size", 32
|
torch.get_device_module().get_device_properties(
|
||||||
|
torch.get_device_module().current_device()
|
||||||
|
),
|
||||||
|
"warp_size",
|
||||||
|
32,
|
||||||
)
|
)
|
||||||
|
if warp_size is None:
|
||||||
|
warp_size = 32
|
||||||
# Autotune for warp counts which are powers of 2 and do not exceed thread per block limit
|
# Autotune for warp counts which are powers of 2 and do not exceed thread per block limit
|
||||||
return [
|
return [
|
||||||
triton.Config({}, num_warps=warp_count)
|
triton.Config({}, num_warps=warp_count)
|
||||||
@@ -820,7 +828,7 @@ def _layer_norm_fwd_impl(
|
|||||||
BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(N))
|
BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(N))
|
||||||
if N > BLOCK_N:
|
if N > BLOCK_N:
|
||||||
raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
|
raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
|
||||||
with torch.cuda.device(x.device.index):
|
with torch.get_device_module().device(x.device.index):
|
||||||
torch.library.wrap_triton(_layer_norm_fwd_1pass_kernel)[(M,)](
|
torch.library.wrap_triton(_layer_norm_fwd_1pass_kernel)[(M,)](
|
||||||
x,
|
x,
|
||||||
out,
|
out,
|
||||||
@@ -1166,3 +1174,31 @@ def triton_one_pass_rms_norm(x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6
|
|||||||
BLOCK_SIZE_SEQ=BLOCK_SIZE_SEQ,
|
BLOCK_SIZE_SEQ=BLOCK_SIZE_SEQ,
|
||||||
)
|
)
|
||||||
return y
|
return y
|
||||||
|
|
||||||
|
|
||||||
|
if current_platform.is_npu():
|
||||||
|
# TODO: remove this when triton ascend bug is fixed
|
||||||
|
def fuse_scale_shift_native(
|
||||||
|
x: torch.Tensor,
|
||||||
|
scale: torch.Tensor,
|
||||||
|
shift: torch.Tensor,
|
||||||
|
block_l: int = 128,
|
||||||
|
block_c: int = 128,
|
||||||
|
):
|
||||||
|
return x * (1 + scale) + shift
|
||||||
|
|
||||||
|
fuse_scale_shift_kernel = fuse_scale_shift_native
|
||||||
|
|
||||||
|
# TODO: remove this when triton ascend bug is fixed
|
||||||
|
def apply_rotary_embedding_native(
|
||||||
|
x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, interleaved: bool = False
|
||||||
|
) -> torch.Tensor:
|
||||||
|
cos = cos.unsqueeze(-2).to(x.dtype)
|
||||||
|
sin = sin.unsqueeze(-2).to(x.dtype)
|
||||||
|
x1 = x[..., ::2]
|
||||||
|
x2 = x[..., 1::2]
|
||||||
|
o1 = x1 * cos - x2 * sin
|
||||||
|
o2 = x2 * cos + x1 * sin
|
||||||
|
return torch.stack((o1, o2), dim=-1).flatten(-2)
|
||||||
|
|
||||||
|
apply_rotary_embedding = apply_rotary_embedding_native
|
||||||
|
|||||||
@@ -145,7 +145,11 @@ class VocabParallelEmbeddingShardIndices:
|
|||||||
assert self.num_added_elements <= self.num_added_elements_padded
|
assert self.num_added_elements <= self.num_added_elements_padded
|
||||||
|
|
||||||
|
|
||||||
@torch.compile(dynamic=True, backend=current_platform.simple_compile_backend)
|
@torch.compile(
|
||||||
|
dynamic=True,
|
||||||
|
backend=current_platform.simple_compile_backend,
|
||||||
|
disable=current_platform.is_npu(),
|
||||||
|
)
|
||||||
def get_masked_input_and_mask(
|
def get_masked_input_and_mask(
|
||||||
input_: torch.Tensor,
|
input_: torch.Tensor,
|
||||||
org_vocab_start_index: int,
|
org_vocab_start_index: int,
|
||||||
|
|||||||
@@ -71,7 +71,7 @@ class GPUWorker:
|
|||||||
def init_device_and_model(self) -> None:
|
def init_device_and_model(self) -> None:
|
||||||
"""Initialize the device and load the model."""
|
"""Initialize the device and load the model."""
|
||||||
setproctitle(f"sgl_diffusion::scheduler_TP{self.local_rank}")
|
setproctitle(f"sgl_diffusion::scheduler_TP{self.local_rank}")
|
||||||
torch.cuda.set_device(self.local_rank)
|
torch.get_device_module().set_device(self.local_rank)
|
||||||
# Set environment variables for distributed initialization
|
# Set environment variables for distributed initialization
|
||||||
os.environ["MASTER_ADDR"] = "localhost"
|
os.environ["MASTER_ADDR"] = "localhost"
|
||||||
os.environ["MASTER_PORT"] = str(self.master_port)
|
os.environ["MASTER_PORT"] = str(self.master_port)
|
||||||
@@ -86,6 +86,7 @@ class GPUWorker:
|
|||||||
ring_degree=self.server_args.ring_degree,
|
ring_degree=self.server_args.ring_degree,
|
||||||
sp_size=self.server_args.sp_degree,
|
sp_size=self.server_args.sp_degree,
|
||||||
dp_size=self.server_args.dp_size,
|
dp_size=self.server_args.dp_size,
|
||||||
|
distributed_init_method=f"tcp://127.0.0.1:{self.master_port}",
|
||||||
dist_timeout=self.server_args.dist_timeout,
|
dist_timeout=self.server_args.dist_timeout,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -160,7 +161,7 @@ class GPUWorker:
|
|||||||
output_batch = None
|
output_batch = None
|
||||||
try:
|
try:
|
||||||
if self.rank == 0:
|
if self.rank == 0:
|
||||||
torch.cuda.reset_peak_memory_stats()
|
torch.get_device_module().reset_peak_memory_stats()
|
||||||
|
|
||||||
start_time = time.monotonic()
|
start_time = time.monotonic()
|
||||||
|
|
||||||
@@ -347,6 +348,7 @@ def run_scheduler_process(
|
|||||||
"""
|
"""
|
||||||
configure_logger(server_args)
|
configure_logger(server_args)
|
||||||
globally_suppress_loggers()
|
globally_suppress_loggers()
|
||||||
|
if current_platform.is_cuda():
|
||||||
set_cuda_arch()
|
set_cuda_arch()
|
||||||
|
|
||||||
port_args = PortArgs.from_server_args(server_args)
|
port_args = PortArgs.from_server_args(server_args)
|
||||||
|
|||||||
@@ -854,7 +854,7 @@ class WanTransformer3DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
|
|
||||||
encoder_hidden_states = (
|
encoder_hidden_states = (
|
||||||
encoder_hidden_states.to(orig_dtype)
|
encoder_hidden_states.to(orig_dtype)
|
||||||
if current_platform.is_mps()
|
if not current_platform.is_amp_supported()
|
||||||
else encoder_hidden_states
|
else encoder_hidden_states
|
||||||
) # cast to orig_dtype for MPS
|
) # cast to orig_dtype for MPS
|
||||||
|
|
||||||
|
|||||||
@@ -264,7 +264,7 @@ class CLIPAttention(nn.Module):
|
|||||||
key_states,
|
key_states,
|
||||||
value_states,
|
value_states,
|
||||||
attn_mask=attn_mask,
|
attn_mask=attn_mask,
|
||||||
is_causal=True,
|
is_causal=attention_mask is None,
|
||||||
scale=self.scale,
|
scale=self.scale,
|
||||||
)
|
)
|
||||||
attn_output = attn_output.transpose(1, 2)
|
attn_output = attn_output.transpose(1, 2)
|
||||||
|
|||||||
@@ -1227,10 +1227,9 @@ class DenoisingStage(PipelineStage):
|
|||||||
raw_latent_shape=batch.raw_latent_shape
|
raw_latent_shape=batch.raw_latent_shape
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
# attn_metadata can be None for SDPA attention backend
|
||||||
return None
|
return None
|
||||||
|
|
||||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
|
||||||
|
|
||||||
return attn_metadata
|
return attn_metadata
|
||||||
|
|
||||||
def _predict_noise(
|
def _predict_noise(
|
||||||
|
|||||||
@@ -101,6 +101,24 @@ def rocm_platform_plugin() -> str | None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def npu_platform_plugin() -> str | None:
|
||||||
|
is_npu = False
|
||||||
|
|
||||||
|
try:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
if torch.npu.is_available():
|
||||||
|
is_npu = True
|
||||||
|
logger.info("NPU is available")
|
||||||
|
except Exception as e:
|
||||||
|
logger.info("NPU detection failed: %s", e)
|
||||||
|
return (
|
||||||
|
"sglang.multimodal_gen.runtime.platforms.npu.NPUPlatformBase"
|
||||||
|
if is_npu
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def musa_platform_plugin() -> str | None:
|
def musa_platform_plugin() -> str | None:
|
||||||
is_musa = False
|
is_musa = False
|
||||||
|
|
||||||
@@ -125,6 +143,7 @@ builtin_platform_plugins = {
|
|||||||
"rocm": rocm_platform_plugin,
|
"rocm": rocm_platform_plugin,
|
||||||
"mps": mps_platform_plugin,
|
"mps": mps_platform_plugin,
|
||||||
"cpu": cpu_platform_plugin,
|
"cpu": cpu_platform_plugin,
|
||||||
|
"npu": npu_platform_plugin,
|
||||||
"musa": musa_platform_plugin,
|
"musa": musa_platform_plugin,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -148,6 +167,11 @@ def resolve_current_platform_cls_qualname() -> str:
|
|||||||
if platform_cls_qualname is not None:
|
if platform_cls_qualname is not None:
|
||||||
return platform_cls_qualname
|
return platform_cls_qualname
|
||||||
|
|
||||||
|
# Fall back to NPU
|
||||||
|
platform_cls_qualname = npu_platform_plugin()
|
||||||
|
if platform_cls_qualname is not None:
|
||||||
|
return platform_cls_qualname
|
||||||
|
|
||||||
# Fall back to MUSA
|
# Fall back to MUSA
|
||||||
platform_cls_qualname = musa_platform_plugin()
|
platform_cls_qualname = musa_platform_plugin()
|
||||||
if platform_cls_qualname is not None:
|
if platform_cls_qualname is not None:
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ import psutil
|
|||||||
import torch
|
import torch
|
||||||
from typing_extensions import ParamSpec
|
from typing_extensions import ParamSpec
|
||||||
|
|
||||||
|
from sglang.multimodal_gen import envs
|
||||||
from sglang.multimodal_gen.runtime.platforms.interface import (
|
from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||||
AttentionBackendEnum,
|
AttentionBackendEnum,
|
||||||
DeviceCapability,
|
DeviceCapability,
|
||||||
@@ -74,6 +75,10 @@ class CudaPlatformBase(Platform):
|
|||||||
dispatch_key: str = "CUDA"
|
dispatch_key: str = "CUDA"
|
||||||
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
|
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_local_torch_device(cls) -> torch.device:
|
||||||
|
return torch.device(f"cuda:{envs.LOCAL_RANK}")
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability | None:
|
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability | None:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|||||||
@@ -47,6 +47,7 @@ class PlatformEnum(enum.Enum):
|
|||||||
TPU = enum.auto()
|
TPU = enum.auto()
|
||||||
CPU = enum.auto()
|
CPU = enum.auto()
|
||||||
MPS = enum.auto()
|
MPS = enum.auto()
|
||||||
|
NPU = enum.auto()
|
||||||
MUSA = enum.auto()
|
MUSA = enum.auto()
|
||||||
OOT = enum.auto()
|
OOT = enum.auto()
|
||||||
UNSPECIFIED = enum.auto()
|
UNSPECIFIED = enum.auto()
|
||||||
@@ -99,6 +100,10 @@ class Platform:
|
|||||||
def is_cuda(self) -> bool:
|
def is_cuda(self) -> bool:
|
||||||
return self.is_cuda_static()
|
return self.is_cuda_static()
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
|
def is_npu(self) -> bool:
|
||||||
|
return self._enum == PlatformEnum.NPU
|
||||||
|
|
||||||
@lru_cache(maxsize=1)
|
@lru_cache(maxsize=1)
|
||||||
def is_rocm(self) -> bool:
|
def is_rocm(self) -> bool:
|
||||||
return self.is_rocm_static()
|
return self.is_rocm_static()
|
||||||
@@ -175,6 +180,15 @@ class Platform:
|
|||||||
def is_hip(self) -> bool:
|
def is_hip(self) -> bool:
|
||||||
return self.is_rocm()
|
return self.is_rocm()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
|
def is_amp_supported(cls) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_local_torch_device(cls) -> torch.device:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_attn_backend_cls_str(
|
def get_attn_backend_cls_str(
|
||||||
cls,
|
cls,
|
||||||
@@ -236,6 +250,8 @@ class Platform:
|
|||||||
def get_device(self, local_rank: int) -> torch.device:
|
def get_device(self, local_rank: int) -> torch.device:
|
||||||
if self.is_cuda() or self.is_rocm():
|
if self.is_cuda() or self.is_rocm():
|
||||||
return torch.device("cuda", local_rank)
|
return torch.device("cuda", local_rank)
|
||||||
|
elif self.is_npu():
|
||||||
|
return torch.device("npu", local_rank)
|
||||||
elif self.is_musa():
|
elif self.is_musa():
|
||||||
return torch.device("musa", local_rank)
|
return torch.device("musa", local_rank)
|
||||||
elif self.is_mps():
|
elif self.is_mps():
|
||||||
@@ -247,6 +263,8 @@ class Platform:
|
|||||||
def get_torch_distributed_backend_str(self) -> str:
|
def get_torch_distributed_backend_str(self) -> str:
|
||||||
if self.is_cuda_alike():
|
if self.is_cuda_alike():
|
||||||
return "nccl"
|
return "nccl"
|
||||||
|
elif self.is_npu():
|
||||||
|
return "hccl"
|
||||||
elif self.is_musa():
|
elif self.is_musa():
|
||||||
return "mccl"
|
return "mccl"
|
||||||
elif self.is_mps():
|
elif self.is_mps():
|
||||||
|
|||||||
@@ -26,6 +26,15 @@ class MpsPlatform(Platform):
|
|||||||
dispatch_key: str = "MPS"
|
dispatch_key: str = "MPS"
|
||||||
device_control_env_var: str = "MPS_VISIBLE_DEVICES"
|
device_control_env_var: str = "MPS_VISIBLE_DEVICES"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
|
def is_amp_supported(cls) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_local_torch_device(cls) -> torch.device:
|
||||||
|
return torch.device("mps")
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability | None:
|
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability | None:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|||||||
@@ -0,0 +1,126 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
# Adapted from vllm-ascend: https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/platform.py
|
||||||
|
|
||||||
|
import os
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen import envs
|
||||||
|
from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||||
|
AttentionBackendEnum,
|
||||||
|
DeviceCapability,
|
||||||
|
Platform,
|
||||||
|
PlatformEnum,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def device_id_to_physical_device_id(device_id: int) -> int:
|
||||||
|
if "ASCEND_RT_VISIBLE_DEVICES" in os.environ:
|
||||||
|
device_ids = os.environ["ASCEND_RT_VISIBLE_DEVICES"].split(",")
|
||||||
|
if device_ids == [""]:
|
||||||
|
msg = (
|
||||||
|
"ASCEND_RT_VISIBLE_DEVICES is set to empty string, which means"
|
||||||
|
" NPU support is disabled"
|
||||||
|
)
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
physical_device_id = device_ids[device_id]
|
||||||
|
return int(physical_device_id)
|
||||||
|
else:
|
||||||
|
return device_id
|
||||||
|
|
||||||
|
|
||||||
|
class NPUPlatformBase(Platform):
|
||||||
|
_enum = PlatformEnum.NPU
|
||||||
|
device_name: str = "npu"
|
||||||
|
device_type: str = "npu"
|
||||||
|
dispatch_key: str = "NPU"
|
||||||
|
device_control_env_var: str = "ASCEND_RT_VISIBLE_DEVICES"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_local_torch_device(cls) -> torch.device:
|
||||||
|
return torch.device(f"npu:{envs.LOCAL_RANK}")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability:
|
||||||
|
return None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_device_name(cls, device_id: int = 0) -> str:
|
||||||
|
return str(torch.npu.get_device_name(device_id))
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||||
|
device_props = torch.npu.get_device_properties(device_id)
|
||||||
|
return int(device_props.total_memory)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
|
||||||
|
if enforce_eager:
|
||||||
|
logger.warning(
|
||||||
|
"To see benefits of async output processing, enable NPU "
|
||||||
|
"graph. Since, enforce-eager is enabled, async output "
|
||||||
|
"processor cannot be used"
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def is_full_nvlink(cls, physical_device_ids: list[int]) -> bool:
|
||||||
|
logger.exception(
|
||||||
|
"NVLink detection not possible, as context support was"
|
||||||
|
" not found. Assuming no NVLink available."
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_available_gpu_memory(
|
||||||
|
cls,
|
||||||
|
device_id: int = 0,
|
||||||
|
distributed: bool = False,
|
||||||
|
empty_cache: bool = True,
|
||||||
|
cpu_group: Any = None,
|
||||||
|
) -> float:
|
||||||
|
if empty_cache:
|
||||||
|
torch.npu.empty_cache()
|
||||||
|
|
||||||
|
free_gpu_memory, _ = torch.npu.mem_get_info(device_id)
|
||||||
|
|
||||||
|
if distributed:
|
||||||
|
import torch.distributed as dist
|
||||||
|
|
||||||
|
tensor = torch.tensor(free_gpu_memory, dtype=torch.float32, device="npu")
|
||||||
|
dist.all_reduce(tensor, op=dist.ReduceOp.MIN, group=cpu_group)
|
||||||
|
free_gpu_memory = float(tensor.item())
|
||||||
|
|
||||||
|
return free_gpu_memory / (1 << 30)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def log_warnings(cls) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_current_memory_usage(
|
||||||
|
cls, device: torch.types.Device | None = None
|
||||||
|
) -> float:
|
||||||
|
torch.npu.reset_peak_memory_stats(device)
|
||||||
|
return float(torch.npu.max_memory_allocated(device))
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_attn_backend_cls_str(
|
||||||
|
cls,
|
||||||
|
selected_backend: AttentionBackendEnum | None,
|
||||||
|
head_size: int,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
) -> str:
|
||||||
|
logger.info("Using Torch SDPA backend.")
|
||||||
|
return (
|
||||||
|
"sglang.multimodal_gen.runtime.layers.attention.backends.sdpa.SDPABackend"
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_device_communicator_cls(cls) -> str:
|
||||||
|
return "sglang.multimodal_gen.runtime.distributed.device_communicators.cuda_communicator.CudaCommunicator" # noqa
|
||||||
@@ -11,6 +11,7 @@ from typing import Any
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
import sglang.multimodal_gen.envs as envs
|
||||||
from sglang.multimodal_gen.runtime.platforms.interface import (
|
from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||||
AttentionBackendEnum,
|
AttentionBackendEnum,
|
||||||
DeviceCapability,
|
DeviceCapability,
|
||||||
@@ -30,6 +31,10 @@ class RocmPlatform(Platform):
|
|||||||
dispatch_key: str = "CUDA"
|
dispatch_key: str = "CUDA"
|
||||||
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
|
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_local_torch_device(cls) -> torch.device:
|
||||||
|
return torch.device(f"cuda:{envs.LOCAL_RANK}")
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability:
|
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability:
|
||||||
major, minor = torch.cuda.get_device_capability(device_id)
|
major, minor = torch.cuda.get_device_capability(device_id)
|
||||||
|
|||||||
@@ -38,6 +38,15 @@ SUITES = {
|
|||||||
],
|
],
|
||||||
}
|
}
|
||||||
|
|
||||||
|
suites_ascend = {
|
||||||
|
"1-npu": [
|
||||||
|
"ascend/test_server_1_npu.py",
|
||||||
|
# add new 1-npu test files here
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
SUITES.update(suites_ascend)
|
||||||
|
|
||||||
|
|
||||||
def parse_args():
|
def parse_args():
|
||||||
parser = argparse.ArgumentParser(description="Run multimodal_gen test suite")
|
parser = argparse.ArgumentParser(description="Run multimodal_gen test suite")
|
||||||
|
|||||||
@@ -0,0 +1,76 @@
|
|||||||
|
{
|
||||||
|
"metadata": {
|
||||||
|
"model": "Diffusion Server",
|
||||||
|
"hardware": "CI A2 64GB pool",
|
||||||
|
"description": "Reference numbers captured from the CI diffusion server baseline run"
|
||||||
|
},
|
||||||
|
"scenarios": {
|
||||||
|
"wan2_1_t2v_1.3b_1_npu": {
|
||||||
|
"stages_ms": {
|
||||||
|
"InputValidationStage": 0.1,
|
||||||
|
"TextEncodingStage": 1609.27,
|
||||||
|
"ConditioningStage": 0.02,
|
||||||
|
"TimestepPreparationStage": 3.46,
|
||||||
|
"LatentPreparationStage": 0.39,
|
||||||
|
"DenoisingStage": 26324.0,
|
||||||
|
"DecodingStage": 817.68,
|
||||||
|
"per_frame_generation": null
|
||||||
|
},
|
||||||
|
"denoise_step_ms": {
|
||||||
|
"0": 195.27,
|
||||||
|
"1": 329.05,
|
||||||
|
"2": 545.43,
|
||||||
|
"3": 541.3,
|
||||||
|
"4": 537.07,
|
||||||
|
"5": 537.21,
|
||||||
|
"6": 537.19,
|
||||||
|
"7": 537.19,
|
||||||
|
"8": 537.27,
|
||||||
|
"9": 537.05,
|
||||||
|
"10": 537.02,
|
||||||
|
"11": 537.11,
|
||||||
|
"12": 537.42,
|
||||||
|
"13": 537.2,
|
||||||
|
"14": 537.16,
|
||||||
|
"15": 537.11,
|
||||||
|
"16": 537.14,
|
||||||
|
"17": 537.19,
|
||||||
|
"18": 537.1,
|
||||||
|
"19": 537.0,
|
||||||
|
"20": 537.26,
|
||||||
|
"21": 537.18,
|
||||||
|
"22": 537.16,
|
||||||
|
"23": 537.24,
|
||||||
|
"24": 537.15,
|
||||||
|
"25": 537.14,
|
||||||
|
"26": 536.99,
|
||||||
|
"27": 537.19,
|
||||||
|
"28": 537.22,
|
||||||
|
"29": 537.23,
|
||||||
|
"30": 537.06,
|
||||||
|
"31": 537.06,
|
||||||
|
"32": 537.18,
|
||||||
|
"33": 537.07,
|
||||||
|
"34": 537.19,
|
||||||
|
"35": 537.28,
|
||||||
|
"36": 537.17,
|
||||||
|
"37": 537.38,
|
||||||
|
"38": 537.31,
|
||||||
|
"39": 537.25,
|
||||||
|
"40": 537.28,
|
||||||
|
"41": 537.26,
|
||||||
|
"42": 537.1,
|
||||||
|
"43": 537.19,
|
||||||
|
"44": 537.19,
|
||||||
|
"45": 537.31,
|
||||||
|
"46": 537.19,
|
||||||
|
"47": 537.16,
|
||||||
|
"48": 537.23,
|
||||||
|
"49": 532.91
|
||||||
|
},
|
||||||
|
"expected_e2e_ms": 28769.9,
|
||||||
|
"expected_avg_denoise_ms": 526.34,
|
||||||
|
"expected_median_denoise_ms": 537.19
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,29 @@
|
|||||||
|
"""
|
||||||
|
Config-driven diffusion performance test with pytest parametrization.
|
||||||
|
|
||||||
|
|
||||||
|
If the actual run is significantly better than the baseline, the improved cases with their updated baseline will be printed
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.multimodal_gen.test.server.ascend.testcase_configs_npu import ONE_NPU_CASES
|
||||||
|
from sglang.multimodal_gen.test.server.test_server_common import ( # noqa: F401
|
||||||
|
DiffusionServerBase,
|
||||||
|
diffusion_server,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.test.server.testcase_configs import DiffusionTestCase
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class TestDiffusionServerOneNpu(DiffusionServerBase):
|
||||||
|
"""Performance tests for 1-NPU diffusion cases."""
|
||||||
|
|
||||||
|
@pytest.fixture(params=ONE_NPU_CASES, ids=lambda c: c.id)
|
||||||
|
def case(self, request) -> DiffusionTestCase:
|
||||||
|
"""Provide a DiffusionTestCase for each 1-NPU test."""
|
||||||
|
return request.param
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
from sglang.multimodal_gen.test.server.testcase_configs import (
|
||||||
|
T2V_PROMPT,
|
||||||
|
DiffusionSamplingParams,
|
||||||
|
DiffusionServerArgs,
|
||||||
|
DiffusionTestCase,
|
||||||
|
)
|
||||||
|
|
||||||
|
ONE_NPU_CASES: list[DiffusionTestCase] = [
|
||||||
|
# === Text to Video (T2V) ===
|
||||||
|
DiffusionTestCase(
|
||||||
|
"wan2_1_t2v_1.3b_1_npu",
|
||||||
|
DiffusionServerArgs(
|
||||||
|
model_path="/root/.cache/modelscope/hub/models/Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||||
|
modality="video",
|
||||||
|
warmup=0,
|
||||||
|
custom_validator="video",
|
||||||
|
),
|
||||||
|
DiffusionSamplingParams(
|
||||||
|
prompt=T2V_PROMPT,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
]
|
||||||
@@ -132,6 +132,24 @@ class BaselineConfig:
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def update(self, path: Path):
|
||||||
|
"""Load baseline configuration from JSON file."""
|
||||||
|
with path.open("r", encoding="utf-8") as fh:
|
||||||
|
data = json.load(fh)
|
||||||
|
|
||||||
|
scenarios_new = {}
|
||||||
|
for name, cfg in data["scenarios"].items():
|
||||||
|
scenarios_new[name] = ScenarioConfig(
|
||||||
|
stages_ms=cfg["stages_ms"],
|
||||||
|
denoise_step_ms={int(k): v for k, v in cfg["denoise_step_ms"].items()},
|
||||||
|
expected_e2e_ms=float(cfg["expected_e2e_ms"]),
|
||||||
|
expected_avg_denoise_ms=float(cfg["expected_avg_denoise_ms"]),
|
||||||
|
expected_median_denoise_ms=float(cfg["expected_median_denoise_ms"]),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.scenarios.update(scenarios_new)
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class DiffusionServerArgs:
|
class DiffusionServerArgs:
|
||||||
@@ -729,4 +747,6 @@ TWO_GPU_CASES_B = [
|
|||||||
]
|
]
|
||||||
|
|
||||||
# Load global configuration
|
# Load global configuration
|
||||||
BASELINE_CONFIG = BaselineConfig.load(Path(__file__).with_name("perf_baselines.json"))
|
BASELINE_CONFIG = BaselineConfig.load(
|
||||||
|
Path(__file__).with_name("perf_baselines.json")
|
||||||
|
).update(Path(__file__).parent / "ascend" / "perf_baselines_npu.json")
|
||||||
|
|||||||
Reference in New Issue
Block a user