[Diffusion] [NPU] Wan2.2-T2V-A14B-Diffusers modelslim quantization support (#17996)

Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
Артем Савкин
2026-03-07 17:26:44 +03:00
committed by GitHub
co-authored by ronnie_zheng
parent f8d4eb7022
commit 5297b02c88
19 changed files with 808 additions and 166 deletions
+39
View File
@@ -64,6 +64,7 @@ jobs:
- ".github/workflows/pr-test-npu.yml" - ".github/workflows/pr-test-npu.yml"
multimodal_gen: multimodal_gen:
- "python/sglang/multimodal_gen/**" - "python/sglang/multimodal_gen/**"
- "python/sglang/srt/**"
- "python/pyproject_npu.toml" - "python/pyproject_npu.toml"
- "scripts/ci/npu/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"
@@ -338,3 +339,41 @@ jobs:
export PATH="/usr/local/Ascend/8.3.RC1/compiler/bishengir/bin:${PATH}" export PATH="/usr/local/Ascend/8.3.RC1/compiler/bishengir/bin:${PATH}"
cd python cd python
python3 sglang/multimodal_gen/test/run_suite.py --suite 2-npu python3 sglang/multimodal_gen/test/run_suite.py --suite 2-npu
multimodal-gen-test-8-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.5.0-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: |
cd python
python3 sglang/multimodal_gen/test/run_suite.py --suite 8-npu
@@ -19,3 +19,8 @@ Compressed-tensors (LLM Compressor) on Ascend support:
- [x] [W4A16 MOE](https://github.com/sgl-project/sglang/pull/12759) - [x] [W4A16 MOE](https://github.com/sgl-project/sglang/pull/12759)
- [x] [W8A8 dynamic linear](https://github.com/sgl-project/sglang/pull/14504) - [x] [W8A8 dynamic linear](https://github.com/sgl-project/sglang/pull/14504)
- [x] [W8A8 dynamic MOE](https://github.com/sgl-project/sglang/pull/14504) - [x] [W8A8 dynamic MOE](https://github.com/sgl-project/sglang/pull/14504)
Diffusion model [modelslim](https://github.com/sgl-project/sglang/pull/17996) quantization on Ascend support:
- [x] W4A4 dynamic linear
- [x] W8A8 static linear
- [x] W8A8 dynamic linear
+1 -1
View File
@@ -287,7 +287,7 @@ def _get_config_info(
for registered_model_hf_id in all_model_hf_paths: for registered_model_hf_id in all_model_hf_paths:
registered_model_name = get_model_short_name(registered_model_hf_id.lower()) registered_model_name = get_model_short_name(registered_model_hf_id.lower())
if registered_model_name == model_short_name: if registered_model_name in model_short_name:
logger.debug( logger.debug(
f"Resolved model name '{registered_model_hf_id}' from partial path match." f"Resolved model name '{registered_model_hf_id}' from partial path match."
) )
@@ -234,6 +234,7 @@ class MinimalA2AAttnOp(DistributedAttention):
attention_type: str, attention_type: str,
topk: float, topk: float,
supported_attention_backends: set[AttentionBackendEnum] | None = None, supported_attention_backends: set[AttentionBackendEnum] | None = None,
prefix: str = "",
): ):
dtype = get_compute_dtype() dtype = get_compute_dtype()
attn_backend = get_attn_backend( attn_backend = get_attn_backend(
@@ -256,6 +257,7 @@ class MinimalA2AAttnOp(DistributedAttention):
num_heads=num_heads, num_heads=num_heads,
head_size=head_size, head_size=head_size,
topk_ratio=topk, topk_ratio=topk,
prefix=f"{prefix}.impl",
) )
super(MinimalA2AAttnOp, self).__init__(local_attn) super(MinimalA2AAttnOp, self).__init__(local_attn)
@@ -20,6 +20,7 @@ from sglang.multimodal_gen.runtime.layers.linear import (
RowParallelLinear, RowParallelLinear,
) )
from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig
from sglang.srt.utils import add_prefix
class MLP(nn.Module): class MLP(nn.Module):
@@ -45,6 +46,7 @@ class MLP(nn.Module):
bias=True, bias=True,
gather_output=False, gather_output=False,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("0.proj", prefix),
) )
self.act = get_act_fn(act_type) self.act = get_act_fn(act_type)
@@ -56,6 +58,7 @@ class MLP(nn.Module):
bias=True, bias=True,
input_is_parallel=True, input_is_parallel=True,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("2", prefix),
) )
def forward(self, x: torch.Tensor) -> torch.Tensor: def forward(self, x: torch.Tensor) -> torch.Tensor:
@@ -6,13 +6,15 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor
QuantizationConfig, QuantizationConfig,
) )
from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config
from sglang.multimodal_gen.runtime.layers.quantization.modelslim import ModelSlimConfig
QuantizationMethods = Literal["fp8"] QuantizationMethods = Literal["fp8", "modelslim"]
QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods)) QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods))
# The customized quantization methods which will be added to this dict. # The customized quantization methods which will be added to this dict.
_CUSTOMIZED_METHOD_TO_QUANT_CONFIG = { _CUSTOMIZED_METHOD_TO_QUANT_CONFIG = {
"modelslim": ModelSlimConfig,
"fp8": Fp8Config, "fp8": Fp8Config,
} }
@@ -0,0 +1,224 @@
from __future__ import annotations
import logging
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Dict, List, Mapping, Optional, cast
import torch
from sglang.multimodal_gen.runtime.layers.linear import (
LinearMethodBase,
UnquantizedLinearMethod,
)
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig,
QuantizeMethodBase,
)
from sglang.srt.layers.quantization.compressed_tensors.utils import should_ignore_layer
from sglang.srt.layers.quantization.modelslim.schemes import (
ModelSlimW4A4Int4,
ModelSlimW8A8Int8,
)
if TYPE_CHECKING:
from sglang.srt.layers.quantization.modelslim.modelslim import ModelSlimConfig
from sglang.srt.layers.quantization.modelslim.schemes import (
ModelSlimLinearScheme,
)
logger = logging.getLogger(__name__)
class ModelSlimConfig(QuantizationConfig):
"""
Config class for ModelSlim Quantization of Diffusion models https://gitcode.com/Ascend/msmodelslim, a NPU-specific quantization type.
The quantization method (W8A8, W4A4, etc.) will be automatically parsed from the `quant_model_description.json` config.
ModelSlim for Diffusion models includes support for various quantization schemes, such as:
- W4A4 dynamic linear
- W8A8 static linear
- W8A8 dynamic linear
"""
def __init__(self, quant_config: Dict[str, Any] = {}):
super().__init__()
self.quant_description = quant_config
ignore = cast(List[str], quant_config.get("ignore", []))
self.ignore = ignore
packed_modules_mapping = quant_config.get("packed_modules_mapping", {})
self.packed_modules_mapping = (
packed_modules_mapping if packed_modules_mapping is not None else {}
)
def get_linear_method(self) -> ModelSlimLinearMethod:
return ModelSlimLinearMethod(self)
@classmethod
def get_supported_act_dtypes(cls) -> List[torch.dtype]:
return [torch.int8, torch.float16, torch.bfloat16]
@classmethod
def get_min_capability(cls) -> int:
return 0
@classmethod
def get_name(cls) -> str:
return "modelslim"
@classmethod
def get_config_filenames(cls) -> List[str]:
filenames = ["quant_model_description.json"]
return filenames
@classmethod
def from_config(cls, config: Dict[str, Any]) -> ModelSlimConfig:
return cls(config)
def get_quant_method(
self,
layer: torch.nn.Module,
prefix: str,
) -> Optional[QuantizeMethodBase]:
from sglang.multimodal_gen.runtime.layers.linear import LinearBase
if isinstance(layer, LinearBase):
if should_ignore_layer(
prefix,
ignore=self.ignore,
fused_mapping=self.packed_modules_mapping,
):
return UnquantizedLinearMethod()
key = "model"
packed_modules_mapping_subset = self.packed_modules_mapping.get(key, {})
prefix_in_quant_config = prefix
proj_name = prefix.split(".")[-1]
if proj_name in packed_modules_mapping_subset:
prefix_in_quant_config = prefix.replace(
proj_name, packed_modules_mapping_subset[proj_name][0]
)
if self.is_layer_skipped(prefix, packed_modules_mapping_subset):
return UnquantizedLinearMethod()
scheme = self.get_scheme(layer=layer, layer_name=prefix_in_quant_config)
layer.scheme = scheme
return ModelSlimLinearMethod(self)
else:
return None
def _get_scheme_from_parts(
self,
layer_name: str,
) -> ModelSlimLinearScheme:
quant_type = self.quant_description.get(layer_name + ".weight", "")
if quant_type == "W8A8_DYNAMIC" or quant_type == "W8A8":
return ModelSlimW8A8Int8(
quant_config=self.quant_description, prefix=layer_name
)
elif quant_type == "W4A4_DYNAMIC":
return ModelSlimW4A4Int4(
quant_config=self.quant_description, prefix=layer_name
)
raise NotImplementedError("No modelslim compatible scheme was found.")
def get_scheme(
self, layer: torch.nn.Module, layer_name: Optional[str] = None
) -> Optional[ModelSlimLinearScheme]:
"""
get_scheme method adjusted for modelslim, taken from
python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py
"""
scheme = self._get_scheme_from_parts(
layer_name=layer_name,
)
# Ascend doesn't support device capability
logger.debug("Using scheme: %s for %s", scheme.__class__.__name__, layer_name)
return scheme
def is_layer_skipped(
self, prefix: str, fused_mapping: Mapping[str, List[str]] = MappingProxyType({})
):
# adapted from vllm.model_executor.layers.quantization.utils.quant_utils.is_layer_skipped
proj_name = prefix.split(".")[-1]
if proj_name in fused_mapping:
shard_prefixes = [
prefix.replace(proj_name, shard_proj_name)
for shard_proj_name in fused_mapping[proj_name]
]
is_skipped = None
for shard_prefix in shard_prefixes:
is_shard_skipped = (
self.quant_description.get(shard_prefix + ".weight", "") == "FLOAT"
)
if is_skipped is None:
is_skipped = is_shard_skipped
elif is_shard_skipped != is_skipped:
raise ValueError(
f"Detected some but not all shards of {prefix} "
"are quantized. All shards of fused layers "
"to have the same precision."
)
else:
is_skipped = self.quant_description.get(prefix + ".weight", "") == "FLOAT"
assert is_skipped is not None
return is_skipped
def get_scaled_act_names(self) -> List[str]:
return []
class ModelSlimLinearMethod(LinearMethodBase):
def __init__(self, quantization_config: ModelSlimConfig):
self.quantization_config = quantization_config
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
layer.scheme.process_weights_after_loading(layer)
def create_weights(
self,
layer: torch.nn.Module,
input_size_per_partition: int,
output_partition_sizes: List[int],
input_size: int,
output_size: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
):
"""
Use the ModelSlimLinearScheme associated with each layer to create
the necessary parameters for the layer. See LinearMethodBase for param
details
"""
weight_loader = extra_weight_attrs.get("weight_loader")
layer.scheme.create_weights(
layer=layer,
input_size=input_size,
input_size_per_partition=input_size_per_partition,
output_partition_sizes=output_partition_sizes,
output_size=output_size,
params_dtype=params_dtype,
weight_loader=weight_loader,
)
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: Optional[torch.Tensor] = None,
):
"""
Use the output of create_weights and the CompressedTensorsScheme
associated with the layer to apply the forward pass with the
layer input. See LinearMethodBase for param details
"""
scheme = layer.scheme
if scheme is None:
raise ValueError("A scheme must be defined for each layer")
return scheme.apply_weights(layer, x, bias=bias)
@@ -22,13 +22,18 @@ from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
get_diffusers_component_config, get_diffusers_component_config,
get_metadata_from_safetensors_file,
get_quant_config,
get_quant_config_from_safetensors_metadata,
maybe_download_model, maybe_download_model,
) )
from sglang.multimodal_gen.runtime.utils.logging_utils import get_log_level, init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import get_log_level, init_logger
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
get_metadata_from_safetensors_file,
get_quant_config,
get_quant_config_from_safetensors_metadata,
)
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
from sglang.srt.utils import is_npu
_is_npu = is_npu()
logger = init_logger(__name__) logger = init_logger(__name__)
@@ -75,9 +80,10 @@ class TransformerLoader(ComponentLoader):
hf_config: Dict[str, List[str]], hf_config: Dict[str, List[str]],
server_args: ServerArgs, server_args: ServerArgs,
safetensors_list: list[str], safetensors_list: list[str],
component_model_path: str,
) -> Optional[dict]: ) -> Optional[dict]:
# priority: model config.json → safetensors metadata → nunchaku config # priority: model config.json → safetensors metadata → nunchaku config
quant_config = get_quant_config(hf_config) quant_config = get_quant_config(hf_config, component_model_path)
if quant_config is None and server_args.transformer_weights_path: if quant_config is None and server_args.transformer_weights_path:
# try to read quantization_config from the safetensors metadata header # try to read quantization_config from the safetensors metadata header
for safetensors_file in safetensors_list: for safetensors_file in safetensors_list:
@@ -129,7 +135,10 @@ class TransformerLoader(ComponentLoader):
safetensors_list = self.get_list_of_safetensors_to_load( safetensors_list = self.get_list_of_safetensors_to_load(
server_args, component_model_path server_args, component_model_path
) )
quant_config = self._resolve_quant_config(config, server_args, safetensors_list)
quant_config = self._resolve_quant_config(
config, server_args, safetensors_list, component_model_path
)
# 3. dit config # 3. dit config
# Config from Diffusers supersedes sgl_diffusion's model config # Config from Diffusers supersedes sgl_diffusion's model config
@@ -33,6 +33,9 @@ from sglang.multimodal_gen.runtime.loader.weight_utils import (
from sglang.multimodal_gen.runtime.platforms import current_platform 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
from sglang.multimodal_gen.utils import set_mixed_precision_policy from sglang.multimodal_gen.utils import set_mixed_precision_policy
from sglang.srt.utils import is_npu
_is_npu = is_npu()
logger = init_logger(__name__) logger = init_logger(__name__)
@@ -142,7 +145,13 @@ def maybe_load_fsdp_model(
if quant_method is not None and hasattr( if quant_method is not None and hasattr(
quant_method, "process_weights_after_loading" quant_method, "process_weights_after_loading"
): ):
if _is_npu:
# Activate the NZ format for storing weights,
# which is a specific optimization for Ascend NPU
torch.npu.config.allow_internal_format = True
quant_method.process_weights_after_loading(module) quant_method.process_weights_after_loading(module)
if _is_npu:
torch.npu.empty_cache()
for n, p in chain(model.named_parameters(), model.named_buffers()): for n, p in chain(model.named_parameters(), model.named_buffers()):
if p.is_meta: if p.is_meta:
@@ -58,6 +58,7 @@ from sglang.multimodal_gen.runtime.platforms import (
from sglang.multimodal_gen.runtime.server_args import get_global_server_args from sglang.multimodal_gen.runtime.server_args import get_global_server_args
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.srt.utils import add_prefix
logger = init_logger(__name__) logger = init_logger(__name__)
_is_cuda = current_platform.is_cuda() _is_cuda = current_platform.is_cuda()
@@ -133,6 +134,7 @@ class WanSelfAttention(nn.Module):
qk_norm=True, qk_norm=True,
eps=1e-6, eps=1e-6,
parallel_attention=False, parallel_attention=False,
prefix: str = "",
supported_attention_backends: set[AttentionBackendEnum] | None = None, supported_attention_backends: set[AttentionBackendEnum] | None = None,
is_cross_attention: bool = False, is_cross_attention: bool = False,
quant_config: QuantizationConfig | None = None, quant_config: QuantizationConfig | None = None,
@@ -150,16 +152,32 @@ class WanSelfAttention(nn.Module):
# layers # layers
self.to_q = ColumnParallelLinear( self.to_q = ColumnParallelLinear(
dim, dim, gather_output=False, quant_config=quant_config dim,
dim,
gather_output=False,
quant_config=quant_config,
prefix=add_prefix("to_q", prefix),
) )
self.to_k = ColumnParallelLinear( self.to_k = ColumnParallelLinear(
dim, dim, gather_output=False, quant_config=quant_config dim,
dim,
gather_output=False,
quant_config=quant_config,
prefix=add_prefix("to_k", prefix),
) )
self.to_v = ColumnParallelLinear( self.to_v = ColumnParallelLinear(
dim, dim, gather_output=False, quant_config=quant_config dim,
dim,
gather_output=False,
quant_config=quant_config,
prefix=add_prefix("to_v", prefix),
) )
self.to_out = RowParallelLinear( self.to_out = RowParallelLinear(
dim, dim, input_is_parallel=True, quant_config=quant_config dim,
dim,
input_is_parallel=True,
quant_config=quant_config,
prefix=add_prefix("to_out.0", prefix),
) )
self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity() self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity() self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
@@ -175,6 +193,7 @@ class WanSelfAttention(nn.Module):
causal=False, causal=False,
supported_attention_backends=supported_attention_backends, supported_attention_backends=supported_attention_backends,
skip_sequence_parallel=is_cross_attention, skip_sequence_parallel=is_cross_attention,
quant_config=quant_config,
) )
def forward(self, x: torch.Tensor, context: torch.Tensor, context_lens: int): def forward(self, x: torch.Tensor, context: torch.Tensor, context_lens: int):
@@ -231,6 +250,7 @@ class WanI2VCrossAttention(WanSelfAttention):
window_size=(-1, -1), window_size=(-1, -1),
qk_norm=True, qk_norm=True,
eps=1e-6, eps=1e-6,
prefix: str = "",
supported_attention_backends: set[AttentionBackendEnum] | None = None, supported_attention_backends: set[AttentionBackendEnum] | None = None,
quant_config: QuantizationConfig | None = None, quant_config: QuantizationConfig | None = None,
) -> None: ) -> None:
@@ -246,10 +266,18 @@ class WanI2VCrossAttention(WanSelfAttention):
) )
self.add_k_proj = ColumnParallelLinear( self.add_k_proj = ColumnParallelLinear(
dim, dim, gather_output=False, quant_config=quant_config dim,
dim,
gather_output=False,
quant_config=quant_config,
prefix=add_prefix("add_k_proj", prefix),
) )
self.add_v_proj = ColumnParallelLinear( self.add_v_proj = ColumnParallelLinear(
dim, dim, gather_output=False, quant_config=quant_config dim,
dim,
gather_output=False,
quant_config=quant_config,
prefix=add_prefix("add_v_proj", prefix),
) )
self.norm_added_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity() self.norm_added_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
@@ -328,17 +356,37 @@ class WanTransformerBlock(nn.Module):
dtype=torch.float32, dtype=torch.float32,
) )
self.to_q = ColumnParallelLinear( self.to_q = ColumnParallelLinear(
dim, dim, bias=True, gather_output=False, quant_config=quant_config dim,
dim,
bias=True,
gather_output=False,
quant_config=quant_config,
prefix=add_prefix("attn1.to_q", prefix),
) )
self.to_k = ColumnParallelLinear( self.to_k = ColumnParallelLinear(
dim, dim, bias=True, gather_output=False, quant_config=quant_config dim,
dim,
bias=True,
gather_output=False,
quant_config=quant_config,
prefix=add_prefix("attn1.to_k", prefix),
) )
self.to_v = ColumnParallelLinear( self.to_v = ColumnParallelLinear(
dim, dim, bias=True, gather_output=False, quant_config=quant_config dim,
dim,
bias=True,
gather_output=False,
quant_config=quant_config,
prefix=add_prefix("attn1.to_v", prefix),
) )
self.to_out = RowParallelLinear( self.to_out = RowParallelLinear(
dim, dim, bias=True, reduce_results=True, quant_config=quant_config dim,
dim,
bias=True,
reduce_results=True,
quant_config=quant_config,
prefix=add_prefix("attn1.to_out.0", prefix),
) )
tp_size = get_tp_world_size() tp_size = get_tp_world_size()
self.local_num_heads = divide(num_heads, tp_size) self.local_num_heads = divide(num_heads, tp_size)
@@ -354,6 +402,7 @@ class WanTransformerBlock(nn.Module):
AttentionBackendEnum.SLA_ATTN, AttentionBackendEnum.SLA_ATTN,
AttentionBackendEnum.SAGE_SLA_ATTN, AttentionBackendEnum.SAGE_SLA_ATTN,
}, },
prefix=add_prefix("attn1", prefix),
) )
else: else:
self.attn1 = USPAttention( self.attn1 = USPAttention(
@@ -361,8 +410,9 @@ class WanTransformerBlock(nn.Module):
head_size=dim // num_heads, head_size=dim // num_heads,
causal=False, causal=False,
supported_attention_backends=self_attn_backends, supported_attention_backends=self_attn_backends,
prefix=add_prefix("attn1", prefix),
quant_config=quant_config,
is_cross_attention=False, is_cross_attention=False,
prefix=f"{prefix}.attn1",
) )
self.hidden_dim = dim self.hidden_dim = dim
@@ -399,6 +449,7 @@ class WanTransformerBlock(nn.Module):
num_heads, num_heads,
qk_norm=qk_norm, qk_norm=qk_norm,
eps=eps, eps=eps,
prefix=add_prefix("attn2", prefix),
supported_attention_backends=cross_attn_backends, supported_attention_backends=cross_attn_backends,
quant_config=quant_config, quant_config=quant_config,
) )
@@ -409,6 +460,7 @@ class WanTransformerBlock(nn.Module):
num_heads, num_heads,
qk_norm=qk_norm, qk_norm=qk_norm,
eps=eps, eps=eps,
prefix=add_prefix("attn2", prefix),
supported_attention_backends=cross_attn_backends, supported_attention_backends=cross_attn_backends,
quant_config=quant_config, quant_config=quant_config,
) )
@@ -421,7 +473,11 @@ class WanTransformerBlock(nn.Module):
# 3. Feed-forward # 3. Feed-forward
self.ffn = MLP( self.ffn = MLP(
dim, ffn_dim, act_type="gelu_pytorch_tanh", quant_config=quant_config dim,
ffn_dim,
act_type="gelu_pytorch_tanh",
prefix=add_prefix("ffn.net", prefix),
quant_config=quant_config,
) )
self.mlp_residual = MulAdd() self.mlp_residual = MulAdd()
@@ -555,27 +611,53 @@ class WanTransformerBlock_VSA(nn.Module):
dtype=torch.float32, dtype=torch.float32,
) )
self.to_q = ColumnParallelLinear( self.to_q = ColumnParallelLinear(
dim, dim, bias=True, gather_output=True, quant_config=quant_config dim,
dim,
bias=True,
gather_output=True,
quant_config=quant_config,
prefix=add_prefix("attn1.to_q", prefix),
) )
self.to_k = ColumnParallelLinear( self.to_k = ColumnParallelLinear(
dim, dim, bias=True, gather_output=True, quant_config=quant_config dim,
dim,
bias=True,
gather_output=True,
quant_config=quant_config,
prefix=add_prefix("attn1.to_k", prefix),
) )
self.to_v = ColumnParallelLinear( self.to_v = ColumnParallelLinear(
dim, dim, bias=True, gather_output=True, quant_config=quant_config dim,
dim,
bias=True,
gather_output=True,
quant_config=quant_config,
prefix=add_prefix("attn1.to_v", prefix),
) )
self.to_gate_compress = ColumnParallelLinear( self.to_gate_compress = ColumnParallelLinear(
dim, dim, bias=True, gather_output=True, quant_config=quant_config dim,
dim,
bias=True,
gather_output=True,
quant_config=quant_config,
prefix=add_prefix("attn1.to_gate_compress", prefix),
) )
self.to_out = ColumnParallelLinear( self.to_out = ColumnParallelLinear(
dim, dim, bias=True, gather_output=True, quant_config=quant_config dim,
dim,
bias=True,
gather_output=True,
quant_config=quant_config,
prefix=add_prefix("attn1.to_out.0", prefix),
) )
self.attn1 = UlyssesAttention_VSA( self.attn1 = UlyssesAttention_VSA(
num_heads=num_heads, num_heads=num_heads,
head_size=dim // num_heads, head_size=dim // num_heads,
causal=False, causal=False,
supported_attention_backends=supported_attention_backends, supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn1", prefix=add_prefix("attn1", prefix),
quant_config=quant_config,
) )
self.hidden_dim = dim self.hidden_dim = dim
self.num_attention_heads = num_heads self.num_attention_heads = num_heads
@@ -609,6 +691,7 @@ class WanTransformerBlock_VSA(nn.Module):
num_heads, num_heads,
qk_norm=qk_norm, qk_norm=qk_norm,
eps=eps, eps=eps,
prefix=add_prefix("attn2", prefix),
supported_attention_backends=cross_attn_backends, supported_attention_backends=cross_attn_backends,
quant_config=quant_config, quant_config=quant_config,
) )
@@ -619,6 +702,7 @@ class WanTransformerBlock_VSA(nn.Module):
num_heads, num_heads,
qk_norm=qk_norm, qk_norm=qk_norm,
eps=eps, eps=eps,
prefix=add_prefix("attn2", prefix),
supported_attention_backends=cross_attn_backends, supported_attention_backends=cross_attn_backends,
quant_config=quant_config, quant_config=quant_config,
) )
@@ -631,7 +715,11 @@ class WanTransformerBlock_VSA(nn.Module):
# 3. Feed-forward # 3. Feed-forward
self.ffn = MLP( self.ffn = MLP(
dim, ffn_dim, act_type="gelu_pytorch_tanh", quant_config=quant_config dim,
ffn_dim,
act_type="gelu_pytorch_tanh",
prefix=add_prefix("ffn.net", prefix),
quant_config=quant_config,
) )
self.mlp_residual = MulAdd() self.mlp_residual = MulAdd()
@@ -784,7 +872,7 @@ class WanTransformer3DModel(CachableDiT, OffloadableDiTMixin):
config.added_kv_proj_dim, config.added_kv_proj_dim,
self._supported_attention_backends self._supported_attention_backends
| {AttentionBackendEnum.VIDEO_SPARSE_ATTN}, | {AttentionBackendEnum.VIDEO_SPARSE_ATTN},
prefix=f"{config.prefix}.blocks.{i}", prefix=f"blocks.{i}",
attention_type=config.attention_type, attention_type=config.attention_type,
sla_topk=config.sla_topk, sla_topk=config.sla_topk,
quant_config=quant_config, quant_config=quant_config,
@@ -805,6 +893,7 @@ class WanTransformer3DModel(CachableDiT, OffloadableDiTMixin):
config.out_channels * math.prod(config.patch_size), config.out_channels * math.prod(config.patch_size),
bias=True, bias=True,
gather_output=True, gather_output=True,
prefix=f"proj_out",
quant_config=quant_config, quant_config=quant_config,
) )
self.scale_shift_table = nn.Parameter( self.scale_shift_table = nn.Parameter(
@@ -26,7 +26,7 @@ import shutil
import time import time
from functools import reduce from functools import reduce
from pathlib import Path from pathlib import Path
from typing import Any, Dict, List, Optional, Union, cast from typing import Any, Optional, Union, cast
from diffusers.loaders.lora_base import ( from diffusers.loaders.lora_base import (
_best_guess_weight_name, # watch out for potetential removal from diffusers _best_guess_weight_name, # watch out for potetential removal from diffusers
@@ -38,13 +38,8 @@ from huggingface_hub.errors import (
) )
from requests.exceptions import ConnectionError as RequestsConnectionError from requests.exceptions import ConnectionError as RequestsConnectionError
from requests.exceptions import RequestException from requests.exceptions import RequestException
from safetensors import safe_open
from transformers import AutoConfig, PretrainedConfig from transformers import AutoConfig, PretrainedConfig
from sglang.multimodal_gen.runtime.layers.quantization import (
QuantizationConfig,
get_quantization_config,
)
from sglang.multimodal_gen.runtime.loader.utils import _clean_hf_config_inplace from sglang.multimodal_gen.runtime.loader.utils import _clean_hf_config_inplace
from sglang.multimodal_gen.runtime.loader.weight_utils import get_lock from sglang.multimodal_gen.runtime.loader.weight_utils import get_lock
from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.platforms import current_platform
@@ -342,80 +337,6 @@ def get_diffusers_component_config(
return combined_config return combined_config
def replace_prefix(key: str, prefix_mapping: dict[str, str]) -> str:
for prefix, new_prefix in prefix_mapping.items():
if key.startswith(prefix):
key = key.replace(prefix, new_prefix, 1)
return key
def get_quant_config(
model_config,
packed_modules_mapping: Dict[str, List[str]] = {},
remap_prefix: Dict[str, str] | None = None,
) -> QuantizationConfig:
if "quantization_config" not in model_config:
return None
quant_cls = get_quantization_config(
model_config["quantization_config"]["quant_method"]
)
# GGUF doesn't have config file
if model_config["quantization_config"]["quant_method"] == "gguf":
return quant_cls.from_config({})
# Read the quantization config from the HF model config, if available.
hf_quant_config = model_config["quantization_config"]
# some vision model may keep quantization_config in their text_config
hf_text_config = getattr(model_config, "text_config", None)
if hf_quant_config is None and hf_text_config is not None:
hf_quant_config = getattr(hf_text_config, "quantization_config", None)
if hf_quant_config is None:
# compressed-tensors uses a compressions_config
hf_quant_config = getattr(model_config, "compression_config", None)
if hf_quant_config is not None:
hf_quant_config["packed_modules_mapping"] = packed_modules_mapping
return quant_cls.from_config(hf_quant_config)
# In case of bitsandbytes/QLoRA, get quant config from the adapter model.
else:
model_name_or_path = model_config["model_path"]
is_local = os.path.isdir(model_name_or_path)
hf_folder = model_name_or_path
possible_config_filenames = quant_cls.get_config_filenames()
# If the quantization config is not found, use the default config.
if not possible_config_filenames:
return quant_cls()
config_files = glob.glob(os.path.join(hf_folder, "*.json"))
quant_config_files = [
f for f in config_files if any(f.endswith(x) for x in possible_config_filenames)
]
if len(quant_config_files) == 0:
raise ValueError(
f"Cannot find the config file for {model_config['quantization_config']['quant_method']}"
)
if len(quant_config_files) > 1:
raise ValueError(
f"Found multiple config files for {model_config['quantization_config']['quant_method']}: "
f"{quant_config_files}"
)
quant_config_file = quant_config_files[0]
with open(quant_config_file) as f:
config = json.load(f)
if remap_prefix is not None:
exclude_modules = [
replace_prefix(key, remap_prefix)
for key in config["quantization"]["exclude_modules"]
]
config["quantization"]["exclude_modules"] = exclude_modules
config["packed_modules_mapping"] = packed_modules_mapping
return quant_cls.from_config(config)
# Models don't use the same configuration key for determining the maximum # Models don't use the same configuration key for determining the maximum
# context length. Store them here so we can sanely check them. # context length. Store them here so we can sanely check them.
# NOTE: The ordering here is important. Some models have two of these and we # NOTE: The ordering here is important. Some models have two of these and we
@@ -897,57 +818,3 @@ def snapshot_download(
} }
hf_kwargs.update(kwargs) hf_kwargs.update(kwargs)
return _hf_snapshot_download(**hf_kwargs) return _hf_snapshot_download(**hf_kwargs)
def get_metadata_from_safetensors_file(file_path: str):
try:
with safe_open(file_path, framework="pt", device="cpu") as f:
metadata = f.metadata()
return metadata
except Exception as e:
logger.warning(e)
def get_quant_config_from_safetensors_metadata(
file_path: str,
) -> Optional[QuantizationConfig]:
"""Extract quantization config from a safetensors file's metadata header.
Returns None if no recognizable quantization metadata is found.
"""
metadata = get_metadata_from_safetensors_file(file_path)
if not metadata:
return None
quant_config_str = metadata.get("_quantization_metadata")
if not quant_config_str:
return None
try:
quant_config_dict = json.loads(quant_config_str)
except Exception as _e:
return None
# handle diffusers fp8 safetensors metadata format
if (
"quant_method" not in quant_config_dict
and "format_version" in quant_config_dict
and "layers" in quant_config_dict
):
layers = quant_config_dict.get("layers", {})
if any(
isinstance(v, dict) and "float8" in v.get("format", "")
for v in layers.values()
):
quant_config_dict["quant_method"] = "fp8"
quant_config_dict["activation_scheme"] = "dynamic"
quant_method = quant_config_dict.get("quant_method")
if not quant_method:
return None
try:
quant_cls = get_quantization_config(quant_method)
config = quant_cls.from_config(quant_config_dict)
logger.debug(f"Get quantization config from safetensors file: {file_path}")
return config
except Exception as _e:
return None
@@ -0,0 +1,169 @@
import glob
import json
import os
from pathlib import Path
from typing import Dict, List, Optional
from safetensors import safe_open
from sglang.multimodal_gen.runtime.layers.quantization import (
QuantizationConfig,
get_quantization_config,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
def find_quant_modelslim_config(model_config, component_model_path):
quant_config_file = Path(component_model_path, "quant_model_description.json")
quant_cfg = None
if quant_config_file.is_file():
with open(quant_config_file) as f:
quant_cfg = json.load(f)
# This field is required for flagless model loading but is not present in
# modelslim model description, so we're adding it here manually.
quant_cfg["quant_method"] = "modelslim"
return quant_cfg
def replace_prefix(key: str, prefix_mapping: dict[str, str]) -> str:
for prefix, new_prefix in prefix_mapping.items():
if key.startswith(prefix):
key = key.replace(prefix, new_prefix, 1)
return key
def get_quant_config(
model_config,
component_model_path: str,
packed_modules_mapping: Dict[str, List[str]] = {},
remap_prefix: Dict[str, str] | None = None,
) -> QuantizationConfig:
quant_cfg = find_quant_modelslim_config(model_config, component_model_path)
if quant_cfg is not None:
quant_cls = get_quantization_config(quant_cfg["quant_method"])
return quant_cls.from_config(quant_cfg)
else:
if "quantization_config" not in model_config:
return None
quant_cls = get_quantization_config(
model_config["quantization_config"]["quant_method"]
)
# GGUF doesn't have config file
if model_config["quantization_config"]["quant_method"] == "gguf":
return quant_cls.from_config({})
# Read the quantization config from the HF model config, if available.
hf_quant_config = model_config["quantization_config"]
# some vision model may keep quantization_config in their text_config
hf_text_config = getattr(model_config, "text_config", None)
if hf_quant_config is None and hf_text_config is not None:
hf_quant_config = getattr(hf_text_config, "quantization_config", None)
if hf_quant_config is None:
# compressed-tensors uses a compressions_config
hf_quant_config = getattr(model_config, "compression_config", None)
if hf_quant_config is not None:
hf_quant_config["packed_modules_mapping"] = packed_modules_mapping
return quant_cls.from_config(hf_quant_config)
# In case of bitsandbytes/QLoRA, get quant config from the adapter model.
else:
model_name_or_path = model_config["model_path"]
is_local = os.path.isdir(model_name_or_path)
hf_folder = model_name_or_path
possible_config_filenames = quant_cls.get_config_filenames()
# If the quantization config is not found, use the default config.
if not possible_config_filenames:
return quant_cls()
config_files = glob.glob(os.path.join(hf_folder, "*.json"))
quant_config_files = [
f
for f in config_files
if any(f.endswith(x) for x in possible_config_filenames)
]
if len(quant_config_files) == 0:
raise ValueError(
f"Cannot find the config file for {model_config['quantization_config']['quant_method']}"
)
if len(quant_config_files) > 1:
raise ValueError(
f"Found multiple config files for {model_config['quantization_config']['quant_method']}: "
f"{quant_config_files}"
)
quant_config_file = quant_config_files[0]
with open(quant_config_file) as f:
config = json.load(f)
if remap_prefix is not None:
exclude_modules = [
replace_prefix(key, remap_prefix)
for key in config["quantization"]["exclude_modules"]
]
config["quantization"]["exclude_modules"] = exclude_modules
config["packed_modules_mapping"] = packed_modules_mapping
return quant_cls.from_config(config)
def handle_fp8_metadata_format(quant_config_dict):
layers = quant_config_dict.get("layers", {})
if any(
isinstance(v, dict) and "float8" in v.get("format", "") for v in layers.values()
):
quant_config_dict["quant_method"] = "fp8"
quant_config_dict["activation_scheme"] = "dynamic"
return quant_config_dict
def get_quant_config_from_safetensors_metadata(
file_path: str,
) -> Optional[QuantizationConfig]:
"""Extract quantization config from a safetensors file's metadata header.
Returns None if no recognizable quantization metadata is found.
"""
metadata = get_metadata_from_safetensors_file(file_path)
if not metadata:
return None
quant_config_str = metadata.get("_quantization_metadata")
if not quant_config_str:
return None
try:
quant_config_dict = json.loads(quant_config_str)
except Exception as _e:
return None
# handle diffusers fp8 safetensors metadata format
if (
"quant_method" not in quant_config_dict
and "format_version" in quant_config_dict
and "layers" in quant_config_dict
):
quant_config_dict = handle_fp8_metadata_format(quant_config_dict)
quant_method = quant_config_dict.get("quant_method")
if not quant_method:
return None
try:
quant_cls = get_quantization_config(quant_method)
config = quant_cls.from_config(quant_config_dict)
logger.debug(f"Get quantization config from safetensors file: {file_path}")
return config
except Exception as _e:
return None
def get_metadata_from_safetensors_file(file_path: str):
try:
with safe_open(file_path, framework="pt", device="cpu") as f:
metadata = f.metadata()
return metadata
except Exception as e:
logger.warning(e)
@@ -61,6 +61,10 @@ suites_ascend = {
"ascend/test_server_2_npu.py", "ascend/test_server_2_npu.py",
# add new 2-npu test files here # add new 2-npu test files here
], ],
"8-npu": [
"ascend/test_server_8_npu.py",
# add new 8-npu test files here
],
} }
SUITES.update(suites_ascend) SUITES.update(suites_ascend)
@@ -201,6 +201,62 @@
"expected_e2e_ms": 38738.17, "expected_e2e_ms": 38738.17,
"expected_avg_denoise_ms": 523.62, "expected_avg_denoise_ms": 523.62,
"expected_median_denoise_ms": 536.23 "expected_median_denoise_ms": 536.23
},
"wan2_2_t2v_14b_w8a8_8npu": {
"stages_ms": {
"InputValidationStage": 0.07,
"TextEncodingStage": 301.21,
"LatentPreparationStage": 0.2,
"TimestepPreparationStage": 2.68,
"DenoisingStage": 83661.46,
"DecodingStage": 232.94,
"per_frame_generation": null
},
"denoise_step_ms": {
"0": 1919.92,
"1": 2099.45,
"2": 2092.11,
"3": 2090.84,
"4": 2089.89,
"5": 2090.6,
"6": 2090.77,
"7": 2091.43,
"8": 2091.24,
"9": 2067.83,
"10": 2078.02,
"11": 2090.75,
"12": 2108.36,
"13": 2096.16,
"14": 2091.74,
"15": 2091.47,
"16": 2091.6,
"17": 2091.94,
"18": 2091.39,
"19": 2090.69,
"20": 2090.27,
"21": 2090.77,
"22": 2090.24,
"23": 2091.65,
"24": 2091.21,
"25": 2126.82,
"26": 2338.39,
"27": 2085.18,
"28": 2084.68,
"29": 2084.71,
"30": 2051.48,
"31": 2104.3,
"32": 2084.58,
"33": 2085.04,
"34": 2085.03,
"35": 2084.58,
"36": 2084.41,
"37": 2085.16,
"38": 2084.88,
"39": 2083.54
},
"expected_e2e_ms": 91733.92,
"expected_avg_denoise_ms": 2091.33,
"expected_median_denoise_ms": 2090.72
} }
} }
} }
@@ -0,0 +1,31 @@
"""
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 (
EIGHT_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 TestDiffusionServerEightNpu(DiffusionServerBase):
"""Performance tests for 8-NPU diffusion cases."""
@pytest.fixture(params=EIGHT_NPU_CASES, ids=lambda c: c.id)
def case(self, request) -> DiffusionTestCase:
"""Provide a DiffusionTestCase for each 8-NPU test."""
return request.param
@@ -43,3 +43,20 @@ TWO_NPU_CASES: list[DiffusionTestCase] = [
T2I_sampling_params, T2I_sampling_params,
), ),
] ]
EIGHT_NPU_CASES: list[DiffusionTestCase] = [
# === Text to Video (T2V) ===
DiffusionTestCase(
"wan2_2_t2v_14b_w8a8_8npu",
DiffusionServerArgs(
model_path="/root/.cache/modelscope/hub/models/Eco-Tech/Wan2.2-T2V-A14B-Diffusers-w8a8",
modality="video",
custom_validator="video",
num_gpus=8,
tp_size=4,
),
DiffusionSamplingParams(
prompt=T2V_PROMPT,
),
),
]
@@ -0,0 +1,115 @@
### Based on https://github.com/huggingface/diffusers/blob/main/scripts/convert_wan_to_diffusers.py
import argparse
import json
import pathlib
from typing import Any, Dict, Tuple
from safetensors.torch import load_file, save_file
TRANSFORMER_KEYS_RENAME_DICT = {
"time_embedding.0": "condition_embedder.time_embedder.linear_1",
"time_embedding.2": "condition_embedder.time_embedder.linear_2",
"text_embedding.0": "condition_embedder.text_embedder.linear_1",
"text_embedding.2": "condition_embedder.text_embedder.linear_2",
"time_projection.1": "condition_embedder.time_proj",
"head.modulation": "scale_shift_table",
"head.head": "proj_out",
"modulation": "scale_shift_table",
"ffn.0": "ffn.net.0.proj",
"ffn.2": "ffn.net.2",
# Hack to swap the layer names
# The original model calls the norms in following order: norm1, norm3, norm2
# We convert it to: norm1, norm2, norm3
"norm2": "norm__placeholder",
"norm3": "norm2",
"norm__placeholder": "norm3",
# For the I2V model
"img_emb.proj.0": "condition_embedder.image_embedder.norm1",
"img_emb.proj.1": "condition_embedder.image_embedder.ff.net.0.proj",
"img_emb.proj.3": "condition_embedder.image_embedder.ff.net.2",
"img_emb.proj.4": "condition_embedder.image_embedder.norm2",
# for the FLF2V model
"img_emb.emb_pos": "condition_embedder.image_embedder.pos_embed",
# Add attention component mappings
"self_attn.q": "attn1.to_q",
"self_attn.k": "attn1.to_k",
"self_attn.v": "attn1.to_v",
"self_attn.o": "attn1.to_out.0",
"self_attn.norm_q": "attn1.norm_q",
"self_attn.norm_k": "attn1.norm_k",
"cross_attn.q": "attn2.to_q",
"cross_attn.k": "attn2.to_k",
"cross_attn.v": "attn2.to_v",
"cross_attn.o": "attn2.to_out.0",
"cross_attn.norm_q": "attn2.norm_q",
"cross_attn.norm_k": "attn2.norm_k",
"attn2.to_k_img": "attn2.add_k_proj",
"attn2.to_v_img": "attn2.add_v_proj",
"attn2.norm_k_img": "attn2.norm_added_k",
}
def get_transformer_config(model_type: str) -> Tuple[Dict[str, Any], ...]:
if model_type == "Wan-T2V-14B":
RENAME_DICT = TRANSFORMER_KEYS_RENAME_DICT
return RENAME_DICT
def update_dict_(dict: Dict[str, Any], old_key: str, new_key: str) -> Dict[str, Any]:
dict[new_key] = dict.pop(old_key)
def load_sharded_safetensors(path: pathlib.Path):
file_path = path
state_dict = {}
state_dict.update(load_file(file_path))
return state_dict
def convert_transformer(model_type: str, model_dir: str, output_dir: str):
pathlib.Path(output_dir).mkdir(parents=True, exist_ok=True)
RENAME_DICT = get_transformer_config(model_type)
original_state_dict = load_sharded_safetensors(
pathlib.Path(model_dir, "*model*.safetensors")
)
with open(pathlib.Path(model_dir, "*quant_model_description*.json")) as f:
original_quant_config = json.load(f)
for key in list(original_state_dict.keys()):
new_key = key[:]
for replace_key, rename_key in RENAME_DICT.items():
new_key = new_key.replace(replace_key, rename_key)
update_dict_(original_state_dict, key, new_key)
update_dict_(original_quant_config, key, new_key)
save_file(
original_state_dict,
pathlib.Path(output_dir, "diffusion_pytorch_model.safetensors"),
)
with open(pathlib.Path(output_dir, "quant_model_description.json"), "w") as f:
json.dump(original_quant_config, f)
def get_args():
parser = argparse.ArgumentParser()
parser.add_argument("--input-path", type=str, required=True)
parser.add_argument("--output-path", type=str, required=True)
return parser.parse_args()
if __name__ == "__main__":
args = get_args()
convert_transformer(
"Wan-T2V-14B",
model_dir=pathlib.Path(args.input_path, "high_noise_model"),
output_dir=pathlib.Path(args.output_path, "transformer"),
)
convert_transformer(
"Wan-T2V-14B",
model_dir=pathlib.Path(args.input_path, "low_noise_model"),
output_dir=pathlib.Path(args.output_path, "transformer_2"),
)
@@ -105,7 +105,7 @@ class NPUW8A8Int8DynamicLinearMethod(_NPULinearMethodBase):
quant_out, quant_out,
layer.weight, layer.weight,
layer.weight_scale, layer.weight_scale,
pertoken_scale=dynamic_scale, pertoken_scale=dynamic_scale.flatten(),
bias=bias, bias=bias,
output_dtype=original_dtype, output_dtype=original_dtype,
) )
@@ -137,7 +137,7 @@ class NPU_W4A4DynamicLinearMethod(_NPULinearMethodBase):
quant_out, quant_out,
layer.weight, layer.weight,
layer.weight_scale, layer.weight_scale,
pertoken_scale=dynamic_scale, pertoken_scale=dynamic_scale.flatten(),
bias=bias, bias=bias,
output_dtype=original_dtype, output_dtype=original_dtype,
) )
@@ -123,9 +123,10 @@ def npu_format_cast(
if envs.SGLANG_NPU_DISABLE_ACL_FORMAT_WEIGHT.get(): if envs.SGLANG_NPU_DISABLE_ACL_FORMAT_WEIGHT.get():
return tensor return tensor
import torch_npu if tensor.device == torch.device("cpu"):
return torch.ops.npu.npu_format_cast(tensor.npu(), acl_format.value).cpu()
return torch_npu.npu_format_cast(tensor, acl_format.value) else:
return torch.ops.npu.npu_format_cast(tensor, acl_format.value)
def get_indexer_weight_stream(): def get_indexer_weight_stream():