[XPU] add gptq(int4) support (#37844)

Signed-off-by: Kunshang Ji <[email protected]>
This commit is contained in:
Kunshang Ji
2026-05-19 11:17:09 +08:00
committed by GitHub
parent 8f16c4a5c0
commit 36dcaf25d8
4 changed files with 55 additions and 13 deletions
+3 -1
View File
@@ -39,7 +39,9 @@ steps:
python3 examples/basic/offline_inference/generate.py --model nvidia/Llama-3.1-8B-Instruct-FP8 --block-size 64 --enforce-eager --quantization modelopt --kv-cache-dtype fp8 --attention-backend TRITON_ATTN --max-model-len 4096 &&
python3 examples/basic/offline_inference/generate.py --model superjob/Qwen3-4B-Instruct-2507-GPTQ-Int4 --block-size 64 --enforce-eager --max-model-len 8192 &&
python3 examples/basic/offline_inference/generate.py --model ibm-research/PowerMoE-3b --block-size 64 --enforce-eager -tp 2 &&
python3 examples/basic/offline_inference/generate.py --model ibm-research/PowerMoE-3b --block-size 64 --enforce-eager -tp 2 --enable-expert-parallel'
python3 examples/basic/offline_inference/generate.py --model ibm-research/PowerMoE-3b --block-size 64 --enforce-eager -tp 2 --enable-expert-parallel &&
python3 examples/basic/offline_inference/generate.py --model superjob/Qwen3-4B-Instruct-2507-GPTQ-Int4 --max-model-len 8192
'
- label: "XPU V1 test"
depends_on:
- image-build-xpu
@@ -50,8 +50,8 @@ class MPLinearKernel(ABC):
assert w_zp_param_name is not None
if c.has_g_idx:
assert w_gidx_param_name is not None
self.w_zp_name = w_zp_param_name
self.w_gidx_name = w_gidx_param_name
self.w_zp_name: str | None = w_zp_param_name
self.w_gidx_name: str | None = w_gidx_param_name
@abstractmethod
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
@@ -61,17 +61,53 @@ class XPUwNa16LinearKernel(MPLinearKernel):
return True, None
def process_weights_after_loading(self, layer: torch.nn.Module):
layer.weight_scale.data = layer.weight_scale.t().contiguous()
# Default names since marlin requires empty parameters for these,
# TODO: remove this requirement from marlin (allow optional tensors)
if self.w_gidx_name is None:
self.w_gidx_name = "g_idx"
if self.w_zp_name is None:
self.w_zp_name = "w_zp"
need_transpose = False
qweight_shape = getattr(layer, self.w_q_name).shape
scale_shape = getattr(layer, self.w_s_name).shape
# gptq marlin and compressed tensors wna16 expect different default
# layouts for weight and scale, so we check the shapes to determine
# if we need to transpose
if qweight_shape[0] != scale_shape[0]:
need_transpose = True
if need_transpose:
getattr(layer, self.w_q_name).data = (
getattr(layer, self.w_q_name).data.t().contiguous()
)
getattr(layer, self.w_s_name).data = getattr(layer, self.w_s_name).data
else:
getattr(layer, self.w_s_name).data = (
getattr(layer, self.w_s_name).data.t().contiguous()
)
if self.config.zero_points:
layer.weight_zero_point.data = layer.weight_zero_point.t().contiguous()
# (FIXME): maybe zero points should also be transposed.
getattr(layer, self.w_zp_name).data = (
getattr(layer, self.w_zp_name).data.t().contiguous()
)
else:
weight_zero_point = torch.Tensor([8]).to(torch.int8).to("xpu")
layer.weight_zero_point = Parameter(weight_zero_point, requires_grad=False)
setattr(
layer, self.w_zp_name, Parameter(weight_zero_point, requires_grad=False)
)
if self.config.has_g_idx:
layer.g_idx.data = layer.g_idx.t().contiguous()
setattr(
layer,
self.w_gidx_name,
Parameter(
getattr(layer, self.w_gidx_name).data.t().contiguous(),
requires_grad=False,
),
)
else:
layer.g_idx = None
setattr(layer, self.w_gidx_name, None)
def apply_weights(
self,
@@ -80,14 +116,15 @@ class XPUwNa16LinearKernel(MPLinearKernel):
bias: torch.Tensor | None = None,
) -> torch.Tensor:
reshaped_x = x.reshape(-1, x.shape[-1])
w_q, w_s, w_zp, w_gidx = self._get_weight_params(layer)
out = torch.ops._xpu_C.int4_gemm_w4a16(
reshaped_x,
layer.weight_packed.t(),
bias,
layer.weight_scale,
layer.weight_zero_point,
w_q.t(),
bias if bias is not None else None,
w_s,
w_zp,
self.config.group_size,
layer.g_idx,
w_gidx,
)
return out
@@ -47,6 +47,9 @@ def query_marlin_supported_quant_types(
if current_platform.is_cpu():
return _query_cpu_marlin_supported_quant_types(has_zp, include_fp_type)
if current_platform.is_xpu():
return [scalar_types.uint4, scalar_types.uint4b8]
if not current_platform.is_rocm():
if device_capability is None:
capability_tuple = current_platform.get_device_capability()