[diffusion][npu][quant] Add MXFP8 quantization support for Wan2.2 Diffusion on Ascend NPU (#20922)

Co-authored-by: ronnie_zheng <zl19940307@163.com>
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
This commit is contained in:
Junlin Wu
2026-05-07 21:30:56 +03:00
committed by GitHub
co-authored by ronnie_zheng github-actions[bot]
parent 7d397ad23d
commit 80a6014243
16 changed files with 706 additions and 144 deletions
@@ -14,9 +14,10 @@ from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
ModelOptFp8Config,
)
from sglang.multimodal_gen.runtime.layers.quantization.modelslim import ModelSlimConfig
from sglang.multimodal_gen.runtime.layers.quantization.mxfp8_npu import MXFP8Config
QuantizationMethods = Literal[
"fp8", "modelopt", "modelopt_fp8", "modelopt_fp4", "modelslim"
"fp8", "modelopt", "modelopt_fp8", "modelopt_fp4", "modelslim", "mxfp8"
]
QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods))
@@ -28,6 +29,7 @@ _CUSTOMIZED_METHOD_TO_QUANT_CONFIG = {
"modelopt_fp4": ModelOptFp4Config,
"modelslim": ModelSlimConfig,
"fp8": Fp8Config,
"mxfp8": MXFP8Config,
}
@@ -119,6 +119,12 @@ class ModelSlimConfig(QuantizationConfig):
return ModelSlimW4A4Int4(
quant_config=self.quant_description, prefix=layer_name
)
elif quant_type == "W8A8_MXFP8":
from sglang.multimodal_gen.runtime.layers.quantization.modelslim_mxfp8_scheme import (
ModelSlimMXFP8Scheme,
)
return ModelSlimMXFP8Scheme()
raise NotImplementedError("No modelslim compatible scheme was found.")
def get_scheme(
@@ -0,0 +1,124 @@
"""ModelSlim MXFP8 scheme for pre-quantized weight inference on Ascend NPU.
Loads weights pre-quantized by msmodelslim (float8_e4m3fn weights,
uint8 scales) and runs MXFP8 matmul at inference.
"""
from typing import List, Optional
import torch
from sglang.multimodal_gen.runtime.platforms import current_platform
_is_npu = current_platform.is_npu()
if _is_npu:
import torch_npu
from sglang.multimodal_gen.runtime.models.parameter import (
GroupQuantScaleParameter,
ModelWeightParameter,
)
from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimLinearScheme
MXFP8_BLOCK_SIZE = 32
class ModelSlimMXFP8Scheme(ModelSlimLinearScheme):
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,
):
weight_loader = extra_weight_attrs.get("weight_loader")
output_size_per_partition = sum(output_partition_sizes)
# msmodelslim exports weight as float8_e4m3fn, shape [out, in]
weight = ModelWeightParameter(
data=torch.empty(
(output_size_per_partition, input_size_per_partition),
dtype=torch.float8_e4m3fn,
),
input_dim=1,
output_dim=0,
weight_loader=weight_loader,
)
layer.register_parameter("weight", weight)
# msmodelslim exports weight_scale as uint8, shape [out, in/32].
# NOTE: This parameter is intentionally named "weight_scale" (not
# "weight_scale_inv" as used in mxfp8_npu.py) because the weight loader
# matches parameter names to checkpoint keys, and msmodelslim checkpoints
# store this tensor under the key "<layer>.weight_scale".
scale_dim = input_size_per_partition // MXFP8_BLOCK_SIZE
weight_scale = GroupQuantScaleParameter(
data=torch.empty(
(output_size_per_partition, scale_dim),
dtype=torch.uint8,
),
input_dim=1,
output_dim=0,
weight_loader=weight_loader,
)
layer.register_parameter("weight_scale", weight_scale)
def process_weights_after_loading(self, layer: torch.nn.Module):
# weight is already float8_e4m3fn, no cast needed
weight = layer.weight.data
layer.weight = torch.nn.Parameter(weight, requires_grad=False)
# Reshape weight_scale: [out, in/32] -> [out, in/32//2, 2]
weight_scale = layer.weight_scale.data
weight_scale = weight_scale.reshape(weight_scale.shape[0], -1, 2)
layer.weight_scale = torch.nn.Parameter(weight_scale, requires_grad=False)
def apply_weights(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
original_dtype = x.dtype
if original_dtype not in (torch.float16, torch.bfloat16):
# npu_dynamic_mx_quant only accepts fp16/bf16 activations
x = x.to(torch.bfloat16)
original_dtype = torch.bfloat16
# npu_dynamic_mx_quant requires a 2D input [tokens, hidden_size].
# Diffusion transformer inputs are typically 3D [batch, seq, hidden] or
# higher. Flattening to 2D merges all leading dimensions into a single
# token axis so the NPU kernel can compute per-token MXFP8 scales, then
# we restore the original shape from the output.
input_shape = x.shape
x_2d = x.reshape(-1, x.shape[-1])
# Dynamic MXFP8 activation quantisation
qx, input_scale = torch_npu.npu_dynamic_mx_quant(
x_2d, dst_type=torch_npu.float8_e4m3fn
)
# MXFP8 matmul
output = torch_npu.npu_quant_matmul(
qx,
layer.weight.transpose(0, 1),
layer.weight_scale.transpose(0, 1),
scale_dtype=torch_npu.float8_e8m0fnu,
pertoken_scale=input_scale,
pertoken_scale_dtype=torch_npu.float8_e8m0fnu,
bias=bias.to(torch.float32) if bias is not None else None,
output_dtype=original_dtype,
group_sizes=[1, 1, MXFP8_BLOCK_SIZE],
)
# Restore original shape
output_shape = list(input_shape[:-1]) + [output.shape[-1]]
output = output.reshape(output_shape)
return output
@@ -0,0 +1,176 @@
"""Online MXFP8 quantization for Diffusion models on Ascend NPU.
Provides ``MXFP8Config`` (registered as ``"mxfp8"``) and
``NPUMXFP8DiffusionLinearMethod`` which quantise FP16/BF16 weights to MXFP8
at load time and use ``npu_dynamic_mx_quant`` + ``npu_quant_matmul`` for
inference, mirroring the LLM-side ``NPUMXFP8LinearMethod``.
"""
from __future__ import annotations
from typing import Any, Dict, List, Optional
import torch
from torch.nn.parameter import Parameter
from sglang.multimodal_gen.runtime.platforms import current_platform
_is_npu = current_platform.is_npu()
if _is_npu:
import torch_npu
from sglang.multimodal_gen.runtime.layers.linear import LinearBase, LinearMethodBase
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig,
QuantizeMethodBase,
)
from sglang.multimodal_gen.runtime.models.parameter import ModelWeightParameter
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
MXFP8_BLOCK_SIZE = 32
class MXFP8Config(QuantizationConfig):
"""Config for online MXFP8 quantization on Ascend NPU (Diffusion)."""
def __init__(self) -> None:
super().__init__()
@classmethod
def get_name(cls) -> str:
return "mxfp8"
@classmethod
def get_supported_act_dtypes(cls) -> List[torch.dtype]:
return [torch.bfloat16, torch.float16]
@classmethod
def get_min_capability(cls) -> int:
return 0 # NPU, not CUDA
@classmethod
def get_config_filenames(cls) -> List[str]:
return []
@classmethod
def from_config(cls, config: Dict[str, Any]) -> "MXFP8Config":
return cls()
def get_quant_method(
self, layer: torch.nn.Module, prefix: str
) -> Optional[QuantizeMethodBase]:
if isinstance(layer, LinearBase):
return NPUMXFP8DiffusionLinearMethod(self)
return None
def get_scaled_act_names(self) -> List[str]:
return []
class NPUMXFP8DiffusionLinearMethod(LinearMethodBase):
"""Ascend NPU MXFP8 linear method for Diffusion models.
Online mode: loads FP16/BF16 weights → quantises to MXFP8 at load time.
Inference: dynamic MXFP8 activation quant + MXFP8 matmul (block_size=32).
"""
def __init__(self, quant_config: MXFP8Config):
self.quant_config = quant_config
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,
):
output_size_per_partition = sum(output_partition_sizes)
weight_loader = extra_weight_attrs.get("weight_loader")
layer.logical_widths = output_partition_sizes
layer.input_size_per_partition = input_size_per_partition
layer.output_size_per_partition = output_size_per_partition
layer.orig_dtype = params_dtype
# Load weights in original dtype; quantise later in process_weights_after_loading
weight = ModelWeightParameter(
data=torch.empty(
output_size_per_partition,
input_size_per_partition,
dtype=params_dtype,
),
input_dim=1,
output_dim=0,
weight_loader=weight_loader,
)
layer.register_parameter("weight", weight)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
weight_fp = layer.weight.data
if weight_fp.dtype not in (torch.float16, torch.bfloat16):
weight_fp = weight_fp.to(torch.bfloat16)
# Move weight to NPU if needed. We intentionally use a conditional
# move rather than an assert because `dit_cpu_offload` defaults to
# True in ServerArgs, which causes fsdp_load to move every parameter
# back to CPU after loading (even when the target device is NPU).
# npu_dynamic_mx_quant requires an NPU tensor, so we must transfer
# here. The quantized fp8 weights produced below will remain on NPU
# for inference; if the model still needs to be offloaded after
# quantization (e.g. very large model on a small NPU), a higher-level
# offload pass can move them back afterwards.
if not weight_fp.is_npu:
weight_fp = weight_fp.to(f"npu:{torch.npu.current_device()}")
# Online MXFP8 quantisation of weights (block_size=32)
qw, w_scale = torch_npu.npu_dynamic_mx_quant(
weight_fp, dst_type=torch_npu.float8_e4m3fn
)
layer.weight = Parameter(qw, requires_grad=False)
layer.weight_scale_inv = Parameter(w_scale, requires_grad=False)
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
original_dtype = x.dtype
if original_dtype not in (torch.float16, torch.bfloat16):
x = x.to(torch.bfloat16)
original_dtype = torch.bfloat16
# Flatten to 2D [tokens, hidden] so npu_dynamic_mx_quant returns 3D scale
input_shape = x.shape
x_2d = x.reshape(-1, x.shape[-1])
# Dynamic MXFP8 activation quantisation
qx, input_scale = torch_npu.npu_dynamic_mx_quant(
x_2d, dst_type=torch_npu.float8_e4m3fn
)
# MXFP8 matmul
output = torch_npu.npu_quant_matmul(
qx,
layer.weight.transpose(0, 1),
layer.weight_scale_inv.transpose(0, 1),
scale_dtype=torch_npu.float8_e8m0fnu,
pertoken_scale=input_scale,
pertoken_scale_dtype=torch_npu.float8_e8m0fnu,
bias=bias.to(torch.float32) if bias is not None else None,
output_dtype=original_dtype,
group_sizes=[1, 1, MXFP8_BLOCK_SIZE],
)
# Restore original shape (replace last dim with output features)
output_shape = list(input_shape[:-1]) + [output.shape[-1]]
output = output.reshape(output_shape)
return output
@@ -613,6 +613,7 @@ def load_model_from_full_model_state_dict(
"bias",
"norm_q",
"norm_k",
"weight_scale",
]
for new_param_name in unused_keys:
meta_sharded_param = meta_sd.get(new_param_name)
@@ -477,8 +477,17 @@ def _resolve_quant_config(
) -> Optional[QuantizationConfig]:
"""
resolve quant config from checkpoints' metadata
priority: model config.json -> safetensors metadata -> format-specific fallback
priority: explicit --quantization flag -> model config.json -> safetensors metadata -> format-specific fallback
"""
# priority: explicit --quantization flag (e.g. mxfp8, mxfp4, modelslim)
if server_args.quantization is not None:
from sglang.multimodal_gen.runtime.layers.quantization import (
get_quantization_config,
)
quant_cls = get_quantization_config(server_args.quantization)
return quant_cls.from_config({})
arch_config = server_args.pipeline_config.dit_config.arch_config
param_names_mapping_dict = arch_config.param_names_mapping
reverse_param_names_mapping_dict = getattr(
@@ -214,6 +214,10 @@ class ServerArgs(DisaggArgsMixin):
disable_autocast: bool | None = None
# Explicit quantization method override (e.g. "mxfp8", "fp8", "modelslim").
# When set, the transformer loader will use this instead of auto-detection.
quantization: str | None = None
# Quantization / Nunchaku SVDQuant configuration
nunchaku_config: NunchakuSVDQuantArgs | NunchakuConfig | None = field(
default_factory=NunchakuSVDQuantArgs, repr=False
@@ -1138,6 +1142,14 @@ class ServerArgs(DisaggArgsMixin):
help="Disable autocast for denoising loop and vae decoding in pipeline sampling",
)
parser.add_argument(
"--quantization",
type=str,
default=None,
help='Quantization method override (e.g. "mxfp8", "fp8", "modelslim"). '
"When set, the transformer loader will use this instead of auto-detection.",
)
# Nunchaku SVDQuant quantization parameters
NunchakuSVDQuantArgs.add_cli_args(parser)
@@ -97,9 +97,16 @@ def _load_quant_cls(quant_cfg: dict):
def find_quant_modelslim_config(model_config, component_model_path):
# Try exact name first, then glob for variant filenames (e.g. after repack)
quant_config_file = Path(component_model_path, "quant_model_description.json")
if not quant_config_file.is_file():
candidates = sorted(
Path(component_model_path).glob("quant_model_description*.json")
)
quant_config_file = candidates[0] if candidates else None
quant_cfg = None
if quant_config_file.is_file():
if quant_config_file is not None and Path(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
@@ -84,6 +84,7 @@ class TestTransformerQuantHelpers(unittest.TestCase):
),
),
nunchaku_config=None,
quantization=None,
tp_size=1,
dit_cpu_offload=False,
text_encoder_cpu_offload=False,
+225 -115
View File
@@ -1,115 +1,225 @@
### 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"),
)
### Based on https://github.com/huggingface/diffusers/blob/main/scripts/convert_wan_to_diffusers.py
import argparse
import json
import pathlib
import shutil
from typing import Any, Dict, List
from safetensors.torch import load_file, save_file
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
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",
}
SUPPORTED_MODEL_TYPES = ["Wan2.2-T2V-A14B", "Wan2.2-I2V-A14B", "Wan2.2-TI2V-5B"]
# Cascade models have two transformers (high_noise + low_noise)
CASCADE_MODEL_TYPES = {"Wan2.2-T2V-A14B", "Wan2.2-I2V-A14B"}
def get_transformer_config(model_type: str) -> Dict[str, Any]:
if model_type in SUPPORTED_MODEL_TYPES:
return TRANSFORMER_KEYS_RENAME_DICT
else:
raise ValueError(
f"Unsupported model_type: {model_type}. Supported: {SUPPORTED_MODEL_TYPES}"
)
def get_transformer_dirs(model_type: str) -> List[str]:
"""Return the list of transformer directory names for a given model type."""
if model_type in CASCADE_MODEL_TYPES:
return ["transformer", "transformer_2"]
return ["transformer"]
def get_quant_subpath(
model_type: str, quant_path: pathlib.Path, transformer_dir: str
) -> pathlib.Path:
"""Return the quant weights subdirectory for a given transformer."""
if model_type in CASCADE_MODEL_TYPES:
sub = (
"high_noise_model"
if transformer_dir == "transformer"
else "low_noise_model"
)
return quant_path / sub
return quant_path
def update_dict_(d: Dict[str, Any], old_key: str, new_key: str) -> None:
d[new_key] = d.pop(old_key)
def load_sharded_safetensors(directory: pathlib.Path, pattern: str) -> dict:
candidates = sorted(directory.glob(pattern))
if not candidates:
raise FileNotFoundError(f"No file matching '{pattern}' found in {directory}")
if len(candidates) > 1:
raise FileNotFoundError(
f"Multiple files matching '{pattern}' found in {directory}: {candidates}"
)
state_dict = {}
state_dict.update(load_file(candidates[0]))
return state_dict
def convert_transformer(
model_type: str, model_dir: pathlib.Path, output_dir: pathlib.Path
) -> None:
"""Convert a single quantized transformer directory into Diffusers format."""
model_path = pathlib.Path(model_dir)
out_path = pathlib.Path(output_dir)
out_path.mkdir(parents=True, exist_ok=True)
RENAME_DICT = get_transformer_config(model_type)
state_dict = load_sharded_safetensors(model_path, "quant_model_weight*.safetensors")
json_candidates = sorted(model_path.glob("quant_model_description*.json"))
if not json_candidates:
raise FileNotFoundError(
f"No quant_model_description*.json found in {model_path}"
)
with open(json_candidates[0]) as f:
quant_config = json.load(f)
for key in list(state_dict.keys()):
new_key = key[:]
for replace_key, rename_key in RENAME_DICT.items():
new_key = new_key.replace(replace_key, rename_key)
if new_key != key:
update_dict_(state_dict, key, new_key)
# The quant JSON only covers quantized layers, not all model keys
if key in quant_config:
update_dict_(quant_config, key, new_key)
save_file(state_dict, out_path / "diffusion_pytorch_model.safetensors")
with open(out_path / "quant_model_description.json", "w") as f:
json.dump(quant_config, f, indent=2)
def repack(
model_type: str,
original_model_path: pathlib.Path,
quant_path: pathlib.Path,
output_path: pathlib.Path,
) -> None:
"""
Full one-step repack workflow:
1. Copy the original HF Diffusers model to output_path, excluding transformer dir(s).
2. For each transformer: convert quant weights and copy config.json from original.
"""
transformer_dirs = get_transformer_dirs(model_type)
# Step 1: Copy original model, skipping transformer dirs (they will be replaced)
logger.debug(f"Step 1: Copying original model to {output_path}")
logger.debug(f" (skipping: {transformer_dirs})")
shutil.copytree(
str(original_model_path),
str(output_path),
ignore=shutil.ignore_patterns(*transformer_dirs),
)
# Step 2+: Convert each transformer
for i, tdir in enumerate(transformer_dirs):
q_path = get_quant_subpath(model_type, quant_path, tdir)
out_tdir = output_path / tdir
logger.debug(
f"\nStep {i + 2}: Converting {tdir} (quant source: {q_path.name})..."
)
convert_transformer(model_type, q_path, out_tdir)
# Copy config.json from the original transformer dir
src_config = original_model_path / tdir / "config.json"
if src_config.is_file():
shutil.copy2(str(src_config), str(out_tdir / "config.json"))
logger.debug(f" Copied config.json from original {tdir}/")
logger.info(f"\nDone! Repacked model saved to: {output_path}")
def get_args():
parser = argparse.ArgumentParser(
description="Repack msmodelslim quantized Wan2.2 weights into HF Diffusers format"
)
parser.add_argument(
"--model-type",
type=str,
required=True,
choices=SUPPORTED_MODEL_TYPES,
help="Model type to convert",
)
parser.add_argument(
"--original-model-path",
type=str,
required=True,
help="Path to the original HF Diffusers model (e.g., /weights/Wan2.2-TI2V-5B-Diffusers)",
)
parser.add_argument(
"--quant-path",
type=str,
required=True,
help="Path to msmodelslim quantized weights directory",
)
parser.add_argument(
"--output-path",
type=str,
required=True,
help="Output path for the repacked model (e.g., /weights/Wan2.2-TI2V-5B-Diffusers-MXFP8)",
)
return parser.parse_args()
if __name__ == "__main__":
args = get_args()
repack(
model_type=args.model_type,
original_model_path=pathlib.Path(args.original_model_path),
quant_path=pathlib.Path(args.quant_path),
output_path=pathlib.Path(args.output_path),
)
+18 -13
View File
@@ -64,10 +64,7 @@ from sglang.srt.layers.quantization.fp8_utils import (
requant_weight_ue8m0_inplace,
)
from sglang.srt.layers.quantization.kv_cache import BaseKVCacheMethod
from sglang.srt.layers.quantization.marlin_utils_fp8 import (
apply_fp8_marlin_linear,
prepare_fp8_layer_for_marlin,
)
from sglang.srt.layers.quantization.marlin_utils_fp8 import prepare_fp8_layer_for_marlin
from sglang.srt.layers.quantization.unquant import (
UnquantizedFusedMoEMethod,
UnquantizedLinearMethod,
@@ -707,7 +704,7 @@ class Fp8LinearMethod(LinearMethodBase):
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
if self.use_marlin:
return apply_fp8_marlin_linear(
return torch.ops.sglang.apply_fp8_marlin_linear(
input=x,
weight=layer.weight,
weight_scale=layer.weight_scale,
@@ -1077,15 +1074,23 @@ class Fp8MoEMethod(FusedMoEMethodBase):
w2_weight_scale, requires_grad=False
)
layer.w2_input_scale = None
if _use_aiter:
if _use_aiter:
# add this section for MI300
# Pre-shuffle weights
layer.w13_weight.data = shuffle_weight(
layer.w13_weight.contiguous(), (16, 16)
)
layer.w2_weight.data = shuffle_weight(
layer.w2_weight.contiguous(), (16, 16)
)
elif _use_aiter:
# Pre-shuffle weights
t = shuffle_weight(layer.w13_weight, (16, 16))
layer.w13_weight.copy_(t)
del t
t = shuffle_weight(layer.w2_weight, (16, 16))
layer.w2_weight.copy_(t)
del t
layer.w13_weight.data = shuffle_weight(
layer.w13_weight.contiguous(), (16, 16)
)
layer.w2_weight.data = shuffle_weight(
layer.w2_weight.contiguous(), (16, 16)
)
elif _is_cpu:
assert (
_is_cpu_amx_available