[PieceWise CUDA Graph] Support awq/gptq model in piecewise cudagraph (#12518)

This commit is contained in:
Xiaoyu Zhang
2025-11-11 11:56:15 +08:00
committed by GitHub
parent 08c805a85f
commit 9caca6a45c
7 changed files with 339 additions and 42 deletions

View File

@@ -17,6 +17,9 @@ class ForwardContext:
def set_attention_layers(self, layers: List[Any]):
self.attention_layers = layers
def set_quant_config(self, quant_config: Any):
self.quant_config = quant_config
_forward_context: Optional[ForwardContext] = None
@@ -28,11 +31,14 @@ def get_forward_context() -> Optional[ForwardContext]:
@contextmanager
def set_forward_context(forward_batch: ForwardBatch, attention_layers: List[Any]):
def set_forward_context(
forward_batch: ForwardBatch, attention_layers: List[Any], quant_config: Any
):
global _forward_context
_forward_context = ForwardContext()
_forward_context.set_forward_batch(forward_batch)
_forward_context.set_attention_layers(attention_layers)
_forward_context.set_quant_config(quant_config)
try:
yield
finally:

View File

@@ -954,3 +954,25 @@ class AWQMoEAscendMethod(AWQMoEMethod):
use_wna16=True,
)
return StandardCombineInput(hidden_states=output)
# Register fake implementations for torch.compile support
if _is_cuda:
@torch.library.register_fake("sgl_kernel::awq_dequantize")
def _(
qweight,
scales,
qzeros,
ch_axis,
group_size,
num_bits,
):
out_shape = qweight.shape[:-1] + (qweight.shape[-1] * 32 // num_bits,)
return qweight.new_empty(out_shape, dtype=scales.dtype)
@torch.library.register_fake("sgl_kernel::awq_marlin_repack")
def _(b_q_weight, size_k, size_n, num_bits):
return b_q_weight.new_empty(
(size_k // 16, size_n * (num_bits // 2)), dtype=b_q_weight.dtype
)

View File

@@ -1094,3 +1094,50 @@ class GPTQMarlinMoEMethod(FusedMoEMethodBase):
is_k_full=self.is_k_full,
).to(orig_dtype)
return StandardCombineInput(hidden_states=output)
# Register fake implementations for torch.compile support
if _is_cuda:
@torch.library.register_fake("sgl_kernel::gptq_gemm")
def _(a, b_q_weight, b_gptq_qzeros, b_gptq_scales, b_g_idx, use_shuffle, bit):
return a.new_empty((a.shape[0], b_q_weight.shape[-1]), dtype=a.dtype)
@torch.library.register_fake("sgl_kernel::gptq_marlin_repack")
def _(b_q_weight, perm, size_k, size_n, num_bits):
return b_q_weight.new_empty(
(size_k // 16, size_n * (num_bits // 2)), dtype=b_q_weight.dtype
)
@torch.library.register_fake("sgl_kernel::gptq_shuffle")
def _(q_weight, q_perm, bit):
return
@torch.library.register_fake("sgl_kernel::moe_wna16_marlin_gemm")
def _(
a,
c,
b_q_weight,
b_scales,
b_zeros,
g_idx,
perm,
workspace,
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
topk_weights,
moe_block_size,
top_k,
mul_topk_weights,
is_ep,
b_q_type_id,
size_m,
size_n,
size_k,
is_k_full,
use_atomic_add,
use_fp32_reduce,
is_zp_float,
):
return c

View File

@@ -31,11 +31,15 @@ if TYPE_CHECKING:
from sglang.srt.layers.linear import LinearBase
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.compilation.piecewise_context_manager import get_forward_context
try:
from vllm import _custom_ops as ops
except ImportError:
ops = None
from sglang.srt.utils import direct_register_custom_op
_is_cuda = is_cuda()
if _is_cuda:
@@ -483,25 +487,44 @@ def apply_gptq_marlin_linear(
dtype=input.dtype,
)
output = gptq_marlin_gemm(
reshaped_x,
None,
weight,
weight_scale,
None,
weight_zp,
g_idx,
g_idx_sort_indices,
workspace,
wtype,
size_m=reshaped_x.shape[0],
size_n=output_size_per_partition,
size_k=input_size_per_partition,
is_k_full=is_k_full,
use_atomic_add=use_atomic_add,
use_fp32_reduce=use_fp32_reduce,
is_zp_float=False,
)
forward_context = get_forward_context()
if forward_context is None:
output = gptq_marlin_gemm(
reshaped_x,
None,
weight,
weight_scale,
None,
weight_zp,
g_idx,
g_idx_sort_indices,
workspace,
wtype,
size_m=reshaped_x.shape[0],
size_n=output_size_per_partition,
size_k=input_size_per_partition,
is_k_full=is_k_full,
use_atomic_add=use_atomic_add,
use_fp32_reduce=use_fp32_reduce,
is_zp_float=False,
)
else:
output = torch.ops.sglang.unified_apply_gptq_marlin_gemm_with_wtype(
input=reshaped_x,
weight=weight,
weight_scale=weight_scale,
weight_zp=weight_zp,
g_idx=g_idx,
g_idx_sort_indices=g_idx_sort_indices,
workspace=workspace,
wtype_id=wtype.id,
output_size_per_partition=output_size_per_partition,
input_size_per_partition=input_size_per_partition,
is_k_full=is_k_full,
use_atomic_add=use_atomic_add,
use_fp32_reduce=use_fp32_reduce,
is_zp_float=False,
)
if bias is not None:
output.add_(bias) # In-place add
@@ -534,24 +557,41 @@ def apply_awq_marlin_linear(
dtype=input.dtype,
)
output = gptq_marlin_gemm(
reshaped_x,
None,
weight,
weight_scale,
None,
weight_zp,
g_idx,
g_idx_sort_indices,
workspace,
quant_type,
size_m=reshaped_x.shape[0],
size_n=output_size_per_partition,
size_k=input_size_per_partition,
use_atomic_add=use_atomic_add,
use_fp32_reduce=use_fp32_reduce,
is_zp_float=False,
)
forward_context = get_forward_context()
if forward_context is None:
output = gptq_marlin_gemm(
reshaped_x,
None,
weight,
weight_scale,
None,
weight_zp,
g_idx,
g_idx_sort_indices,
workspace,
quant_type,
size_m=reshaped_x.shape[0],
size_n=output_size_per_partition,
size_k=input_size_per_partition,
use_atomic_add=use_atomic_add,
use_fp32_reduce=use_fp32_reduce,
is_zp_float=False,
)
else:
output = torch.ops.sglang.unified_apply_gptq_marlin_gemm(
input=reshaped_x,
weight=weight,
weight_scale=weight_scale,
weight_zp=weight_zp,
g_idx=g_idx,
g_idx_sort_indices=g_idx_sort_indices,
workspace=workspace,
output_size_per_partition=output_size_per_partition,
input_size_per_partition=input_size_per_partition,
use_atomic_add=use_atomic_add,
use_fp32_reduce=use_fp32_reduce,
is_zp_float=False,
)
if bias is not None:
output.add_(bias) # In-place add
@@ -818,3 +858,140 @@ class MarlinLinearMethod(LinearMethodBase):
output.add_(bias) # In-place add
return output
def unified_apply_gptq_marlin_gemm(
input: torch.Tensor,
weight: torch.Tensor,
weight_scale: torch.Tensor,
weight_zp: torch.Tensor,
g_idx: torch.Tensor,
g_idx_sort_indices: torch.Tensor,
workspace: torch.Tensor,
output_size_per_partition: int,
input_size_per_partition: int,
use_atomic_add: bool,
use_fp32_reduce: bool,
is_zp_float: bool,
) -> torch.Tensor:
quant_config = get_forward_context().quant_config
quant_type = quant_config.quant_type
return gptq_marlin_gemm(
input,
None,
weight,
weight_scale,
None,
weight_zp,
g_idx,
g_idx_sort_indices,
workspace,
quant_type,
size_m=input.shape[0],
size_n=output_size_per_partition,
size_k=input_size_per_partition,
use_atomic_add=use_atomic_add,
use_fp32_reduce=use_fp32_reduce,
is_zp_float=is_zp_float,
)
def fake_unified_apply_gptq_marlin_gemm(
input: torch.Tensor,
weight: torch.Tensor,
weight_scale: torch.Tensor,
weight_zp: torch.Tensor,
g_idx: torch.Tensor,
g_idx_sort_indices: torch.Tensor,
workspace: torch.Tensor,
output_size_per_partition: int,
input_size_per_partition: int,
use_atomic_add: bool,
use_fp32_reduce: bool,
is_zp_float: bool,
) -> torch.Tensor:
return input.new_empty(
(input.shape[0], output_size_per_partition), dtype=input.dtype
)
direct_register_custom_op(
op_name="unified_apply_gptq_marlin_gemm",
op_func=unified_apply_gptq_marlin_gemm,
mutates_args=[],
fake_impl=fake_unified_apply_gptq_marlin_gemm,
)
def unified_apply_gptq_marlin_gemm_with_wtype(
input: torch.Tensor,
weight: torch.Tensor,
weight_scale: torch.Tensor,
weight_zp: torch.Tensor,
g_idx: torch.Tensor,
g_idx_sort_indices: torch.Tensor,
workspace: torch.Tensor,
wtype_id: int,
output_size_per_partition: int,
input_size_per_partition: int,
is_k_full: bool,
use_atomic_add: bool,
use_fp32_reduce: bool,
is_zp_float: bool,
) -> torch.Tensor:
# Reconstruct ScalarType from id
wtype = None
for attr_name in dir(scalar_types):
if not attr_name.startswith("_"):
st = getattr(scalar_types, attr_name)
if hasattr(st, "id") and st.id == wtype_id:
wtype = st
break
return gptq_marlin_gemm(
input,
None,
weight,
weight_scale,
None,
weight_zp,
g_idx,
g_idx_sort_indices,
workspace,
wtype,
size_m=input.shape[0],
size_n=output_size_per_partition,
size_k=input_size_per_partition,
is_k_full=is_k_full,
use_atomic_add=use_atomic_add,
use_fp32_reduce=use_fp32_reduce,
is_zp_float=is_zp_float,
)
def fake_unified_apply_gptq_marlin_gemm_with_wtype(
input: torch.Tensor,
weight: torch.Tensor,
weight_scale: torch.Tensor,
weight_zp: torch.Tensor,
g_idx: torch.Tensor,
g_idx_sort_indices: torch.Tensor,
workspace: torch.Tensor,
wtype_id: int,
output_size_per_partition: int,
input_size_per_partition: int,
is_k_full: bool,
use_atomic_add: bool,
use_fp32_reduce: bool,
is_zp_float: bool,
) -> torch.Tensor:
return input.new_empty(
(input.shape[0], output_size_per_partition), dtype=input.dtype
)
direct_register_custom_op(
op_name="unified_apply_gptq_marlin_gemm_with_wtype",
op_func=unified_apply_gptq_marlin_gemm_with_wtype,
mutates_args=[],
fake_impl=fake_unified_apply_gptq_marlin_gemm_with_wtype,
)

View File

@@ -163,6 +163,7 @@ class PiecewiseCudaGraphRunner:
self.model_runner.server_args.piecewise_cuda_graph_tokens,
self.model_runner.server_args.piecewise_cuda_graph_compiler,
)
self.quant_config = getattr(self.model_runner.model, "quant_config", None)
# Batch sizes to capture
self.capture_num_tokens = self.compile_config.get_capture_sizes()
@@ -263,7 +264,9 @@ class PiecewiseCudaGraphRunner:
# Attention backend
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
with set_forward_context(forward_batch, self.attention_layers):
with set_forward_context(
forward_batch, self.attention_layers, self.quant_config
):
_ = self.model_runner.model.forward(
forward_batch.input_ids,
forward_batch.positions,
@@ -405,7 +408,9 @@ class PiecewiseCudaGraphRunner:
set_is_extend_in_batch(False)
kwargs = {}
with set_forward_context(forward_batch, self.attention_layers):
with set_forward_context(
forward_batch, self.attention_layers, self.quant_config
):
self.model_runner.model.forward(
forward_batch.input_ids,
forward_batch.positions,
@@ -503,7 +508,9 @@ class PiecewiseCudaGraphRunner:
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
static_forward_batch = self.replay_prepare(forward_batch, **kwargs)
# Replay
with set_forward_context(static_forward_batch, self.attention_layers):
with set_forward_context(
static_forward_batch, self.attention_layers, self.quant_config
):
with set_compiled(True):
output = self.model_runner.model.forward(
static_forward_batch.input_ids,