[NPU][Diffusion] FA MXFP8 and modelslim w4a4f8 and w8a8f8 support for Wan2.2 and FLUX (#39438)

This commit is contained in:
Arseniy Mironov
2026-09-17 13:13:05 +03:00
committed by GitHub
parent 6ca866ea29
commit a1b4ec02ae
9 changed files with 514 additions and 68 deletions
@@ -187,6 +187,16 @@ description: "Configure SGLang diffusion behavior with environment variables."
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>false</td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>false</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Experimental opt-in for fused W8A8 FP8 GEMM in diffusion weight-only FP8 linears. When disabled, FP8 weights are dequantized to the compute dtype before matmul. Enabling this dynamically quantizes activations to FP8 and may change output quality.</td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Experimental opt-in for fused W8A8 FP8 GEMM in diffusion weight-only FP8 linears. When disabled, FP8 weights are dequantized to the compute dtype before matmul. Enabling this dynamically quantizes activations to FP8 and may change output quality.</td>
</tr> </tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_DIFFUSION_ENABLE_MXFP8_ATTENTION</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>false</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Enable Ascend MXFP8 FA for supported non-causal self-attention layers. Applies to online <code>MXFP8Config</code> and offline ModelSlim <code>W8A8_MXFP8</code> checkpoints. Unsupported calls continue to use the regular attention path.</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_DIFFUSION_MXFP8_FA_HEAD_CHUNK_SIZE</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>4</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Maximum number of attention heads processed by each Ascend MXFP8 FA call. Smaller chunks may improve performance for large workloads but add kernel launches; the optimal value depends on the model and input shape. Set to 0 to disable head splitting.</td>
</tr>
</tbody> </tbody>
</table> </table>
@@ -30,6 +30,10 @@ class WanVideoArchConfig(DiTArchConfig):
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$": r"blocks.\1.ffn.fc_in.\2", r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2", r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
r"^blocks\.(\d+)\.norm2\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2", r"^blocks\.(\d+)\.norm2\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2",
r"^blocks\.(\d+)\.attn1\.to_q_rot$": r"blocks.\1.q_rot",
r"^blocks\.(\d+)\.attn1\.to_k_rot$": r"blocks.\1.k_rot",
r"^blocks\.(\d+)\.attn1\.q_rot$": r"blocks.\1.q_rot",
r"^blocks\.(\d+)\.attn1\.k_rot$": r"blocks.\1.k_rot",
} }
) )
+15
View File
@@ -87,6 +87,7 @@ if TYPE_CHECKING:
SGLANG_LINGBOT_LAZY_VAE_ENCODE_BLACK_FRAMES: int | None = None SGLANG_LINGBOT_LAZY_VAE_ENCODE_BLACK_FRAMES: int | None = None
SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND: str | None = None SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND: str | None = None
SGLANG_DIFFUSION_ENABLE_W8A8_FP8_GEMM: bool = False SGLANG_DIFFUSION_ENABLE_W8A8_FP8_GEMM: bool = False
SGLANG_DIFFUSION_MXFP8_FA_HEAD_CHUNK_SIZE: int = 4
SGLANG_DIFFUSION_FP8_WEIGHT_DEQUANT_CACHE: bool = True SGLANG_DIFFUSION_FP8_WEIGHT_DEQUANT_CACHE: bool = True
SGLANG_DIFFUSION_ENABLE_COSMOS3_STEP_MIXED_PRECISION: bool = True SGLANG_DIFFUSION_ENABLE_COSMOS3_STEP_MIXED_PRECISION: bool = True
SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_FIRST_STEPS: int = 3 SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_FIRST_STEPS: int = 3
@@ -228,6 +229,20 @@ environment_variables: dict[str, Callable[[], Any]] = {
"SGLANG_DIFFUSION_ATTENTION_BACKEND": _lazy_str( "SGLANG_DIFFUSION_ATTENTION_BACKEND": _lazy_str(
"SGLANG_DIFFUSION_ATTENTION_BACKEND" "SGLANG_DIFFUSION_ATTENTION_BACKEND"
), ),
# MXFP8 Attention quantization
# Applies to both online ``MXFP8Config`` and offline ``ModelSlimConfig`` (W8A8_MXFP8)
# Q/K/V are getting offline rotating in case of rotation matrices in quant_config
# Otherwise rotation matrix are generating online
"SGLANG_DIFFUSION_ENABLE_MXFP8_ATTENTION": _lazy_bool(
"SGLANG_DIFFUSION_ENABLE_MXFP8_ATTENTION", "false"
),
# Number of attention heads processed by each MXFP8 FA call.
# Smaller chunks can improve performance for large head counts
# The default value set to 4 is better for video generation
# For image generation task depends on image quality and the model config
"SGLANG_DIFFUSION_MXFP8_FA_HEAD_CHUNK_SIZE": _lazy_int(
"SGLANG_DIFFUSION_MXFP8_FA_HEAD_CHUNK_SIZE", 4
),
# Use dedicated multiprocess context for workers. # Use dedicated multiprocess context for workers.
# Both spawn and fork work # Both spawn and fork work
"SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD": _lazy_str( "SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD": _lazy_str(
@@ -1,21 +1,62 @@
from collections.abc import Sequence from collections.abc import Sequence
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any from itertools import pairwise
from typing import Any, ClassVar
import torch import torch
from sglang.multimodal_gen import envs
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import ( from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
AttentionBackend, AttentionBackend,
AttentionImpl, AttentionImpl,
AttentionMetadata, AttentionMetadata,
AttentionMetadataBuilder, AttentionMetadataBuilder,
) )
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum,
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__)
def resolve_mx_fa_scheme(quant_config) -> str | None:
"""Resolve the opt-in MXFP8 attention scheme for an NPU quant config."""
if (
quant_config is None
or not current_platform.is_npu()
or not envs.SGLANG_DIFFUSION_ENABLE_MXFP8_ATTENTION
):
return None
if type(quant_config).__name__ not in ("MXFP8Config", "ModelSlimConfig"):
return None
if torch.npu.get_soc_version() < 260:
logger.warning_once(
"MXFP8 attention is disabled because MXFP8 quantization is only "
"supported on Ascend 950 (A5) devices."
)
return None
required_ops = (
"npu_dynamic_mx_quant",
"npu_fused_infer_attention_score_v2",
)
missing_ops = [name for name in required_ops if not hasattr(torch.ops.npu, name)]
required_dtypes = ("float8_e4m3fn", "float8_e8m0fnu")
missing_dtypes = [name for name in required_dtypes if not hasattr(torch, name)]
if missing_ops or missing_dtypes:
missing_features = missing_ops + missing_dtypes
logger.warning_once(
"MXFP8 attention is disabled because the installed torch_npu does not "
f"provide the required APIs: {', '.join(missing_features)}. "
"Please install torch==2.10.0, torch_npu>=2.10.0.post4, and CANN>=9.1.1."
)
return None
return "MXFP8"
def _packed_boundaries( def _packed_boundaries(
cu_seqlens: torch.Tensor, cu_seqlens: torch.Tensor,
cu_seqlens_host: Sequence[int] | None, cu_seqlens_host: Sequence[int] | None,
@@ -45,7 +86,7 @@ def _packed_boundaries(
f"{name} must end at the packed token count {total_tokens}, " f"{name} must end at the packed token count {total_tokens}, "
f"got {boundaries[-1]}" f"got {boundaries[-1]}"
) )
if any(stop < start for start, stop in zip(boundaries[:-1], boundaries[1:])): if any(stop < start for start, stop in pairwise(boundaries)):
raise ValueError(f"{name} must be non-decreasing") raise ValueError(f"{name} must be non-decreasing")
return boundaries return boundaries
@@ -99,12 +140,8 @@ def fused_infer_attention_varlen(
if len(q_boundaries) != len(k_boundaries): if len(q_boundaries) != len(k_boundaries):
raise ValueError("cu_seqlens_q and cu_seqlens_k must describe the same batch") raise ValueError("cu_seqlens_q and cu_seqlens_k must describe the same batch")
q_nonempty = [ q_nonempty = [stop > start for start, stop in pairwise(q_boundaries)]
stop > start for start, stop in zip(q_boundaries[:-1], q_boundaries[1:]) k_nonempty = [stop > start for start, stop in pairwise(k_boundaries)]
]
k_nonempty = [
stop > start for start, stop in zip(k_boundaries[:-1], k_boundaries[1:])
]
if q_nonempty != k_nonempty: if q_nonempty != k_nonempty:
raise NotImplementedError( raise NotImplementedError(
"NPU packed attention does not support a sequence that is empty only " "NPU packed attention does not support a sequence that is empty only "
@@ -197,6 +234,21 @@ class AscendFABackend(AttentionBackend):
class AscendFAImpl(AttentionImpl): class AscendFAImpl(AttentionImpl):
# FA v2 uses per-token-group quantization (mode 6) for Q/K and
# per-channel-group quantization (mode 8) for V in the packed TND path.
_MXFP8_LAYOUT = "TND"
_MXFP8_QK_QUANT_AXIS = -1
_MXFP8_V_QUANT_AXIS = 0
_MXFP8_QK_QUANT_MODE = 6
_MXFP8_V_QUANT_MODE = 8
# Online Q/K rotations are deterministic CPU FP32 tensors shared
# by all backend instances and keyed by head size. Applying the same
# orthogonal matrix R preserves scores:
# (Q @ R) @ (K @ R).T = Q @ R @ R.T @ K.T = Q @ K.T.
# Offline checkpoint rotations do not use this generated-matrix cache.
_rot_matrices: ClassVar[dict[int, torch.Tensor]] = {}
def __init__( def __init__(
self, self,
num_heads: int, num_heads: int,
@@ -209,6 +261,50 @@ class AscendFAImpl(AttentionImpl):
) -> None: ) -> None:
self.causal = causal self.causal = causal
self.softmax_scale = softmax_scale self.softmax_scale = softmax_scale
quant_config = extra_impl_args.get("quant_config")
self._quant_scheme = resolve_mx_fa_scheme(quant_config)
self.use_offline_qk_rotation = (
quant_config.use_offline_qk_rotation
if hasattr(quant_config, "use_offline_qk_rotation")
else False
)
self._is_cross_attention = bool(
extra_impl_args.get("is_cross_attention", False)
)
if self._quant_scheme is not None:
self._head_size = head_size
self._mxfp8_head_chunk_size = envs.SGLANG_DIFFUSION_MXFP8_FA_HEAD_CHUNK_SIZE
self._rot_device: torch.Tensor | None = None
if not self.use_offline_qk_rotation:
self._ensure_rot_matrix(head_size)
@classmethod
def _ensure_rot_matrix(cls, head_size: int) -> None:
if head_size in cls._rot_matrices:
return
generator = torch.Generator(device="cpu")
generator.manual_seed(42)
rotation, _ = torch.linalg.qr(
torch.randn(
head_size,
head_size,
generator=generator,
device="cpu",
dtype=torch.float32,
)
)
cls._rot_matrices[head_size] = rotation
def _get_rotation(self, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
if (
self._rot_device is None
or self._rot_device.device != device
or self._rot_device.dtype != dtype
):
self._rot_device = self._rot_matrices[self._head_size].to(
device=device, dtype=dtype
)
return self._rot_device
def forward( def forward(
self, self,
@@ -218,17 +314,43 @@ class AscendFAImpl(AttentionImpl):
attn_metadata: AttentionMetadata, attn_metadata: AttentionMetadata,
return_softmax_lse: bool = False, return_softmax_lse: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
if (
self._quant_scheme == "MXFP8"
and not self.causal
and not self._is_cross_attention
and query.shape[1:3] == key.shape[1:3]
and key.shape == value.shape
and (query.shape[0] * query.shape[1]) % 64 == 0
):
batch_size, query_length, num_heads, head_size = query.shape
key_length = key.shape[1]
actual_seq_qlen = [
query_length * batch_index for batch_index in range(1, batch_size + 1)
]
actual_seq_kvlen = [
key_length * batch_index for batch_index in range(1, batch_size + 1)
]
output = self._forward_mxfp8_tnd(
query.reshape(-1, num_heads, head_size),
key.reshape(-1, key.shape[2], head_size),
value.reshape(-1, value.shape[2], head_size),
actual_seq_qlen=actual_seq_qlen,
actual_seq_kvlen=actual_seq_kvlen,
return_softmax_lse=return_softmax_lse,
)
return output.reshape(batch_size, query_length, num_heads, head_size)
mask = None mask = None
num_heads, num_key_value_heads = query.shape[2], key.shape[2] num_heads, num_key_value_heads = query.shape[2], key.shape[2]
if self.causal: if self.causal:
seq_len = query.shape[1] seq_len = query.shape[1]
mask = torch.triu( mask = torch.triu(
torch.ones(seq_len, seq_len, device=query.device), diagonal=1 torch.ones(seq_len, seq_len, device=query.device), diagonal=1
).bool() ).bool()[None]
# transpose to bs, heads, seq_len, head_dim # transpose to bs, heads, seq_len, head_dim
query = query.transpose(1, 2) query = query.transpose(1, 2)
key = key.transpose(1, 2) key = key.transpose(1, 2).contiguous()
value = value.transpose(1, 2) value = value.transpose(1, 2).contiguous()
output, lse = torch.ops.npu.npu_fused_infer_attention_score( output, lse = torch.ops.npu.npu_fused_infer_attention_score(
query, query,
key, key,
@@ -256,6 +378,30 @@ class AscendFAImpl(AttentionImpl):
cu_seqlens_host: tuple[int, ...] | None = None, cu_seqlens_host: tuple[int, ...] | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
del max_seqlen del max_seqlen
if (
self._quant_scheme == "MXFP8"
and not self.causal
and not self._is_cross_attention
and query.shape == key.shape
and key.shape == value.shape
and query.shape[0] % 64 == 0
):
boundaries = _packed_boundaries(
cu_seqlens, cu_seqlens_host, query.shape[0], "cu_seqlens"
)
actual_seq_lengths = [
stop for start, stop in pairwise(boundaries) if stop > start
]
if not actual_seq_lengths:
return torch.empty_like(query)
return self._forward_mxfp8_tnd(
query,
key,
value,
actual_seq_qlen=actual_seq_lengths,
actual_seq_kvlen=actual_seq_lengths,
)
if self.causal: if self.causal:
bounds = ( bounds = (
cu_seqlens_host cu_seqlens_host
@@ -263,7 +409,7 @@ class AscendFAImpl(AttentionImpl):
else tuple(int(item) for item in cu_seqlens.tolist()) else tuple(int(item) for item in cu_seqlens.tolist())
) )
output = torch.empty_like(query) output = torch.empty_like(query)
for start, stop in zip(bounds[:-1], bounds[1:]): for start, stop in pairwise(bounds):
if start == stop: if start == stop:
continue continue
segment = self.forward( segment = self.forward(
@@ -286,6 +432,111 @@ class AscendFAImpl(AttentionImpl):
softmax_scale=self.softmax_scale, softmax_scale=self.softmax_scale,
) )
def _forward_mxfp8_tnd(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
*,
actual_seq_qlen: Sequence[int],
actual_seq_kvlen: Sequence[int],
return_softmax_lse: bool = False,
) -> torch.Tensor:
if return_softmax_lse:
raise NotImplementedError(
"MXFP8 attention does not support returning softmax LSE"
)
logger.info_once("Using MXFP8 quantized Ascend Flash Attention.")
if not self.use_offline_qk_rotation:
rotation = self._get_rotation(query.device, query.dtype)
query = torch.matmul(query, rotation)
key = torch.matmul(key, rotation)
num_heads = query.shape[1]
num_kv_heads = key.shape[1]
if num_heads != num_kv_heads:
raise NotImplementedError("MXFP8 attention currently requires MHA")
head_chunk_size = self._mxfp8_head_chunk_size
if head_chunk_size > 0 and num_heads > head_chunk_size:
num_groups, remainder = divmod(num_heads, head_chunk_size)
head_groups = [head_chunk_size] * num_groups
if remainder:
head_groups.append(remainder)
outputs = [
self._run_mxfp8_attention(
query_chunk,
key_chunk,
value_chunk,
actual_seq_qlen=actual_seq_qlen,
actual_seq_kvlen=actual_seq_kvlen,
)
for query_chunk, key_chunk, value_chunk in zip(
query.split(head_groups, dim=1),
key.split(head_groups, dim=1),
value.split(head_groups, dim=1),
)
]
return torch.cat(outputs, dim=1)
return self._run_mxfp8_attention(
query,
key,
value,
actual_seq_qlen=actual_seq_qlen,
actual_seq_kvlen=actual_seq_kvlen,
)
def _run_mxfp8_attention(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
*,
actual_seq_qlen: Sequence[int],
actual_seq_kvlen: Sequence[int],
) -> torch.Tensor:
quant_dtype = torch.float8_e4m3fn
scale_dtype = torch.float8_e8m0fnu
query = query.contiguous()
key = key.contiguous()
value = value.contiguous()
query_fp8, query_scale = torch.ops.npu.npu_dynamic_mx_quant(
query, dst_type=quant_dtype, axis=self._MXFP8_QK_QUANT_AXIS
)
key_fp8, key_scale = torch.ops.npu.npu_dynamic_mx_quant(
key, dst_type=quant_dtype, axis=self._MXFP8_QK_QUANT_AXIS
)
value_fp8, value_scale = torch.ops.npu.npu_dynamic_mx_quant(
value, dst_type=quant_dtype, axis=self._MXFP8_V_QUANT_AXIS
)
return torch.ops.npu.npu_fused_infer_attention_score_v2(
query_fp8,
key_fp8,
value_fp8,
input_layout=self._MXFP8_LAYOUT,
num_query_heads=query.shape[1],
num_key_value_heads=key.shape[1],
softmax_scale=self.softmax_scale,
dequant_scale_query=query_scale,
dequant_scale_key=key_scale,
dequant_scale_value=value_scale,
actual_seq_qlen=actual_seq_qlen,
actual_seq_kvlen=actual_seq_kvlen,
sparse_mode=0,
query_quant_mode=self._MXFP8_QK_QUANT_MODE,
key_quant_mode=self._MXFP8_QK_QUANT_MODE,
value_quant_mode=self._MXFP8_V_QUANT_MODE,
query_dtype=quant_dtype,
key_dtype=quant_dtype,
value_dtype=quant_dtype,
dequant_scale_query_dtype=scale_dtype,
dequant_scale_key_dtype=scale_dtype,
dequant_scale_value_dtype=scale_dtype,
out_dtype=query.dtype,
)[0]
def forward_ring_kv_chunk( def forward_ring_kv_chunk(
self, self,
query: torch.Tensor, query: torch.Tensor,
@@ -144,7 +144,11 @@ class ModelSlimConfig(QuantizationConfig):
ModelSlimMXFP4Scheme, ModelSlimMXFP4Scheme,
) )
return ModelSlimMXFP4Scheme() return ModelSlimMXFP4Scheme(
quant_config=self.quant_description,
prefix=prefix,
quant_type=quant_type,
)
raise NotImplementedError( raise NotImplementedError(
f"No modelslim compatible scheme was found for layer '{layer_name}'. " f"No modelslim compatible scheme was found for layer '{layer_name}'. "
f"quant_description['{layer_name}.weight'] = '{quant_type}'" f"quant_description['{layer_name}.weight'] = '{quant_type}'"
@@ -32,9 +32,42 @@ MXFP4_BLOCK_SIZE = 32
# L1 (dual) scale groups this many L0 blocks together. # L1 (dual) scale groups this many L0 blocks together.
# L1 block covers 16 * 32 = 512 elements. # L1 block covers 16 * 32 = 512 elements.
MXFP4_DUAL_LEVEL_RATIO = 16 MXFP4_DUAL_LEVEL_RATIO = 16
MXFP4_PACK_FACTOR = 2
class ModelSlimMXFP4Scheme(ModelSlimLinearScheme): class ModelSlimMXFP4Scheme(ModelSlimLinearScheme):
def __init__(
self,
quant_config: dict,
prefix: str,
quant_type: str,
):
self.quant_config = quant_config
self.prefix = prefix
self.quant_type = quant_type
self.is_dual_scale = quant_type == "W4A4_MXFP4_DUALSCALE"
self.dual_scale_key = prefix + ".weight_dual_scale"
self.mul_scale_key = prefix + ".mul_scale"
self.legacy_mul_scale_key = prefix + ".div.mul_scale"
self.has_mul_scale = (
self.legacy_mul_scale_key in quant_config
and self.mul_scale_key in quant_config
)
self.single_level_kernel = None
if not self.is_dual_scale:
from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import (
NPUSingleLevelMXFP4OfflineLinearMethod,
)
self.single_level_kernel = NPUSingleLevelMXFP4OfflineLinearMethod()
else:
if self.is_dual_scale_key not in self.quant_config:
raise ValueError(
f"Dual-level MXFP4 quantization requires missing '{self.dual_scale_key}' in quant_config."
"Check that the model was exported with dual-level quantization."
)
def create_weights( def create_weights(
self, self,
layer: torch.nn.Module, layer: torch.nn.Module,
@@ -53,8 +86,11 @@ class ModelSlimMXFP4Scheme(ModelSlimLinearScheme):
# (npu_dtype_cast → float4_e2m1fn_x2) happens in process_weights_after_loading. # (npu_dtype_cast → float4_e2m1fn_x2) happens in process_weights_after_loading.
weight = ModelWeightParameter( weight = ModelWeightParameter(
data=torch.empty( data=torch.empty(
(output_size_per_partition, input_size_per_partition), (
dtype=torch.float8_e4m3fn, output_size_per_partition,
input_size_per_partition // MXFP4_PACK_FACTOR,
),
dtype=torch.uint8,
), ),
input_dim=1, input_dim=1,
output_dim=0, output_dim=0,
@@ -74,40 +110,51 @@ class ModelSlimMXFP4Scheme(ModelSlimLinearScheme):
weight_loader=weight_loader, weight_loader=weight_loader,
) )
layer.register_parameter("weight_scale", weight_scale) layer.register_parameter("weight_scale", weight_scale)
if self.is_dual_scale:
# L0 (coarse) scale for dual-level quantization matmul.
# Each L0 block covers MXFP4_DUAL_LEVEL_RATIO L1 blocks = 16 * 32 = 512 elements.
dual_scale_dim = scale_dim // MXFP4_DUAL_LEVEL_RATIO # in/32 / 16 = in/512
weight_dual_scale = GroupQuantScaleParameter(
data=torch.empty(
(output_size_per_partition, dual_scale_dim, 1),
dtype=torch.float32,
),
input_dim=1,
output_dim=0,
weight_loader=weight_loader,
)
weight_dual_scale.missing_param_init = "error"
layer.register_parameter("weight_dual_scale", weight_dual_scale)
# L0 (coarse) scale for dual-level quantization matmul. if self.has_mul_scale:
# Each L0 block covers MXFP4_DUAL_LEVEL_RATIO L1 blocks = 16 * 32 = 512 elements. # Smooth quant activation scale (mul_scale) from NonFusionSmoothQuantWrapper.
dual_scale_dim = scale_dim // MXFP4_DUAL_LEVEL_RATIO # in/32 / 16 = in/512 # msmodelslim exports this as `<prefix>.div.mul_scale` with shape [in].
weight_dual_scale = GroupQuantScaleParameter( # After repack, it becomes `<prefix>.mul_scale`.
data=torch.empty( # This is CRITICAL: the offline-quantized weights were calibrated with
(output_size_per_partition, dual_scale_dim, 1), # x * mul_scale applied to the activation. Omitting it causes mosaic output.
dtype=torch.float32, mul_scale = BasevLLMParameter(
), data=torch.empty(
input_dim=1, (input_size_per_partition,),
output_dim=0, dtype=torch.float32,
weight_loader=weight_loader, ),
) weight_loader=weight_loader,
layer.register_parameter("weight_dual_scale", weight_dual_scale) )
mul_scale.missing_param_init = "error"
# Smooth quant activation scale (mul_scale) from NonFusionSmoothQuantWrapper. layer.register_parameter("mul_scale", mul_scale)
# msmodelslim exports this as `<prefix>.div.mul_scale` with shape [in].
# After repack, it becomes `<prefix>.mul_scale`.
# This is CRITICAL: the offline-quantized weights were calibrated with
# x * mul_scale applied to the activation. Omitting it causes mosaic output.
mul_scale = BasevLLMParameter(
data=torch.empty(
(input_size_per_partition,),
dtype=torch.float32,
),
weight_loader=weight_loader,
)
# If mul_scale is not in the checkpoint (e.g. non-smooth-quant model
# or old repack without .div. handling), initialize to ones so that
# x * 1.0 = x (no-op). fsdp_load.py checks this attribute.
mul_scale.missing_param_init = "ones"
layer.register_parameter("mul_scale", mul_scale)
def process_weights_after_loading(self, layer: torch.nn.Module): def process_weights_after_loading(self, layer: torch.nn.Module):
if not self.is_dual_scale:
self.single_level_kernel.process_weights_after_loading(layer)
if self.has_mul_scale:
mul_scale = layer.mul_scale.data
if not mul_scale.is_npu:
mul_scale = mul_scale.to(f"npu:{torch.npu.current_device()}")
layer.mul_scale = torch.nn.Parameter(mul_scale, requires_grad=False)
layer.use_mul_scale = not torch.all(mul_scale == 1.0).item()
else:
layer.use_mul_scale = False
return
# Cast weight from fp8 container to FP4 packed format # Cast weight from fp8 container to FP4 packed format
weight = layer.weight.data weight = layer.weight.data
if not weight.is_npu: if not weight.is_npu:
@@ -127,23 +174,29 @@ class ModelSlimMXFP4Scheme(ModelSlimLinearScheme):
weight_scale = weight_scale.reshape(weight_scale.shape[0], -1, 2) weight_scale = weight_scale.reshape(weight_scale.shape[0], -1, 2)
layer.weight_scale = torch.nn.Parameter(weight_scale, requires_grad=False) layer.weight_scale = torch.nn.Parameter(weight_scale, requires_grad=False)
# Transform weight_dual_scale: [out, in/512, 1] -> [in/512, out] if self.is_dual_scale:
weight_dual_scale = layer.weight_dual_scale.data # Transform weight_dual_scale: [out, in/512, 1] -> [in/512, out]
if not weight_dual_scale.is_npu: weight_dual_scale = layer.weight_dual_scale.data
weight_dual_scale = weight_dual_scale.to( if not weight_dual_scale.is_npu:
f"npu:{torch.npu.current_device()}" weight_dual_scale = weight_dual_scale.to(
f"npu:{torch.npu.current_device()}"
)
weight_dual_scale = (
weight_dual_scale.squeeze(-1).transpose(0, 1).contiguous()
)
layer.weight_dual_scale = torch.nn.Parameter(
weight_dual_scale, requires_grad=False
) )
weight_dual_scale = weight_dual_scale.squeeze(-1).transpose(0, 1).contiguous()
layer.weight_dual_scale = torch.nn.Parameter(
weight_dual_scale, requires_grad=False
)
# Move mul_scale to NPU if present and not already there if self.has_mul_scale:
mul_scale = layer.mul_scale.data # Move mul_scale to NPU if present and not already there
if not mul_scale.is_npu: mul_scale = layer.mul_scale.data
mul_scale = mul_scale.to(f"npu:{torch.npu.current_device()}") if not mul_scale.is_npu:
layer.mul_scale = torch.nn.Parameter(mul_scale, requires_grad=False) mul_scale = mul_scale.to(f"npu:{torch.npu.current_device()}")
layer.use_mul_scale = not torch.all(mul_scale == 1.0).item() layer.mul_scale = torch.nn.Parameter(mul_scale, requires_grad=False)
layer.use_mul_scale = not torch.all(mul_scale == 1.0).item()
else:
layer.use_mul_scale = False
def apply_weights( def apply_weights(
self, self,
@@ -151,6 +204,10 @@ class ModelSlimMXFP4Scheme(ModelSlimLinearScheme):
x: torch.Tensor, x: torch.Tensor,
bias: Optional[torch.Tensor] = None, bias: Optional[torch.Tensor] = None,
) -> torch.Tensor: ) -> torch.Tensor:
if not self.is_dual_scale:
if getattr(layer, "use_mul_scale", False):
x = x * layer.mul_scale.to(x.dtype)
return self.single_level_kernel.apply(layer, x, bias)
original_dtype = x.dtype original_dtype = x.dtype
if original_dtype not in (torch.float16, torch.bfloat16): if original_dtype not in (torch.float16, torch.bfloat16):
@@ -165,7 +222,7 @@ class ModelSlimMXFP4Scheme(ModelSlimLinearScheme):
# The offline-quantized weights were calibrated under x * mul_scale, # The offline-quantized weights were calibrated under x * mul_scale,
# so we MUST apply it here for scale alignment. # so we MUST apply it here for scale alignment.
mul_scale = layer.mul_scale mul_scale = layer.mul_scale
if getattr(layer, "use_mul_scale", True): if getattr(layer, "use_mul_scale", False):
x_2d = x_2d * mul_scale.to(x_2d.dtype) x_2d = x_2d * mul_scale.to(x_2d.dtype)
# Dual-level MXFP4 activation quantization # Dual-level MXFP4 activation quantization
@@ -201,7 +201,20 @@ def _maybe_dequantize_fp8(
scale_key = target_param_name.rsplit(".", 1)[0] + ".weight_scale" scale_key = target_param_name.rsplit(".", 1)[0] + ".weight_scale"
scale_tensor = param_sd.get(scale_key) scale_tensor = param_sd.get(scale_key)
if scale_tensor is not None: if scale_tensor is not None:
full_tensor = full_tensor.to(torch.float32) * scale_tensor.float() if (
scale_tensor.dtype == torch.uint8
and full_tensor.ndim == scale_tensor.ndim
and full_tensor.shape[:1] == scale_tensor.shape[:1]
and full_tensor.shape[-1] == scale_tensor.shape[-1] * 32
):
scale = torch.exp2(scale_tensor.float() - 127.0)
blocked_shape = (*full_tensor.shape[:-1], scale_tensor.shape[-1], 32)
full_tensor = (
full_tensor.float().reshape(blocked_shape) * scale.unsqueeze(-1)
).reshape(full_tensor.shape)
else:
full_tensor = full_tensor.to(torch.float32) * scale_tensor.float()
logger.debug( logger.debug(
"Auto-dequantized FP8 weight %s using %s", "Auto-dequantized FP8 weight %s using %s",
target_param_name, target_param_name,
@@ -89,7 +89,10 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
) )
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
from sglang.multimodal_gen.runtime.models.dits.common import get_qkv_projections from sglang.multimodal_gen.runtime.models.dits.common import get_qkv_projections
from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum,
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__) # pylint: disable=invalid-name logger = init_logger(__name__) # pylint: disable=invalid-name
@@ -637,12 +640,43 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
prefix=f"{prefix}.to_add_out" if prefix else "", prefix=f"{prefix}.to_add_out" if prefix else "",
) )
# TODO Need to create mxfp8 attention scheme and port the code below
from sglang.multimodal_gen import envs
quant_description = getattr(quant_config, "quant_description", {})
self.use_offline_qk_rotation = (
quant_description.get(f"{prefix}.q_rot") == "FLOAT"
and quant_description.get(f"{prefix}.k_rot") == "FLOAT"
and envs.SGLANG_DIFFUSION_ENABLE_MXFP8_ATTENTION
)
if self.use_offline_qk_rotation:
self.register_buffer(
"q_rot",
torch.empty(
self.head_dim,
self.head_dim,
dtype=torch.bfloat16,
),
persistent=True,
)
self.register_buffer(
"k_rot",
torch.empty(
self.head_dim,
self.head_dim,
dtype=torch.bfloat16,
),
persistent=True,
)
quant_config.use_offline_qk_rotation = True
self.attn = USPAttention( self.attn = USPAttention(
num_heads=self.local_heads if self.shard_qkv else num_heads, num_heads=self.local_heads if self.shard_qkv else num_heads,
head_size=self.head_dim, head_size=self.head_dim,
dropout_rate=0, dropout_rate=0,
softmax_scale=None, softmax_scale=None,
causal=False, causal=False,
quant_config=quant_config,
) )
def forward( def forward(
@@ -737,6 +771,19 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
allow_inplace=True, allow_inplace=True,
) )
# Offline rotations belong to the MXFP8 FA contract.
if (
self.use_offline_qk_rotation
and self.attn.backend is AttentionBackendEnum.FA
and query.shape[1:3] == key.shape[1:3]
and key.shape == value.shape
and (query.shape[0] * query.shape[1]) % 64 == 0
):
self.q_rot = self.q_rot.to(device=query.device, dtype=query.dtype)
self.k_rot = self.k_rot.to(device=key.device, dtype=key.dtype)
query = torch.matmul(query, self.q_rot)
key = torch.matmul(key, self.k_rot)
x = self.attn( x = self.attn(
query, query,
key, key,
@@ -492,6 +492,11 @@ class WanTransformerBlock(nn.Module):
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("to_out", prefix), prefix=add_prefix("to_out", prefix),
) )
self.hidden_dim = dim
self.num_attention_heads = num_heads
self.dim_head = dim // num_heads
self.use_offline_qk_rotation = False
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)
self_attn_backends = supported_attention_backends self_attn_backends = supported_attention_backends
@@ -509,6 +514,36 @@ class WanTransformerBlock(nn.Module):
prefix=add_prefix("attn1", prefix), prefix=add_prefix("attn1", prefix),
) )
else: else:
# TODO Need to create mxfp8 attention scheme and port the code below
from sglang.multimodal_gen import envs
quant_description = getattr(quant_config, "quant_description", {})
self.use_offline_qk_rotation = (
quant_description.get(f"{prefix}.attn1.q_rot") == "FLOAT"
and quant_description.get(f"{prefix}.attn1.k_rot") == "FLOAT"
and envs.SGLANG_DIFFUSION_ENABLE_MXFP8_ATTENTION
)
if self.use_offline_qk_rotation:
self.register_buffer(
"q_rot",
torch.empty(
self.dim_head,
self.dim_head,
dtype=torch.bfloat16,
),
persistent=True,
)
self.register_buffer(
"k_rot",
torch.empty(
self.dim_head,
self.dim_head,
dtype=torch.bfloat16,
),
persistent=True,
)
quant_config.use_offline_qk_rotation = True
self.attn1 = USPAttention( self.attn1 = USPAttention(
num_heads=self.local_num_heads, num_heads=self.local_num_heads,
head_size=dim // num_heads, head_size=dim // num_heads,
@@ -519,9 +554,6 @@ class WanTransformerBlock(nn.Module):
is_cross_attention=False, is_cross_attention=False,
) )
self.hidden_dim = dim
self.num_attention_heads = num_heads
self.dim_head = dim // num_heads
if qk_norm == "rms_norm": if qk_norm == "rms_norm":
self.norm_q = RMSNorm(self.dim_head, eps=eps) self.norm_q = RMSNorm(self.dim_head, eps=eps)
self.norm_k = RMSNorm(self.dim_head, eps=eps) self.norm_k = RMSNorm(self.dim_head, eps=eps)
@@ -683,6 +715,19 @@ class WanTransformerBlock(nn.Module):
_apply_rotary_emb(query, cos, sin, is_neox_style=False), _apply_rotary_emb(query, cos, sin, is_neox_style=False),
_apply_rotary_emb(key, cos, sin, is_neox_style=False), _apply_rotary_emb(key, cos, sin, is_neox_style=False),
) )
if (
self.use_offline_qk_rotation
and self.attn1.backend is AttentionBackendEnum.FA
and query.shape[1:3] == key.shape[1:3]
and key.shape == value.shape
and (query.shape[0] * query.shape[1]) % 64 == 0
):
self.q_rot = self.q_rot.to(device=query.device, dtype=query.dtype)
self.k_rot = self.k_rot.to(device=key.device, dtype=key.dtype)
query = torch.matmul(query, self.q_rot)
key = torch.matmul(key, self.k_rot)
attn_output = self.attn1(query, key, value) attn_output = self.attn1(query, key, value)
attn_output = attn_output.flatten(2) attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output) attn_output, _ = self.to_out(attn_output)