[AMD] Fuse RMSNorm + FP8 per-token quant for GLM-4.7-FP8 (#21403)
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
aeeff58cd4
commit
7e4e1dcd7a
@@ -80,14 +80,68 @@ _is_gfx95_supported = is_gfx95_supported()
|
|||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
_use_ag_after_qlora = envs.SGLANG_USE_AG_AFTER_QLORA.get()
|
_use_ag_after_qlora = envs.SGLANG_USE_AG_AFTER_QLORA.get()
|
||||||
|
|
||||||
if _use_aiter and _is_gfx95_supported:
|
if _use_aiter:
|
||||||
from aiter.ops.triton.fused_fp8_quant import fused_rms_fp8_group_quant
|
from aiter.ops.rmsnorm import add_rmsnorm_quant as _aiter_add_rmsnorm_quant
|
||||||
|
from aiter.ops.rmsnorm import rmsnorm_quant as _aiter_rmsnorm_quant
|
||||||
|
|
||||||
from sglang.srt.layers.quantization.rocm_mxfp4_utils import fused_rms_mxfp4_quant
|
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype as _aiter_fp8_dtype
|
||||||
|
|
||||||
|
if _is_gfx95_supported:
|
||||||
|
from aiter.ops.triton.fused_fp8_quant import fused_rms_fp8_group_quant
|
||||||
|
|
||||||
|
from sglang.srt.layers.quantization.rocm_mxfp4_utils import (
|
||||||
|
fused_rms_mxfp4_quant,
|
||||||
|
)
|
||||||
elif _is_npu:
|
elif _is_npu:
|
||||||
from sglang.srt.hardware_backend.npu.cmo import prepare_weight_cache
|
from sglang.srt.hardware_backend.npu.cmo import prepare_weight_cache
|
||||||
|
|
||||||
|
|
||||||
|
def _fused_rmsnorm_fp8_per_token_quant(
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
weight: torch.Tensor,
|
||||||
|
epsilon: float,
|
||||||
|
residual: Optional[torch.Tensor] = None,
|
||||||
|
):
|
||||||
|
"""Fused (optional residual-add +) RMSNorm + FP8 per-token quantization.
|
||||||
|
|
||||||
|
Only used with the aiter (ROCm) backend.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
residual: if provided, computes hidden_states + residual before RMSNorm
|
||||||
|
and returns updated residual_out as second element.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
If residual is None: (out_fp8, scale)
|
||||||
|
If residual provided: ((out_fp8, scale), residual_out)
|
||||||
|
"""
|
||||||
|
M, N = hidden_states.shape
|
||||||
|
out_fp8 = torch.empty((M, N), dtype=_aiter_fp8_dtype, device=hidden_states.device)
|
||||||
|
scale = torch.empty(M, dtype=torch.float32, device=hidden_states.device)
|
||||||
|
if residual is not None:
|
||||||
|
residual_out = torch.empty_like(hidden_states)
|
||||||
|
_aiter_add_rmsnorm_quant(
|
||||||
|
out_fp8,
|
||||||
|
hidden_states,
|
||||||
|
residual,
|
||||||
|
residual_out,
|
||||||
|
scale,
|
||||||
|
weight,
|
||||||
|
epsilon,
|
||||||
|
0, # group_size=0 → per-token
|
||||||
|
)
|
||||||
|
return (out_fp8, scale.unsqueeze(1)), residual_out
|
||||||
|
else:
|
||||||
|
_aiter_rmsnorm_quant(
|
||||||
|
out_fp8,
|
||||||
|
hidden_states,
|
||||||
|
scale,
|
||||||
|
weight,
|
||||||
|
epsilon,
|
||||||
|
0, # group_size=0 → per-token
|
||||||
|
)
|
||||||
|
return (out_fp8, scale.unsqueeze(1))
|
||||||
|
|
||||||
|
|
||||||
# TODO: According to the discussion in https://github.com/flashinfer-ai/flashinfer/issues/1223#issuecomment-3047256465
|
# TODO: According to the discussion in https://github.com/flashinfer-ai/flashinfer/issues/1223#issuecomment-3047256465
|
||||||
# We set the max token num to 128 for allreduce fusion with min-latency case(use_oneshot=True).
|
# We set the max token num to 128 for allreduce fusion with min-latency case(use_oneshot=True).
|
||||||
FUSE_ALLREDUCE_MAX_BATCH_SIZE = 2048
|
FUSE_ALLREDUCE_MAX_BATCH_SIZE = 2048
|
||||||
@@ -147,7 +201,6 @@ class ScatterMode(Enum):
|
|||||||
|
|
||||||
|
|
||||||
class AttentionInputs:
|
class AttentionInputs:
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
@@ -311,8 +364,8 @@ class LayerScatterModes:
|
|||||||
if context.is_layer_sparse:
|
if context.is_layer_sparse:
|
||||||
return (
|
return (
|
||||||
ScatterMode.SCATTERED
|
ScatterMode.SCATTERED
|
||||||
|
# Token dispatch/combine will be handled outside of LayerCommunicator for these modes.
|
||||||
if (
|
if (
|
||||||
# Token dispatch/combine will be handled outside of LayerCommunicator for these modes.
|
|
||||||
not get_moe_a2a_backend().is_none()
|
not get_moe_a2a_backend().is_none()
|
||||||
or should_use_flashinfer_cutlass_moe_fp4_allgather()
|
or should_use_flashinfer_cutlass_moe_fp4_allgather()
|
||||||
)
|
)
|
||||||
@@ -484,7 +537,7 @@ class LayerCommunicator:
|
|||||||
None,
|
None,
|
||||||
None,
|
None,
|
||||||
)
|
)
|
||||||
elif _use_aiter and _is_gfx95_supported and ("fp8" in quant_format):
|
elif _use_aiter and _is_gfx95_supported and (quant_format == "fp8"):
|
||||||
# aiter (ROCm gfx95) fused RMSNorm + FP8 group quant.
|
# aiter (ROCm gfx95) fused RMSNorm + FP8 group quant.
|
||||||
# When NSA is active, also preserve the unquantized bf16
|
# When NSA is active, also preserve the unquantized bf16
|
||||||
# output as a 3-tuple (fp8, scale, bf16) so the NSA
|
# output as a 3-tuple (fp8, scale, bf16) so the NSA
|
||||||
@@ -509,10 +562,16 @@ class LayerCommunicator:
|
|||||||
_unq_bf16,
|
_unq_bf16,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
elif _use_aiter and (quant_format == "fp8_per_token"):
|
||||||
|
hidden_states = _fused_rmsnorm_fp8_per_token_quant(
|
||||||
|
hidden_states,
|
||||||
|
self.input_layernorm.weight.data,
|
||||||
|
self.input_layernorm.variance_epsilon,
|
||||||
|
)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
hidden_states = self.input_layernorm(hidden_states)
|
hidden_states = self.input_layernorm(hidden_states)
|
||||||
else:
|
else:
|
||||||
|
|
||||||
if _use_aiter and _is_gfx95_supported and ("mxfp4" in quant_format):
|
if _use_aiter and _is_gfx95_supported and ("mxfp4" in quant_format):
|
||||||
hidden_states, *_, residual = fused_rms_mxfp4_quant(
|
hidden_states, *_, residual = fused_rms_mxfp4_quant(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
@@ -523,7 +582,7 @@ class LayerCommunicator:
|
|||||||
None,
|
None,
|
||||||
residual,
|
residual,
|
||||||
)
|
)
|
||||||
elif _use_aiter and _is_gfx95_supported and ("fp8" in quant_format):
|
elif _use_aiter and _is_gfx95_supported and (quant_format == "fp8"):
|
||||||
# aiter (ROCm gfx95) fused RMSNorm + FP8 group quant
|
# aiter (ROCm gfx95) fused RMSNorm + FP8 group quant
|
||||||
# with residual addition. When NSA is active, pack
|
# with residual addition. When NSA is active, pack
|
||||||
# the unquantized bf16 as a 3-tuple (fp8, scale, bf16).
|
# the unquantized bf16 as a 3-tuple (fp8, scale, bf16).
|
||||||
@@ -548,6 +607,15 @@ class LayerCommunicator:
|
|||||||
hidden_states[1],
|
hidden_states[1],
|
||||||
_unq_bf16,
|
_unq_bf16,
|
||||||
)
|
)
|
||||||
|
elif _use_aiter and (quant_format == "fp8_per_token"):
|
||||||
|
if post_residual_addition is not None:
|
||||||
|
residual = residual + post_residual_addition
|
||||||
|
hidden_states, residual = _fused_rmsnorm_fp8_per_token_quant(
|
||||||
|
hidden_states,
|
||||||
|
self.input_layernorm.weight.data,
|
||||||
|
self.input_layernorm.variance_epsilon,
|
||||||
|
residual=residual,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
hidden_states, residual = self.input_layernorm(
|
hidden_states, residual = self.input_layernorm(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ from __future__ import annotations
|
|||||||
import logging
|
import logging
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple
|
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -1608,7 +1608,7 @@ def can_auto_enable_marlin_fp8() -> bool:
|
|||||||
|
|
||||||
|
|
||||||
def apply_fp8_ptpc_linear(
|
def apply_fp8_ptpc_linear(
|
||||||
input: torch.Tensor,
|
input: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]],
|
||||||
weight: torch.Tensor,
|
weight: torch.Tensor,
|
||||||
weight_scale: torch.Tensor,
|
weight_scale: torch.Tensor,
|
||||||
input_scale: Optional[torch.Tensor] = None,
|
input_scale: Optional[torch.Tensor] = None,
|
||||||
@@ -1619,6 +1619,19 @@ def apply_fp8_ptpc_linear(
|
|||||||
pad_output: Optional[bool] = None,
|
pad_output: Optional[bool] = None,
|
||||||
compressed_tensor_quant: bool = False,
|
compressed_tensor_quant: bool = False,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
"""FP8 per-token per-channel linear. Only used with the aiter (ROCm) backend."""
|
||||||
|
# Handle pre-quantized (fp8_tensor, scale) tuple from fused RMSNorm+Quant
|
||||||
|
if isinstance(input, tuple):
|
||||||
|
q_input, x_scale = input
|
||||||
|
q_input = q_input.view(-1, q_input.shape[-1])
|
||||||
|
output_shape = [*q_input.shape[:-1], weight.shape[0]]
|
||||||
|
output = aiter.gemm_a8w8_bpreshuffle(
|
||||||
|
q_input, weight, x_scale, weight_scale, None, torch.bfloat16
|
||||||
|
)
|
||||||
|
if bias is not None:
|
||||||
|
output = output + bias
|
||||||
|
return output.view(*output_shape)
|
||||||
|
|
||||||
# View input as 2D matrix for fp8 methods
|
# View input as 2D matrix for fp8 methods
|
||||||
input_2d = input.view(-1, input.shape[-1])
|
input_2d = input.view(-1, input.shape[-1])
|
||||||
|
|
||||||
|
|||||||
@@ -289,7 +289,9 @@ class Glm4MoeAttention(nn.Module):
|
|||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
):
|
):
|
||||||
if hidden_states.shape[0] == 0:
|
# hidden_states can be a (fp8_tensor, scale) tuple from fused RMSNorm+Quant
|
||||||
|
hs = hidden_states[0] if isinstance(hidden_states, tuple) else hidden_states
|
||||||
|
if hs.shape[0] == 0:
|
||||||
return hidden_states, forward_batch, None
|
return hidden_states, forward_batch, None
|
||||||
qkv, _ = self.qkv_proj(hidden_states)
|
qkv, _ = self.qkv_proj(hidden_states)
|
||||||
|
|
||||||
@@ -865,6 +867,51 @@ class Glm4MoeDecoderLayer(nn.Module):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Detect if QKV uses aiter FP8 per-token quant so we can fuse
|
||||||
|
# RMSNorm + FP8 quant into a single kernel in prepare_attn
|
||||||
|
self.attn_quant_format = ""
|
||||||
|
self._detect_attn_quant_format()
|
||||||
|
|
||||||
|
def _detect_fp8_per_token_quant(self, linear_layer, label: str) -> str:
|
||||||
|
"""Check if a linear layer uses aiter FP8 per-token quantization."""
|
||||||
|
from sglang.srt.utils import get_bool_env_var, is_hip
|
||||||
|
|
||||||
|
if not (get_bool_env_var("SGLANG_USE_AITER") and is_hip()):
|
||||||
|
return ""
|
||||||
|
if not hasattr(linear_layer, "quant_method"):
|
||||||
|
return ""
|
||||||
|
scheme = getattr(linear_layer, "scheme", None) or getattr(
|
||||||
|
linear_layer.quant_method, "scheme", None
|
||||||
|
)
|
||||||
|
if scheme is not None:
|
||||||
|
from compressed_tensors.quantization import QuantizationStrategy
|
||||||
|
|
||||||
|
from sglang.srt.layers.quantization.compressed_tensors.schemes.compressed_tensors_w8a8_fp8 import (
|
||||||
|
CompressedTensorsW8A8Fp8,
|
||||||
|
)
|
||||||
|
|
||||||
|
if (
|
||||||
|
isinstance(scheme, CompressedTensorsW8A8Fp8)
|
||||||
|
and scheme.strategy == QuantizationStrategy.CHANNEL
|
||||||
|
):
|
||||||
|
logger.info(
|
||||||
|
"layer_%d Fused RMSNorm+Quant %s: ENABLED (fp8_per_token)",
|
||||||
|
self.layer_id,
|
||||||
|
label,
|
||||||
|
)
|
||||||
|
return "fp8_per_token"
|
||||||
|
logger.info(
|
||||||
|
"layer_%d Fused RMSNorm+Quant %s: skipped",
|
||||||
|
self.layer_id,
|
||||||
|
label,
|
||||||
|
)
|
||||||
|
return ""
|
||||||
|
|
||||||
|
def _detect_attn_quant_format(self):
|
||||||
|
self.attn_quant_format = self._detect_fp8_per_token_quant(
|
||||||
|
self.self_attn.qkv_proj, "attn"
|
||||||
|
)
|
||||||
|
|
||||||
def _is_layer_sparse(self, layer_id: int, is_nextn: bool) -> bool:
|
def _is_layer_sparse(self, layer_id: int, is_nextn: bool) -> bool:
|
||||||
return is_nextn or (
|
return is_nextn or (
|
||||||
self.config.n_routed_experts is not None
|
self.config.n_routed_experts is not None
|
||||||
@@ -880,7 +927,10 @@ class Glm4MoeDecoderLayer(nn.Module):
|
|||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|
||||||
hidden_states, residual = self.layer_communicator.prepare_attn(
|
hidden_states, residual = self.layer_communicator.prepare_attn(
|
||||||
hidden_states, residual, forward_batch
|
hidden_states,
|
||||||
|
residual,
|
||||||
|
forward_batch,
|
||||||
|
quant_format=self.attn_quant_format,
|
||||||
)
|
)
|
||||||
|
|
||||||
hidden_states = self.self_attn(
|
hidden_states = self.self_attn(
|
||||||
@@ -927,7 +977,12 @@ class Glm4MoeDecoderLayer(nn.Module):
|
|||||||
tbo_subbatch_index: Optional[int] = None,
|
tbo_subbatch_index: Optional[int] = None,
|
||||||
):
|
):
|
||||||
state.hidden_states_after_comm_pre_attn, state.residual_after_input_ln = (
|
state.hidden_states_after_comm_pre_attn, state.residual_after_input_ln = (
|
||||||
self.layer_communicator.prepare_attn(hidden_states, residual, forward_batch)
|
self.layer_communicator.prepare_attn(
|
||||||
|
hidden_states,
|
||||||
|
residual,
|
||||||
|
forward_batch,
|
||||||
|
quant_format=self.attn_quant_format,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
state.update(
|
state.update(
|
||||||
dict(
|
dict(
|
||||||
|
|||||||
Reference in New Issue
Block a user