[Perf] refactor piecewise cuda graph support of Qwen3-Next (#17613)

This commit is contained in:
Minglei Zhu
2026-02-13 17:30:50 -08:00
committed by GitHub
parent 3a1c388b43
commit 8be18c655d
5 changed files with 80 additions and 34 deletions

View File

@@ -14,6 +14,7 @@ import triton
import triton.language as tl
from einops import rearrange
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
cdiv,
cpu_has_amx_support,
@@ -26,6 +27,9 @@ from sglang.srt.utils import (
_is_npu = is_npu()
_use_cpu = is_cpu() and cpu_has_amx_support()
# Maximum rows per Triton block for layernorm gated kernel
MAX_ROWS_PER_BLOCK = 4
def rms_norm_ref(
x,
@@ -173,9 +177,17 @@ def _get_sm_count(device: torch.device) -> int:
def calc_rows_per_block(M: int, device: torch.device) -> int:
# When piecewise cuda graph is enabled, use a constant value to avoid
# torch.compile creating guards on the dynamic batch dimension.
try:
if get_global_server_args().enable_piecewise_cuda_graph:
return MAX_ROWS_PER_BLOCK
except ValueError:
# Global server args not initialized (e.g., in unit tests)
pass
sm_count = _get_sm_count(device)
rows_per_block = next_power_of_2(cdiv(M, 2 * sm_count))
rows_per_block = min(rows_per_block, 4)
rows_per_block = min(rows_per_block, MAX_ROWS_PER_BLOCK)
return rows_per_block

View File

@@ -19,6 +19,10 @@ from typing import TYPE_CHECKING, Optional, Tuple, Union
import torch
from torch import nn
from sglang.srt.compilation.compilation_config import register_split_op
from sglang.srt.compilation.piecewise_context_manager import get_forward_context
from sglang.srt.utils.custom_op import register_custom_op
if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@@ -70,10 +74,60 @@ class RadixLinearAttention(nn.Module):
a: torch.Tensor,
b: torch.Tensor,
) -> torch.Tensor:
return forward_batch.attn_backend.forward(
layer=self,
forward_batch=forward_batch,
mixed_qkv=mixed_qkv,
a=a,
b=b,
)
if forward_batch.forward_mode.is_extend() and get_forward_context() is not None:
# Output shape from linear attention: (1, seq_len, num_v_heads, head_v_dim)
seq_len = mixed_qkv.shape[0]
output = torch.empty(
(1, seq_len, self.num_v_heads, self.head_v_dim),
dtype=mixed_qkv.dtype,
device=mixed_qkv.device,
)
unified_linear_attention_with_output(
mixed_qkv,
a,
b,
output,
self.layer_id,
)
return output
else:
return forward_batch.attn_backend.forward(
layer=self,
forward_batch=forward_batch,
mixed_qkv=mixed_qkv,
a=a,
b=b,
)
@register_custom_op(mutates_args=["output"])
@register_split_op()
def unified_linear_attention_with_output(
mixed_qkv: torch.Tensor,
a: torch.Tensor,
b: torch.Tensor,
output: torch.Tensor,
layer_id: int,
) -> None:
"""
Custom op wrapper for linear attention computation only.
"""
context = get_forward_context()
forward_batch = context.forward_batch
attention_layers = context.attention_layers
attention_layer = attention_layers[layer_id]
ret = forward_batch.attn_backend.forward(
layer=attention_layer,
forward_batch=forward_batch,
mixed_qkv=mixed_qkv,
a=a,
b=b,
)
assert (
output.numel() == ret.numel()
), f"Output tensor element mismatch: {output.numel()} != {ret.numel()}"
output.view(ret.shape).copy_(ret)
return

View File

@@ -2159,7 +2159,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
elif hasattr(layer, "attn"):
self.attention_layers.append(layer.attn)
elif hasattr(layer, "linear_attn"):
self.attention_layers.append(layer.linear_attn)
if hasattr(layer.linear_attn, "attn"):
self.attention_layers.append(layer.linear_attn.attn)
else:
self.attention_layers.append(layer.linear_attn)
# For InternVL model
elif hasattr(layer, "attention"):
if hasattr(layer.attention, "attn"):

View File

@@ -316,7 +316,7 @@ class Qwen3GatedDeltaNet(nn.Module):
prefix=add_prefix("out_proj", prefix),
)
self.linear_attn = RadixLinearAttention(
self.attn = RadixLinearAttention(
layer_id=layer_id,
num_q_heads=self.num_k_heads // self.attn_tp_size,
num_k_heads=self.num_k_heads // self.attn_tp_size,
@@ -405,23 +405,6 @@ class Qwen3GatedDeltaNet(nn.Module):
hidden_states: torch.Tensor,
forward_batch: ForwardBatch,
):
if forward_batch.forward_mode.is_extend() and get_forward_context() is not None:
output = torch.empty_like(hidden_states)
gdn_with_output(
hidden_states,
output,
self.layer_id,
)
return output
else:
return self._forward(hidden_states, forward_batch)
def _forward(
self,
hidden_states: torch.Tensor,
forward_batch: ForwardBatch,
):
seq_len, _ = hidden_states.shape
is_cuda_graph = forward_batch.forward_mode.is_cuda_graph()
projected_states_qkvz, projected_states_ba = self._forward_input_proj(
@@ -460,7 +443,7 @@ class Qwen3GatedDeltaNet(nn.Module):
lambda x: x.reshape(x.shape[0], -1), (query, key, value)
)
mixed_qkv = torch.cat((query, key, value), dim=-1)
core_attn_out = self.linear_attn(
core_attn_out = self.attn(
forward_batch,
mixed_qkv=mixed_qkv,
a=a,