[Kernel] Migrate top-level srt/layers stray kernels to sglang.kernels (RFC #29630, Phase 2.5, 3/7) (#30787)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
e489685509
commit
e9ef06c560
@@ -123,7 +123,9 @@ def fused_rope_inplace(
|
||||
inverse: if True, apply inverse rotation (conjugate freqs)
|
||||
"""
|
||||
if _is_hip or _is_xpu:
|
||||
from sglang.srt.layers.deepseek_v4_rope import apply_rotary_emb_triton
|
||||
from sglang.kernels.ops.attention.deepseek_v4_rope import (
|
||||
apply_rotary_emb_triton,
|
||||
)
|
||||
|
||||
apply_rotary_emb_triton(q, freqs_cis, positions=positions, inverse=inverse)
|
||||
if k is not None:
|
||||
|
||||
@@ -41,3 +41,24 @@ for _mod, _fn in _TRITON_KERNELS:
|
||||
del _mod, _fn
|
||||
|
||||
__all__ = []
|
||||
|
||||
|
||||
# RoPE / QK-norm fusion kernels migrated from srt/layers top-level strays
|
||||
# (RFC #29630, Phase 2.5); registered for inventory.
|
||||
for _mod, _fn in [
|
||||
("deepseek_v4_rope", "precompute_freqs_cis"),
|
||||
("fused_qk_norm_rope_store", "fused_qk_norm_rope_swa_store"),
|
||||
("fused_qk_rmsnorm_rope_gate", "fused_qk_gemma_rmsnorm_rope_gate"),
|
||||
("fused_qk_norm", "fused_qk_norm"),
|
||||
("rotary_triton", "triton_mrope_fused"),
|
||||
("rotary_triton", "triton_ernie45_rope_fused_inplace"),
|
||||
("mrope", "apply_interleaved_rope_triton"),
|
||||
]:
|
||||
register_kernel(
|
||||
KernelSpec(
|
||||
op=f"attention.{_fn}",
|
||||
backend=KernelBackend.TRITON,
|
||||
target=f"sglang.kernels.ops.attention.{_mod}:{_fn}",
|
||||
)
|
||||
)
|
||||
del _mod, _fn
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
"""Interleaved M-RoPE Triton kernel, migrated from
|
||||
``sglang.srt.layers.rotary_embedding.mrope`` (RFC #29630, Phase 2.5).
|
||||
"""
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit
|
||||
def apply_interleaved_rope_kernel(
|
||||
x_ptr,
|
||||
out_ptr,
|
||||
S: tl.constexpr,
|
||||
D: tl.constexpr,
|
||||
stride_x_m,
|
||||
stride_x_s,
|
||||
stride_out_s,
|
||||
section_1_end,
|
||||
section_2_end,
|
||||
BLOCK_S: tl.constexpr,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
start_s = tl.program_id(0) * BLOCK_S
|
||||
s_offsets = start_s + tl.arange(0, BLOCK_S)
|
||||
|
||||
dim_offset = tl.program_id(1) * BLOCK_SIZE
|
||||
dim_indices = dim_offset + tl.arange(0, BLOCK_SIZE)
|
||||
|
||||
mask_s = s_offsets < S
|
||||
mask_d = dim_indices < D
|
||||
mask = mask_s[:, None] & mask_d[None, :]
|
||||
|
||||
val_ptr = (
|
||||
x_ptr + 0 * stride_x_m + s_offsets[:, None] * stride_x_s + dim_indices[None, :]
|
||||
)
|
||||
val = tl.load(val_ptr, mask=mask, other=0.0)
|
||||
|
||||
cond_a = (dim_indices[None, :] % 3 == 1) & (
|
||||
dim_indices[None, :] < section_1_end * 3
|
||||
)
|
||||
val_a_ptr = (
|
||||
x_ptr + 1 * stride_x_m + s_offsets[:, None] * stride_x_s + dim_indices[None, :]
|
||||
)
|
||||
val_a = tl.load(val_a_ptr, mask=mask & cond_a, other=0.0)
|
||||
|
||||
cond_b = (dim_indices[None, :] % 3 == 2) & (
|
||||
dim_indices[None, :] < section_2_end * 3
|
||||
)
|
||||
val_b_ptr = (
|
||||
x_ptr + 2 * stride_x_m + s_offsets[:, None] * stride_x_s + dim_indices[None, :]
|
||||
)
|
||||
val_b = tl.load(val_b_ptr, mask=mask & cond_b, other=0.0)
|
||||
|
||||
val = tl.where(cond_a, val_a, val)
|
||||
val = tl.where(cond_b, val_b, val)
|
||||
|
||||
out_ptr = out_ptr + s_offsets[:, None] * stride_out_s + dim_indices[None, :]
|
||||
tl.store(out_ptr, val, mask=mask)
|
||||
|
||||
|
||||
def apply_interleaved_rope_triton(x: torch.Tensor, mrope_section: list) -> torch.Tensor:
|
||||
x = x.contiguous()
|
||||
M, S, D = x.shape
|
||||
|
||||
out = torch.empty((S, D), dtype=x.dtype, device=x.device)
|
||||
|
||||
BLOCK_S = 64
|
||||
BLOCK_SIZE = 128
|
||||
|
||||
grid = (triton.cdiv(S, BLOCK_S), triton.cdiv(D, BLOCK_SIZE))
|
||||
|
||||
section_1_end = mrope_section[1]
|
||||
section_2_end = mrope_section[2]
|
||||
|
||||
apply_interleaved_rope_kernel[grid](
|
||||
x,
|
||||
out,
|
||||
S,
|
||||
D,
|
||||
x.stride(0),
|
||||
x.stride(1),
|
||||
out.stride(0),
|
||||
section_1_end,
|
||||
section_2_end,
|
||||
BLOCK_S=BLOCK_S,
|
||||
BLOCK_SIZE=BLOCK_SIZE,
|
||||
)
|
||||
return out
|
||||
@@ -336,3 +336,26 @@ __all__ = [
|
||||
"gemma_rmsnorm",
|
||||
"gemma_fused_add_rmsnorm",
|
||||
]
|
||||
|
||||
|
||||
from sglang.kernels.registry import register_kernel
|
||||
from sglang.kernels.spec import KernelSpec
|
||||
|
||||
# Triton / TileLang kernels migrated from srt/layers top-level strays
|
||||
# (RFC #29630, Phase 2.5); registered for inventory.
|
||||
_PHASE25_KERNELS = [
|
||||
("elementwise", "fused_dual_residual_rmsnorm", "triton"),
|
||||
("elementwise", "fused_rmsnorm", "triton"),
|
||||
("gemma4_fused_ops", "gemma4_fused_routing", "triton"),
|
||||
("gemma4_fused_ops", "gemma_qkv_rmsnorm", "triton"),
|
||||
("mhc_head", "fused_hc_head", "triton"),
|
||||
]
|
||||
for _mod, _fn, _bk in _PHASE25_KERNELS:
|
||||
register_kernel(
|
||||
KernelSpec(
|
||||
op=f"layernorm.{_fn}",
|
||||
backend=KernelBackend(_bk),
|
||||
target=f"sglang.kernels.ops.layernorm.{_mod}:{_fn}",
|
||||
)
|
||||
)
|
||||
del _mod, _fn, _bk
|
||||
|
||||
@@ -28,3 +28,20 @@ for _mod, _fn in _TRITON_KERNELS:
|
||||
del _mod, _fn
|
||||
|
||||
__all__ = []
|
||||
|
||||
|
||||
# Migrated from srt/layers (RFC #29630, Phase 2.5).
|
||||
register_kernel(
|
||||
KernelSpec(
|
||||
op="memory.gpu_tensor_hash",
|
||||
backend=KernelBackend.TRITON,
|
||||
target="sglang.kernels.ops.memory.gpu_tensor_hash:gpu_tensor_hash",
|
||||
)
|
||||
)
|
||||
register_kernel(
|
||||
KernelSpec(
|
||||
op="memory.memcpy_triton",
|
||||
backend=KernelBackend.TRITON,
|
||||
target="sglang.kernels.ops.memory.memcpy_triton:memcpy_triton",
|
||||
)
|
||||
)
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
"""Offset/size-driven device memcpy kernel, migrated from
|
||||
``sglang.srt.layers.dp_attention`` (RFC #29630, Phase 2.5).
|
||||
"""
|
||||
|
||||
import functools
|
||||
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit
|
||||
def memcpy_triton_kernel(
|
||||
dst_ptr,
|
||||
src_ptr,
|
||||
offset_ptr,
|
||||
sz_ptr,
|
||||
offset_src: tl.constexpr,
|
||||
chunk_size, # multiplied for offset and sz
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
pid = tl.program_id(axis=0).to(tl.int64)
|
||||
offset = tl.load(offset_ptr).to(tl.int64) * chunk_size
|
||||
sz = tl.load(sz_ptr).to(tl.int64) * chunk_size
|
||||
|
||||
start_index = pid * BLOCK_SIZE
|
||||
offs = tl.arange(0, BLOCK_SIZE)
|
||||
mask = start_index + offs < sz
|
||||
|
||||
if offset_src:
|
||||
data = tl.load(src_ptr + offset + start_index + offs, mask=mask)
|
||||
tl.store(dst_ptr + start_index + offs, data, mask=mask)
|
||||
else:
|
||||
data = tl.load(src_ptr + start_index + offs, mask=mask)
|
||||
tl.store(dst_ptr + offset + start_index + offs, data, mask=mask)
|
||||
|
||||
|
||||
def prod(x):
|
||||
return functools.reduce(lambda a, b: a * b, x, 1)
|
||||
|
||||
|
||||
def memcpy_triton(dst, src, dim, offset, sz, offset_src):
|
||||
max_size = min(src.numel(), dst.numel())
|
||||
assert dim == 0, "dim != 0 unsupported"
|
||||
assert src.shape[1:] == dst.shape[1:], "src and dst must have same shape"
|
||||
chunk_size = prod(src.shape[1:])
|
||||
BLOCK_SIZE = 8192
|
||||
grid = (triton.cdiv(max_size, BLOCK_SIZE),)
|
||||
|
||||
memcpy_triton_kernel[grid](dst, src, offset, sz, offset_src, chunk_size, BLOCK_SIZE)
|
||||
@@ -54,3 +54,13 @@ def top_p_renorm_probs(
|
||||
|
||||
|
||||
__all__ = ["top_k_renorm_probs", "top_p_renorm_probs"]
|
||||
|
||||
|
||||
# Migrated from srt/layers/utils/hash.py (RFC #29630, Phase 2.5).
|
||||
register_kernel(
|
||||
KernelSpec(
|
||||
op="sampling.murmur_hash32",
|
||||
backend=KernelBackend.TRITON,
|
||||
target="sglang.kernels.ops.sampling.murmur_hash:murmur_hash32",
|
||||
)
|
||||
)
|
||||
|
||||
@@ -312,7 +312,7 @@ class CompressorAscendBackendMixin(CompressorBackendMixin):
|
||||
):
|
||||
return self._forward_compress_native(compressor, x, forward_batch)
|
||||
|
||||
from sglang.srt.layers.deepseek_v4_rope import (
|
||||
from sglang.kernels.ops.attention.deepseek_v4_rope import (
|
||||
get_fused_compressor_rope_cos_sin,
|
||||
)
|
||||
|
||||
@@ -649,7 +649,7 @@ class CompressorAscendBackendMixin(CompressorBackendMixin):
|
||||
# Use the same contig cache as the outer rope path; .real/.imag on a
|
||||
# complex tensor are strided views and aclnnIndex over them triggers
|
||||
# StridedSlice (see _get_contig_freqs_real_imag in deepseek_v4_rope.py).
|
||||
from sglang.srt.layers.deepseek_v4_rope import (
|
||||
from sglang.kernels.ops.attention.deepseek_v4_rope import (
|
||||
_get_contig_freqs_real_imag,
|
||||
)
|
||||
|
||||
@@ -890,7 +890,7 @@ class C4IndexerAscendBackendMixin(C4IndexerBackendMixin):
|
||||
def _compute_q_npu(
|
||||
self, c4_indexer, q_lora: torch.Tensor, positions: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
from sglang.srt.layers.deepseek_v4_rope import v4_rope_inplace_npu
|
||||
from sglang.kernels.ops.attention.deepseek_v4_rope import v4_rope_inplace_npu
|
||||
|
||||
bs = q_lora.shape[0]
|
||||
q, _ = c4_indexer.wq_b(q_lora)
|
||||
|
||||
@@ -9,6 +9,11 @@ import torch.nn as nn
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.kernels.ops.attention.deepseek_v4_rope import (
|
||||
apply_rotary_emb_triton,
|
||||
fused_norm_rope_inplace_triton,
|
||||
fused_softmax_pool_triton,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.dsa.dsa_indexer import rotate_activation
|
||||
from sglang.srt.layers.attention.dsv4.compressor import Compressor as _CompressorBase
|
||||
@@ -16,14 +21,9 @@ from sglang.srt.layers.attention.dsv4.fused_compress_triton import (
|
||||
fused_ape_pool_norm_rope,
|
||||
)
|
||||
from sglang.srt.layers.attention.nsa.nsa_indexer import rotate_activation
|
||||
from sglang.srt.layers.deepseek_v4_rope import (
|
||||
apply_rotary_emb_triton,
|
||||
fused_norm_rope_inplace_triton,
|
||||
fused_softmax_pool_triton,
|
||||
)
|
||||
|
||||
try:
|
||||
from sglang.srt.layers.deepseek_v4_rope import fused_softmax_pool_triton
|
||||
from sglang.kernels.ops.attention.deepseek_v4_rope import fused_softmax_pool_triton
|
||||
except ImportError:
|
||||
fused_softmax_pool_triton = None
|
||||
from sglang.srt.mem_cache.deepseek_v4_compress_state import (
|
||||
|
||||
@@ -562,12 +562,14 @@ class CompressorBackendMixin:
|
||||
layer_id: int,
|
||||
) -> None:
|
||||
"""HIP-specific forward path using PyTorch/Triton fallbacks."""
|
||||
from sglang.kernels.ops.attention.deepseek_v4_rope import (
|
||||
fused_norm_rope_inplace_triton,
|
||||
)
|
||||
from sglang.srt.layers.attention.dsv4.quant_k_cache import (
|
||||
quant_to_nope_fp8_rope_bf16_pack_triton,
|
||||
)
|
||||
from sglang.srt.layers.attention.nsa.nsa_indexer import rotate_activation
|
||||
from sglang.srt.layers.attention.nsa.triton_kernel import act_quant
|
||||
from sglang.srt.layers.deepseek_v4_rope import fused_norm_rope_inplace_triton
|
||||
|
||||
compress_ratio = compressor.ratio
|
||||
head_dim = compressor.head_dim
|
||||
|
||||
@@ -1,14 +1,11 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import logging
|
||||
from contextlib import contextmanager
|
||||
from enum import IntEnum, auto
|
||||
from typing import TYPE_CHECKING, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
GroupCoordinator,
|
||||
@@ -376,45 +373,7 @@ def get_dp_local_slice_cpu(
|
||||
return local_start_pos, local_num_tokens
|
||||
|
||||
|
||||
@triton.jit
|
||||
def memcpy_triton_kernel(
|
||||
dst_ptr,
|
||||
src_ptr,
|
||||
offset_ptr,
|
||||
sz_ptr,
|
||||
offset_src: tl.constexpr,
|
||||
chunk_size, # multiplied for offset and sz
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
pid = tl.program_id(axis=0).to(tl.int64)
|
||||
offset = tl.load(offset_ptr).to(tl.int64) * chunk_size
|
||||
sz = tl.load(sz_ptr).to(tl.int64) * chunk_size
|
||||
|
||||
start_index = pid * BLOCK_SIZE
|
||||
offs = tl.arange(0, BLOCK_SIZE)
|
||||
mask = start_index + offs < sz
|
||||
|
||||
if offset_src:
|
||||
data = tl.load(src_ptr + offset + start_index + offs, mask=mask)
|
||||
tl.store(dst_ptr + start_index + offs, data, mask=mask)
|
||||
else:
|
||||
data = tl.load(src_ptr + start_index + offs, mask=mask)
|
||||
tl.store(dst_ptr + offset + start_index + offs, data, mask=mask)
|
||||
|
||||
|
||||
def prod(x):
|
||||
return functools.reduce(lambda a, b: a * b, x, 1)
|
||||
|
||||
|
||||
def memcpy_triton(dst, src, dim, offset, sz, offset_src):
|
||||
max_size = min(src.numel(), dst.numel())
|
||||
assert dim == 0, "dim != 0 unsupported"
|
||||
assert src.shape[1:] == dst.shape[1:], "src and dst must have same shape"
|
||||
chunk_size = prod(src.shape[1:])
|
||||
BLOCK_SIZE = 8192
|
||||
grid = (triton.cdiv(max_size, BLOCK_SIZE),)
|
||||
|
||||
memcpy_triton_kernel[grid](dst, src, offset, sz, offset_src, chunk_size, BLOCK_SIZE)
|
||||
from sglang.kernels.ops.memory.memcpy_triton import memcpy_triton
|
||||
|
||||
|
||||
def _dp_gather_via_all_reduce(
|
||||
|
||||
@@ -7,11 +7,11 @@ from typing import List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.rotary_embedding.base import RotaryEmbedding
|
||||
from sglang.srt.layers.rotary_embedding.triton_kernels import (
|
||||
from sglang.kernels.ops.attention.rotary_triton import (
|
||||
triton_ernie45_rope_fused_inplace,
|
||||
triton_mrope_fused,
|
||||
)
|
||||
from sglang.srt.layers.rotary_embedding.base import RotaryEmbedding
|
||||
from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb
|
||||
from sglang.srt.layers.rotary_embedding.yarn import (
|
||||
yarn_find_correction_range,
|
||||
@@ -41,94 +41,10 @@ if _is_npu:
|
||||
if _is_xpu:
|
||||
from sgl_kernel import multimodal_rotary_embedding
|
||||
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.kernels.ops.attention.mrope import apply_interleaved_rope_triton
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
|
||||
@triton.jit
|
||||
def apply_interleaved_rope_kernel(
|
||||
x_ptr,
|
||||
out_ptr,
|
||||
S: tl.constexpr,
|
||||
D: tl.constexpr,
|
||||
stride_x_m,
|
||||
stride_x_s,
|
||||
stride_out_s,
|
||||
section_1_end,
|
||||
section_2_end,
|
||||
BLOCK_S: tl.constexpr,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
start_s = tl.program_id(0) * BLOCK_S
|
||||
s_offsets = start_s + tl.arange(0, BLOCK_S)
|
||||
|
||||
dim_offset = tl.program_id(1) * BLOCK_SIZE
|
||||
dim_indices = dim_offset + tl.arange(0, BLOCK_SIZE)
|
||||
|
||||
mask_s = s_offsets < S
|
||||
mask_d = dim_indices < D
|
||||
mask = mask_s[:, None] & mask_d[None, :]
|
||||
|
||||
val_ptr = (
|
||||
x_ptr + 0 * stride_x_m + s_offsets[:, None] * stride_x_s + dim_indices[None, :]
|
||||
)
|
||||
val = tl.load(val_ptr, mask=mask, other=0.0)
|
||||
|
||||
cond_a = (dim_indices[None, :] % 3 == 1) & (
|
||||
dim_indices[None, :] < section_1_end * 3
|
||||
)
|
||||
val_a_ptr = (
|
||||
x_ptr + 1 * stride_x_m + s_offsets[:, None] * stride_x_s + dim_indices[None, :]
|
||||
)
|
||||
val_a = tl.load(val_a_ptr, mask=mask & cond_a, other=0.0)
|
||||
|
||||
cond_b = (dim_indices[None, :] % 3 == 2) & (
|
||||
dim_indices[None, :] < section_2_end * 3
|
||||
)
|
||||
val_b_ptr = (
|
||||
x_ptr + 2 * stride_x_m + s_offsets[:, None] * stride_x_s + dim_indices[None, :]
|
||||
)
|
||||
val_b = tl.load(val_b_ptr, mask=mask & cond_b, other=0.0)
|
||||
|
||||
val = tl.where(cond_a, val_a, val)
|
||||
val = tl.where(cond_b, val_b, val)
|
||||
|
||||
out_ptr = out_ptr + s_offsets[:, None] * stride_out_s + dim_indices[None, :]
|
||||
tl.store(out_ptr, val, mask=mask)
|
||||
|
||||
|
||||
def apply_interleaved_rope_triton(x: torch.Tensor, mrope_section: list) -> torch.Tensor:
|
||||
x = x.contiguous()
|
||||
M, S, D = x.shape
|
||||
|
||||
out = torch.empty((S, D), dtype=x.dtype, device=x.device)
|
||||
|
||||
BLOCK_S = 64
|
||||
BLOCK_SIZE = 128
|
||||
|
||||
grid = (triton.cdiv(S, BLOCK_S), triton.cdiv(D, BLOCK_SIZE))
|
||||
|
||||
section_1_end = mrope_section[1]
|
||||
section_2_end = mrope_section[2]
|
||||
|
||||
apply_interleaved_rope_kernel[grid](
|
||||
x,
|
||||
out,
|
||||
S,
|
||||
D,
|
||||
x.stride(0),
|
||||
x.stride(1),
|
||||
out.stride(0),
|
||||
section_1_end,
|
||||
section_2_end,
|
||||
BLOCK_S=BLOCK_S,
|
||||
BLOCK_SIZE=BLOCK_SIZE,
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def apply_interleaved_rope(x: torch.Tensor, mrope_section: list) -> torch.Tensor:
|
||||
x_t = x[0].clone()
|
||||
x_t[..., 1 : mrope_section[1] * 3 : 3] = x[1, ..., 1 : mrope_section[1] * 3 : 3]
|
||||
|
||||
@@ -5,12 +5,12 @@ import torch
|
||||
import torch.distributed as dist
|
||||
from torch import nn
|
||||
|
||||
from sglang.kernels.ops.sampling.murmur_hash import murmur_hash32
|
||||
from sglang.srt.distributed import get_tp_group
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.layers.utils.hash import murmur_hash32
|
||||
from sglang.srt.layers.utils.logprob import get_token_ids_logprobs, get_top_logprobs
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
|
||||
@@ -16,8 +16,8 @@ import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from sglang.kernels.ops.memory.gpu_tensor_hash import gpu_tensor_hash
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.multimodal import gpu_tensor_hash
|
||||
from sglang.srt.managers.io_struct import (
|
||||
BaseBatchReq,
|
||||
TokenizedEmbeddingReqInput,
|
||||
|
||||
@@ -25,6 +25,9 @@ from sglang.jit_kernel.dsv4 import (
|
||||
fused_rope_inplace,
|
||||
sglang_per_token_group_quant_fp8_dsv4_wo_a,
|
||||
)
|
||||
from sglang.kernels.ops.attention.deepseek_v4_rope import (
|
||||
v4_rope_inplace_npu,
|
||||
)
|
||||
from sglang.srt.compilation.compilation_config import register_split_op
|
||||
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
|
||||
from sglang.srt.distributed import (
|
||||
@@ -47,9 +50,6 @@ from sglang.srt.layers.communicator_dsa_cp import (
|
||||
dsa_cp_gather_hidden_states,
|
||||
dsa_cp_reduce_scatter_hidden_states,
|
||||
)
|
||||
from sglang.srt.layers.deepseek_v4_rope import (
|
||||
v4_rope_inplace_npu,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
_tbo_event,
|
||||
attn_tp_all_gather,
|
||||
@@ -132,7 +132,11 @@ if not _is_hip:
|
||||
if _is_xpu:
|
||||
from sgl_kernel import hc_split_sinkhorn
|
||||
else:
|
||||
from sglang.srt.layers.mhc import hc_split_sinkhorn, mhc_fused_post_pre, npu_hc_pre
|
||||
from sglang.kernels.ops.layernorm.mhc import (
|
||||
hc_split_sinkhorn,
|
||||
mhc_fused_post_pre,
|
||||
npu_hc_pre,
|
||||
)
|
||||
|
||||
from sglang.srt.utils import (
|
||||
LazyValue,
|
||||
@@ -476,7 +480,7 @@ class MqaAttentionBase(nn.Module):
|
||||
tp_size=self.attn_tp_size,
|
||||
)
|
||||
|
||||
from sglang.srt.layers.deepseek_v4_rope import precompute_freqs_cis
|
||||
from sglang.kernels.ops.attention.deepseek_v4_rope import precompute_freqs_cis
|
||||
|
||||
rope_theta, rope_scaling = get_rope_config(config)
|
||||
self.rope_scaling = rope_scaling
|
||||
@@ -801,7 +805,7 @@ class MQALayer(MqaAttentionBase):
|
||||
else self.wkv(x_linear)[0]
|
||||
)
|
||||
|
||||
from sglang.srt.layers.fused_qk_norm_rope_store import (
|
||||
from sglang.kernels.ops.attention.fused_qk_norm_rope_store import (
|
||||
fused_qk_norm_rope_swa_store,
|
||||
)
|
||||
|
||||
@@ -915,7 +919,7 @@ class MQALayer(MqaAttentionBase):
|
||||
False,
|
||||
)
|
||||
|
||||
from sglang.srt.layers.fused_qk_norm_rope_store import (
|
||||
from sglang.kernels.ops.attention.fused_qk_norm_rope_store import (
|
||||
fused_qk_norm_rope_swa_store,
|
||||
)
|
||||
|
||||
@@ -1366,7 +1370,7 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
return y, post, comb, False
|
||||
|
||||
if envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get():
|
||||
from sglang.srt.layers.mhc import mhc_pre
|
||||
from sglang.kernels.ops.layernorm.mhc import mhc_pre
|
||||
|
||||
norm_kwargs = {}
|
||||
if norm is not None:
|
||||
@@ -1450,7 +1454,7 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
return torch.ops.custom.npu_hc_post(x, residual, post, comb)
|
||||
|
||||
if envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get():
|
||||
from sglang.srt.layers.mhc import mhc_post
|
||||
from sglang.kernels.ops.layernorm.mhc import mhc_post
|
||||
|
||||
return mhc_post(x, residual, post, comb)
|
||||
|
||||
@@ -2048,7 +2052,7 @@ class DeepseekV4Model(nn.Module):
|
||||
hc_base: torch.Tensor,
|
||||
):
|
||||
if x.numel() > 0:
|
||||
from sglang.srt.layers.mhc_head import fused_hc_head
|
||||
from sglang.kernels.ops.layernorm.mhc_head import fused_hc_head
|
||||
|
||||
return fused_hc_head(
|
||||
x.contiguous(),
|
||||
@@ -2315,7 +2319,7 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
# batched/contiguous-load rope kernels (faster on gfx95; .
|
||||
# Module-level toggles default OFF; flipped True here for DSV4
|
||||
if _is_hip:
|
||||
from sglang.srt.layers.deepseek_v4_rope import set_batched_rope
|
||||
from sglang.kernels.ops.attention.deepseek_v4_rope import set_batched_rope
|
||||
from sglang.srt.layers.quantization.fp8_utils import set_force_ck_w8a8
|
||||
|
||||
set_force_ck_w8a8(True)
|
||||
@@ -2600,7 +2604,7 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
if layer is None:
|
||||
return
|
||||
|
||||
from sglang.srt.layers.mhc import prewarm_mhc_pre
|
||||
from sglang.kernels.ops.layernorm.mhc import prewarm_mhc_pre
|
||||
|
||||
tic = time.perf_counter()
|
||||
prewarm_mhc_pre(
|
||||
|
||||
@@ -24,16 +24,16 @@ from transformers import (
|
||||
PreTrainedModel,
|
||||
)
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
)
|
||||
from sglang.srt.layers.gemma4_fused_ops import (
|
||||
from sglang.kernels.ops.layernorm.gemma4_fused_ops import (
|
||||
gemma4_fused_routing,
|
||||
gemma_dual_rmsnorm_residual_scalar,
|
||||
gemma_qkv_rmsnorm,
|
||||
gemma_rmsnorm_residual_scalar,
|
||||
gemma_routing_post_topk,
|
||||
)
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
)
|
||||
from sglang.srt.layers.layernorm import Gemma4RMSNorm, RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
QKVParallelLinear,
|
||||
|
||||
@@ -22,16 +22,16 @@ import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
from sglang.kernels.ops.layernorm.elementwise import (
|
||||
fused_dual_residual_rmsnorm,
|
||||
fused_rmsnorm,
|
||||
gelu_and_mul_triton,
|
||||
)
|
||||
from sglang.kernels.ops.moe.router import fused_moe_router_shim
|
||||
from sglang.srt.distributed import (
|
||||
tensor_model_parallel_all_reduce,
|
||||
)
|
||||
from sglang.srt.layers.activation import GeluAndMul
|
||||
from sglang.srt.layers.elementwise import (
|
||||
fused_dual_residual_rmsnorm,
|
||||
fused_rmsnorm,
|
||||
gelu_and_mul_triton,
|
||||
)
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
MergedColumnParallelLinear,
|
||||
|
||||
@@ -27,6 +27,7 @@ import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
from sglang.kernels.ops.layernorm.elementwise import fused_gate_sigmoid_mul_add
|
||||
from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
@@ -48,7 +49,6 @@ from sglang.srt.layers.cp.utils import is_cp_v2_active
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.layers.elementwise import fused_gate_sigmoid_mul_add
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
MergedColumnParallelLinear,
|
||||
|
||||
@@ -25,6 +25,7 @@ import triton
|
||||
from sglang.jit_kernel.triton.gdn_fused_proj import (
|
||||
fused_qkvzba_split_reshape_cat_contiguous,
|
||||
)
|
||||
from sglang.kernels.ops.layernorm.elementwise import fused_sigmoid_mul
|
||||
|
||||
# Configs
|
||||
from sglang.srt.configs.qwen3_5 import (
|
||||
@@ -45,7 +46,6 @@ from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.layers.elementwise import fused_sigmoid_mul
|
||||
|
||||
# Layers - Others
|
||||
from sglang.srt.layers.layernorm import GemmaRMSNorm
|
||||
@@ -142,7 +142,7 @@ def _disable_shared_experts_fusion() -> bool:
|
||||
|
||||
|
||||
if _is_cuda:
|
||||
from sglang.srt.layers.fused_qk_rmsnorm_rope_gate import (
|
||||
from sglang.kernels.ops.attention.fused_qk_rmsnorm_rope_gate import (
|
||||
fused_qk_gemma_rmsnorm_rope_gate,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user