mirror of
https://github.com/vllm-project/vllm.git
synced 2026-08-13 09:18:12 +00:00
[XPU] add gptq(int4) support (#37844)
Signed-off-by: Kunshang Ji <[email protected]>
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user