fix(xpu): enable compressed-tensors FP8 W8A8 on XPU (RedHatAI FP8-dynamic models) (#33057)

This commit is contained in:
vikram singh shekhawat
2026-08-24 12:53:48 +08:00
committed by GitHub
parent b498efce52
commit f98b60de80
6 changed files with 42 additions and 10 deletions
@@ -35,6 +35,7 @@ from sglang.srt.utils import (
is_cuda,
is_hip,
is_musa,
is_xpu,
log_info_on_rank0,
)
from sglang.srt.utils.custom_op import register_custom_op
@@ -44,6 +45,7 @@ _is_hip = is_hip()
_is_cuda = is_cuda()
_is_cpu = is_cpu()
_is_musa = is_musa()
_is_xpu = is_xpu()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
if _is_cuda:
@@ -54,6 +56,8 @@ if _is_cuda:
from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
per_tensor_quant_fp8 as sgl_per_tensor_quant_fp8,
)
elif _is_xpu:
from sgl_kernel import sgl_per_tensor_quant_fp8, sgl_per_token_quant_fp8
if _is_musa:
from sgl_kernel import sgl_per_token_quant_fp8
@@ -69,11 +69,12 @@ from sglang.srt.layers.quantization.unquant import (
UnquantizedFusedMoEMethod,
UnquantizedLinearMethod,
)
from sglang.srt.utils import is_cuda, is_hip, is_npu
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
_is_cuda = is_cuda()
_is_npu = is_npu()
_is_hip = is_hip()
_is_xpu = is_xpu()
if TYPE_CHECKING:
from sglang.srt.layers.moe.token_dispatcher import (
@@ -395,6 +396,14 @@ class CompressedTensorsConfig(QuantizationConfig):
return []
def _check_scheme_supported(self, min_capability: int, error: bool = True) -> bool:
if _is_xpu:
if error:
raise RuntimeError(
f"Quantization scheme requiring compute capability "
f"{min_capability} is not supported on XPU."
)
return False
capability_tuple = DeviceCapability(*torch.cuda.get_device_capability())
if capability_tuple is not None:
@@ -700,9 +709,12 @@ class CompressedTensorsConfig(QuantizationConfig):
)
if self._is_fp8_w8a8(weight_quant, input_quant):
is_fp8_w8a8_supported = self._check_scheme_supported(
CompressedTensorsW8A8Fp8.get_min_capability(), error=False
)
if _is_xpu:
is_fp8_w8a8_supported = True
else:
is_fp8_w8a8_supported = self._check_scheme_supported(
CompressedTensorsW8A8Fp8.get_min_capability(), error=False
)
if is_fp8_w8a8_supported:
return CompressedTensorsW8A8Fp8(
weight_quant=weight_quant,
@@ -939,7 +951,13 @@ class CompressedTensorsConfig(QuantizationConfig):
# Raise error if device does not support the scheme
# (e.g. fp8 needs ada lovelace)
# Note: NPU devices do not support min_capability function
if not _is_npu:
if _is_xpu:
if not isinstance(scheme, CompressedTensorsW8A8Fp8):
raise RuntimeError(
f"{scheme.__class__.__name__} is not supported on XPU "
"(no XPU kernel implementation)."
)
elif not _is_npu:
self._check_scheme_supported(scheme.get_min_capability())
logger.debug("Using scheme: %s for %s", scheme.__class__.__name__, layer_name)
return scheme
+4 -1
View File
@@ -2422,7 +2422,10 @@ class Fp8MoEMethod(FusedMoEMethodBase):
if quant_info is not None:
return self.runner.run(dispatch_output, quant_info)
if use_intel_xpu_backend():
if use_intel_xpu_backend() and not (
getattr(self, "runner", None) is not None
and self.runner.runner_backend.is_triton()
):
# sgl-kernel-xpu path
from sgl_kernel import fused_experts
@@ -51,6 +51,7 @@ from sglang.srt.utils import (
is_sm90_supported,
is_sm100_supported,
is_sm120_supported,
is_xpu,
offloader,
)
from sglang.srt.utils.custom_op import register_custom_op
@@ -59,6 +60,7 @@ logger = logging.getLogger(__name__)
_is_hip = is_hip()
_is_cuda = is_cuda()
_is_xpu = is_xpu()
_is_fp8_fnuz = is_fp8_fnuz()
_is_sm90_supported = is_sm90_supported()
_is_sm100_supported = is_sm100_supported()
@@ -1858,7 +1860,9 @@ def apply_fp8_linear(
elif compressed_tensor_quant:
# Maybe apply padding to output, see comment in __init__
num_token_padding = output_padding
if channelwise_cutlass:
if channelwise_cutlass or (_is_xpu and weight_scale.numel() == weight.shape[1]):
# On XPU, sgl-kernel-xpu's native quant kernels require output_q
# to exactly match input's shape; padded output isn't supported.
num_token_padding = None
# For static per-tensor activation scales when using inductor compiler,
# use pure PyTorch ops instead of the opaque sgl_kernel quant kernel.
+3 -1
View File
@@ -716,7 +716,9 @@ class InternVLChatModel(nn.Module):
ckpt_gate_proj_name="gate_proj",
ckpt_down_proj_name="down_proj",
ckpt_up_proj_name="up_proj",
num_experts=self.config.num_experts,
# InternVLChatConfig has no top-level num_experts; the MoE config lives on
# the nested llm_config (the Qwen3MoE text backbone).
num_experts=self.config.llm_config.num_experts,
)
elif "Qwen3ForCausalLM" in self.config.llm_config.architectures:
stacked_params_mapping = [
+3 -2
View File
@@ -756,7 +756,7 @@ class RuntimeContext:
self.forward = ForwardFlags()
def get_stream(self, name: str) -> Any:
"""Named process-level CUDA side stream: get-or-create, shared by
"""Named process-level side stream: get-or-create, shared by
name (the keyed-lazy pattern of the persistent buffers). Creation is
a driver call that must stay outside cuda-graph capture — call sites
lease their stream at init/warmup time."""
@@ -764,7 +764,8 @@ class RuntimeContext:
if stream is None:
import torch
stream = torch.cuda.Stream()
device = self._server_args.device if self._server_args else "cuda"
stream = torch.get_device_module(device).Stream()
self.resources.streams[name] = stream
return stream