[PCG] fix piecewise cuda graph for Qwen3.5 (#19220)
This commit is contained in:
@@ -72,6 +72,13 @@ if _is_cuda:
|
|||||||
N = mat_b.shape[-1]
|
N = mat_b.shape[-1]
|
||||||
return mat_a.new_empty((M, N), dtype=out_dtype)
|
return mat_a.new_empty((M, N), dtype=out_dtype)
|
||||||
|
|
||||||
|
@torch.library.register_fake("sgl_kernel::fp8_blockwise_scaled_mm")
|
||||||
|
def _fp8_blockwise_scaled_mm_abstract(mat_a, mat_b, scales_a, scales_b, out_dtype):
|
||||||
|
# mat_a: [M, K], mat_b: [K, N] or [N, K] depending on callsite layout; output is [M, N].
|
||||||
|
M = mat_a.shape[-2]
|
||||||
|
N = mat_b.shape[-1]
|
||||||
|
return mat_a.new_empty((M, N), dtype=out_dtype)
|
||||||
|
|
||||||
|
|
||||||
use_vllm_cutlass_w8a8_fp8_kernel = get_bool_env_var("USE_VLLM_CUTLASS_W8A8_FP8_KERNEL")
|
use_vllm_cutlass_w8a8_fp8_kernel = get_bool_env_var("USE_VLLM_CUTLASS_W8A8_FP8_KERNEL")
|
||||||
use_triton_w8a8_fp8_kernel = get_bool_env_var("USE_TRITON_W8A8_FP8_KERNEL")
|
use_triton_w8a8_fp8_kernel = get_bool_env_var("USE_TRITON_W8A8_FP8_KERNEL")
|
||||||
|
|||||||
@@ -22,9 +22,6 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from einops import rearrange
|
from einops import rearrange
|
||||||
|
|
||||||
# Model Executor
|
|
||||||
from sglang.srt.compilation.piecewise_context_manager import get_forward_context
|
|
||||||
|
|
||||||
# Configs
|
# Configs
|
||||||
from sglang.srt.configs.qwen3_5 import (
|
from sglang.srt.configs.qwen3_5 import (
|
||||||
Qwen3_5Config,
|
Qwen3_5Config,
|
||||||
@@ -72,7 +69,6 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock
|
from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock
|
||||||
|
|
||||||
# Models
|
# Models
|
||||||
from sglang.srt.models.qwen3_next import gdn_with_output
|
|
||||||
from sglang.srt.models.qwen3_vl import Qwen3VLForConditionalGeneration
|
from sglang.srt.models.qwen3_vl import Qwen3VLForConditionalGeneration
|
||||||
|
|
||||||
# Utils
|
# Utils
|
||||||
@@ -253,22 +249,6 @@ class Qwen3_5GatedDeltaNet(nn.Module):
|
|||||||
self,
|
self,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
):
|
|
||||||
output = torch.empty_like(hidden_states)
|
|
||||||
if forward_batch.forward_mode.is_extend() and get_forward_context() is not None:
|
|
||||||
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,
|
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Forward pass with three parts:
|
Forward pass with three parts:
|
||||||
@@ -287,7 +267,7 @@ class Qwen3_5GatedDeltaNet(nn.Module):
|
|||||||
b = b.contiguous()
|
b = b.contiguous()
|
||||||
a = a.contiguous()
|
a = a.contiguous()
|
||||||
|
|
||||||
core_attn_out = self.attn.forward(
|
core_attn_out = self.attn(
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
mixed_qkv=mixed_qkv,
|
mixed_qkv=mixed_qkv,
|
||||||
a=a,
|
a=a,
|
||||||
|
|||||||
@@ -5,8 +5,6 @@ from typing import Any, Iterable, Optional, Set, Tuple
|
|||||||
import torch
|
import torch
|
||||||
from torch import nn
|
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.configs.qwen3_next import Qwen3NextConfig
|
from sglang.srt.configs.qwen3_next import Qwen3NextConfig
|
||||||
from sglang.srt.distributed import get_pp_group
|
from sglang.srt.distributed import get_pp_group
|
||||||
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
||||||
@@ -53,7 +51,6 @@ from sglang.srt.utils import (
|
|||||||
make_layers,
|
make_layers,
|
||||||
set_weight_attrs,
|
set_weight_attrs,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.custom_op import register_custom_op
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
@@ -1149,25 +1146,3 @@ class Qwen3NextForCausalLM(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
EntryClass = Qwen3NextForCausalLM
|
EntryClass = Qwen3NextForCausalLM
|
||||||
|
|
||||||
|
|
||||||
@register_custom_op(mutates_args=["output"])
|
|
||||||
@register_split_op()
|
|
||||||
def gdn_with_output(
|
|
||||||
hidden_states: torch.Tensor,
|
|
||||||
output: torch.Tensor,
|
|
||||||
layer_id: int,
|
|
||||||
) -> None:
|
|
||||||
context = get_forward_context()
|
|
||||||
forward_batch = context.forward_batch
|
|
||||||
attention_layers = context.attention_layers
|
|
||||||
attention_layer = attention_layers[layer_id]
|
|
||||||
|
|
||||||
ret = attention_layer._forward(hidden_states, forward_batch)
|
|
||||||
|
|
||||||
assert (
|
|
||||||
output.numel() == ret.numel()
|
|
||||||
), f"Output tensor element mismatch: {output.numel()} != {ret.numel()}"
|
|
||||||
|
|
||||||
output.view(ret.shape).copy_(ret)
|
|
||||||
return
|
|
||||||
|
|||||||
@@ -1233,6 +1233,7 @@ class Qwen3VLForConditionalGeneration(nn.Module):
|
|||||||
def should_apply_lora(self, module_name: str) -> bool:
|
def should_apply_lora(self, module_name: str) -> bool:
|
||||||
return bool(self._lora_pattern.match(module_name))
|
return bool(self._lora_pattern.match(module_name))
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
input_ids: torch.Tensor,
|
input_ids: torch.Tensor,
|
||||||
|
|||||||
Reference in New Issue
Block a user