Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
a5a71c6c26
commit
1a35440c4a
@@ -43,6 +43,26 @@ del _mod, _fn
|
||||
__all__ = []
|
||||
|
||||
|
||||
# Generic attention kernels migrated in Phase 2.5 (RFC #29630).
|
||||
for _mod, _fn in [
|
||||
("utils", "mla_quantize_and_rope_for_fp8"),
|
||||
("utils", "launch_reshape_and_cache_flash"),
|
||||
("utils", "launch_reshape_and_cache_shuffle_5d"),
|
||||
("flash_mla_sm120", "flash_mla_with_kvcache_sm120"),
|
||||
("dcp_kernels", "create_dcp_kv_indices"),
|
||||
("dcp_kernels", "correct_attn_out"),
|
||||
("pa_page_table", "_build_pa_page_table"),
|
||||
("nsa_triton_decode", "triton_sparse_attn_decode"),
|
||||
]:
|
||||
register_kernel(
|
||||
KernelSpec(
|
||||
op=f"attention.{_fn.lstrip('_')}",
|
||||
backend=KernelBackend.TRITON,
|
||||
target=f"sglang.kernels.ops.attention.{_mod}:{_fn}",
|
||||
)
|
||||
)
|
||||
del _mod, _fn
|
||||
|
||||
# RoPE / QK-norm fusion kernels migrated from srt/layers top-level strays
|
||||
# (RFC #29630, Phase 2.5); registered for inventory.
|
||||
for _mod, _fn in [
|
||||
|
||||
+1
-1
@@ -235,7 +235,7 @@ def flash_mla_with_kvcache_sm120(**kwargs):
|
||||
)
|
||||
|
||||
if _sm120_default_backend == "triton":
|
||||
from sglang.srt.layers.attention.flash_mla_sm120_triton import (
|
||||
from sglang.kernels.ops.attention.flash_mla_sm120_triton import (
|
||||
flash_mla_sparse_decode_triton,
|
||||
)
|
||||
|
||||
+1
-1
@@ -9,7 +9,7 @@ from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.nsa.triton_decode.triton_mla_kernels_decode_optimized import (
|
||||
from sglang.kernels.ops.attention.nsa_triton_decode.triton_mla_kernels_decode_optimized import (
|
||||
triton_sparse_attn_decode,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
"""Paged-attention page-table builder, migrated from
|
||||
``sglang.srt.layers.attention.flashattention_backend`` (RFC #29630, Phase 2.5).
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _build_pa_page_table_kernel(
|
||||
req_to_token_ptr,
|
||||
req_pool_indices_ptr,
|
||||
seq_lens_ptr,
|
||||
prefill_lens_ptr,
|
||||
dst_page_table_ptr,
|
||||
kv_lens_ptr,
|
||||
window_size: tl.constexpr,
|
||||
req_to_token_stride,
|
||||
dst_stride,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
"""Build PA-SWA page_table directly from req_to_token.
|
||||
|
||||
For each request, dst row = [0..prefill_len) ∪ [decode_start..seq_len).
|
||||
decode_start = max(prefill_len, seq_len - window_size)
|
||||
|
||||
prefill_lens_ptr is the full pool-sized buffer, prefill_len is loaded
|
||||
via indirect indexing using req_idx.
|
||||
"""
|
||||
bid = tl.program_id(0)
|
||||
req_idx = tl.load(req_pool_indices_ptr + bid)
|
||||
sl = tl.load(seq_lens_ptr + bid).to(tl.int32)
|
||||
pf = tl.load(prefill_lens_ptr + req_idx).to(tl.int32)
|
||||
|
||||
decode_start = tl.maximum(pf, sl - window_size)
|
||||
gap = tl.where(decode_start > pf, decode_start - pf, 0)
|
||||
kv_len = sl - gap
|
||||
|
||||
tl.store(kv_lens_ptr + bid, kv_len)
|
||||
|
||||
src_base = req_idx * req_to_token_stride
|
||||
dst_base = bid * dst_stride
|
||||
|
||||
for start in tl.range(0, kv_len, BLOCK_SIZE):
|
||||
offs = start + tl.arange(0, BLOCK_SIZE)
|
||||
mask = offs < kv_len
|
||||
pos = tl.where(offs < pf, offs, offs + gap)
|
||||
kv_loc = tl.load(
|
||||
req_to_token_ptr + src_base + pos,
|
||||
mask=mask,
|
||||
other=0,
|
||||
)
|
||||
tl.store(dst_page_table_ptr + dst_base + offs, kv_loc.to(tl.int32), mask=mask)
|
||||
|
||||
|
||||
def _build_pa_page_table(
|
||||
req_to_token: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
prefill_lens: torch.Tensor,
|
||||
window_size: int,
|
||||
bs: int,
|
||||
pa_max_len: int,
|
||||
device: torch.device,
|
||||
dst_page_table: Optional[torch.Tensor] = None,
|
||||
dst_kv_lens: Optional[torch.Tensor] = None,
|
||||
):
|
||||
"""Build prefill-aware page_table from req_to_token.
|
||||
|
||||
When dst_page_table/dst_kv_lens are None, allocates new tensors (non-CUDA-graph).
|
||||
When provided, writes in-place into existing buffers (CUDA-graph replay).
|
||||
|
||||
prefill_lens is the full pool-sized buffer; the kernel indexes it via
|
||||
req_pool_indices values (indirect indexing, avoids external gather).
|
||||
|
||||
Returns (page_table, kv_lens).
|
||||
"""
|
||||
if dst_page_table is None:
|
||||
dst_page_table = torch.zeros(bs, pa_max_len, dtype=torch.int32, device=device)
|
||||
if dst_kv_lens is None:
|
||||
dst_kv_lens = torch.empty(bs, dtype=torch.int32, device=device)
|
||||
if bs > 0 and pa_max_len > 0:
|
||||
_build_pa_page_table_kernel[(bs,)](
|
||||
req_to_token,
|
||||
req_pool_indices.contiguous(),
|
||||
seq_lens.to(torch.int32),
|
||||
prefill_lens,
|
||||
dst_page_table,
|
||||
dst_kv_lens,
|
||||
window_size,
|
||||
req_to_token.stride(0),
|
||||
dst_page_table.stride(0),
|
||||
BLOCK_SIZE=256,
|
||||
)
|
||||
return dst_page_table, dst_kv_lens
|
||||
@@ -14,17 +14,17 @@ from typing import TYPE_CHECKING, Optional
|
||||
import torch
|
||||
import triton
|
||||
|
||||
from sglang.kernels.ops.kvcache.aiter_unified_attention import (
|
||||
scatter_ragged_to_page_table_kernel,
|
||||
scatter_req_to_token_to_page_table_kernel,
|
||||
)
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.attention.utils import (
|
||||
from sglang.kernels.ops.attention.utils import (
|
||||
assert_buffer_fits,
|
||||
create_flashinfer_kv_indices_triton,
|
||||
create_flashmla_kv_indices_triton,
|
||||
get_num_kv_index_blocks_flashmla,
|
||||
)
|
||||
from sglang.kernels.ops.kvcache.aiter_unified_attention import (
|
||||
scatter_ragged_to_page_table_kernel,
|
||||
scatter_req_to_token_to_page_table_kernel,
|
||||
)
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
@@ -60,16 +60,16 @@ except ImportError:
|
||||
"aiter is AMD specific kernel library. Please make sure aiter is installed on your AMD device."
|
||||
)
|
||||
|
||||
from sglang.kernels.ops.attention.utils import (
|
||||
launch_reshape_and_cache_flash,
|
||||
pad_sequence_with_mask,
|
||||
)
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype
|
||||
from sglang.srt.configs.model_config import AttentionArch
|
||||
from sglang.srt.layers.attention.aiter_utils import (
|
||||
forward_decode_vectorized_5d,
|
||||
forward_extend_vectorized_5d,
|
||||
)
|
||||
from sglang.srt.layers.attention.utils import (
|
||||
launch_reshape_and_cache_flash,
|
||||
pad_sequence_with_mask,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||
from sglang.srt.utils import get_bool_env_var
|
||||
|
||||
@@ -33,8 +33,8 @@ except ImportError: # pragma: no cover - import-time guard mirrors aiter_backen
|
||||
pa_decode_gluon = None
|
||||
get_recommended_splits = None
|
||||
|
||||
from sglang.kernels.ops.attention.utils import launch_gather_shuffle_5d_to_linear
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype
|
||||
from sglang.srt.layers.attention.utils import launch_gather_shuffle_5d_to_linear
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.attention.aiter_backend import AiterAttnBackend
|
||||
|
||||
@@ -13,11 +13,11 @@ from typing import TYPE_CHECKING, Optional, Union
|
||||
import torch
|
||||
import triton
|
||||
|
||||
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
|
||||
from sglang.srt.layers.attention.utils import (
|
||||
from sglang.kernels.ops.attention.utils import (
|
||||
create_flashmla_kv_indices_triton,
|
||||
get_num_kv_index_blocks_flashmla,
|
||||
)
|
||||
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.utils import is_cuda
|
||||
|
||||
|
||||
@@ -1671,7 +1671,7 @@ class DeepseekV4AttnBackend(
|
||||
)
|
||||
|
||||
if _is_sm120:
|
||||
from sglang.srt.layers.attention.flash_mla_sm120 import (
|
||||
from sglang.kernels.ops.attention.flash_mla_sm120 import (
|
||||
flash_mla_with_kvcache_sm120,
|
||||
)
|
||||
|
||||
|
||||
@@ -11,8 +11,8 @@ from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.attention.utils import seqlens_expand_triton
|
||||
from sglang.srt.layers.attention.dsa.utils import compute_dsa_seqlens
|
||||
from sglang.srt.layers.attention.utils import seqlens_expand_triton
|
||||
from sglang.srt.utils import is_cuda, is_hip
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
@@ -18,6 +18,11 @@ from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
from sglang.kernels.ops.attention.utils import (
|
||||
concat_mla_absorb_q_general,
|
||||
mla_quantize_and_rope_for_fp8,
|
||||
seqlens_expand_triton,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.attention.dsa.dequant_k_cache import dequantize_k_cache_paged
|
||||
@@ -47,11 +52,6 @@ from sglang.srt.layers.attention.dsa.utils import (
|
||||
pad_dsa_cache_seqlens,
|
||||
should_use_dsa_fused_topk,
|
||||
)
|
||||
from sglang.srt.layers.attention.utils import (
|
||||
concat_mla_absorb_q_general,
|
||||
mla_quantize_and_rope_for_fp8,
|
||||
seqlens_expand_triton,
|
||||
)
|
||||
from sglang.srt.layers.utils.cp_utils import (
|
||||
cp_all_gather_rerange_output,
|
||||
cp_split_and_rebuild_position,
|
||||
|
||||
@@ -5,19 +5,18 @@ from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.kernels.ops.attention.metadata import (
|
||||
normal_decode_set_metadata,
|
||||
prepare_swa_spec_page_table_triton,
|
||||
)
|
||||
from sglang.kernels.ops.attention.pa_page_table import _build_pa_page_table
|
||||
from sglang.kernels.ops.attention.utils import assert_buffer_fits
|
||||
from sglang.kernels.ops.kvcache.trtllm_mha_page_table import (
|
||||
build_trtllm_mha_page_table,
|
||||
)
|
||||
from sglang.srt.configs.model_config import AttentionArch
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.attention.utils import assert_buffer_fits
|
||||
from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy
|
||||
from sglang.srt.layers.cp.utils import is_cp_v2_active
|
||||
from sglang.srt.layers.radix_attention import AttentionType
|
||||
@@ -50,95 +49,6 @@ def _should_disable_scheduler_metadata_precompute(server_args) -> bool:
|
||||
return bool(server_args.enable_prefill_cp or server_args.enable_dp_attention)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _build_pa_page_table_kernel(
|
||||
req_to_token_ptr,
|
||||
req_pool_indices_ptr,
|
||||
seq_lens_ptr,
|
||||
prefill_lens_ptr,
|
||||
dst_page_table_ptr,
|
||||
kv_lens_ptr,
|
||||
window_size: tl.constexpr,
|
||||
req_to_token_stride,
|
||||
dst_stride,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
"""Build PA-SWA page_table directly from req_to_token.
|
||||
|
||||
For each request, dst row = [0..prefill_len) ∪ [decode_start..seq_len).
|
||||
decode_start = max(prefill_len, seq_len - window_size)
|
||||
|
||||
prefill_lens_ptr is the full pool-sized buffer, prefill_len is loaded
|
||||
via indirect indexing using req_idx.
|
||||
"""
|
||||
bid = tl.program_id(0)
|
||||
req_idx = tl.load(req_pool_indices_ptr + bid)
|
||||
sl = tl.load(seq_lens_ptr + bid).to(tl.int32)
|
||||
pf = tl.load(prefill_lens_ptr + req_idx).to(tl.int32)
|
||||
|
||||
decode_start = tl.maximum(pf, sl - window_size)
|
||||
gap = tl.where(decode_start > pf, decode_start - pf, 0)
|
||||
kv_len = sl - gap
|
||||
|
||||
tl.store(kv_lens_ptr + bid, kv_len)
|
||||
|
||||
src_base = req_idx * req_to_token_stride
|
||||
dst_base = bid * dst_stride
|
||||
|
||||
for start in tl.range(0, kv_len, BLOCK_SIZE):
|
||||
offs = start + tl.arange(0, BLOCK_SIZE)
|
||||
mask = offs < kv_len
|
||||
pos = tl.where(offs < pf, offs, offs + gap)
|
||||
kv_loc = tl.load(
|
||||
req_to_token_ptr + src_base + pos,
|
||||
mask=mask,
|
||||
other=0,
|
||||
)
|
||||
tl.store(dst_page_table_ptr + dst_base + offs, kv_loc.to(tl.int32), mask=mask)
|
||||
|
||||
|
||||
def _build_pa_page_table(
|
||||
req_to_token: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
prefill_lens: torch.Tensor,
|
||||
window_size: int,
|
||||
bs: int,
|
||||
pa_max_len: int,
|
||||
device: torch.device,
|
||||
dst_page_table: Optional[torch.Tensor] = None,
|
||||
dst_kv_lens: Optional[torch.Tensor] = None,
|
||||
):
|
||||
"""Build prefill-aware page_table from req_to_token.
|
||||
|
||||
When dst_page_table/dst_kv_lens are None, allocates new tensors (non-CUDA-graph).
|
||||
When provided, writes in-place into existing buffers (CUDA-graph replay).
|
||||
|
||||
prefill_lens is the full pool-sized buffer; the kernel indexes it via
|
||||
req_pool_indices values (indirect indexing, avoids external gather).
|
||||
|
||||
Returns (page_table, kv_lens).
|
||||
"""
|
||||
if dst_page_table is None:
|
||||
dst_page_table = torch.zeros(bs, pa_max_len, dtype=torch.int32, device=device)
|
||||
if dst_kv_lens is None:
|
||||
dst_kv_lens = torch.empty(bs, dtype=torch.int32, device=device)
|
||||
if bs > 0 and pa_max_len > 0:
|
||||
_build_pa_page_table_kernel[(bs,)](
|
||||
req_to_token,
|
||||
req_pool_indices.contiguous(),
|
||||
seq_lens.to(torch.int32),
|
||||
prefill_lens,
|
||||
dst_page_table,
|
||||
dst_kv_lens,
|
||||
window_size,
|
||||
req_to_token.stride(0),
|
||||
dst_page_table.stride(0),
|
||||
BLOCK_SIZE=256,
|
||||
)
|
||||
return dst_page_table, dst_kv_lens
|
||||
|
||||
|
||||
@dataclass
|
||||
class FlashAttentionMetadata:
|
||||
"""Metadata to be init once in the model forward pass,
|
||||
|
||||
@@ -19,13 +19,13 @@ from typing import TYPE_CHECKING, Callable, List, Optional, Union
|
||||
import torch
|
||||
|
||||
from sglang.kernel_api_logging import debug_kernel_api
|
||||
from sglang.srt.dllm.config import DllmConfig
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.attention.utils import (
|
||||
from sglang.kernels.ops.attention.utils import (
|
||||
assert_buffer_fits,
|
||||
create_flashinfer_kv_indices_triton,
|
||||
)
|
||||
from sglang.srt.dllm.config import DllmConfig
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.radix_attention import AttentionType
|
||||
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
|
||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||
|
||||
@@ -17,12 +17,12 @@ from typing import TYPE_CHECKING, Callable, Optional, Union
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.attention.utils import assert_buffer_fits
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.attention.flashinfer_backend import (
|
||||
create_flashinfer_kv_indices_triton,
|
||||
)
|
||||
from sglang.srt.layers.attention.utils import assert_buffer_fits
|
||||
from sglang.srt.layers.dcp import (
|
||||
DecodeContextParallelMetadata,
|
||||
update_local_kv_lens_for_dcp,
|
||||
|
||||
@@ -12,12 +12,12 @@ import torch
|
||||
import triton
|
||||
from sgl_kernel.flash_mla import flash_mla_with_kvcache, get_mla_metadata
|
||||
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import scaled_fp8_quant
|
||||
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
|
||||
from sglang.srt.layers.attention.utils import (
|
||||
from sglang.kernels.ops.attention.utils import (
|
||||
create_flashmla_kv_indices_triton,
|
||||
get_num_kv_index_blocks_flashmla,
|
||||
)
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import scaled_fp8_quant
|
||||
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
|
||||
@@ -38,7 +38,7 @@ def flash_mla_with_kvcache_entrypoint(backend: str, **kwargs):
|
||||
return dpsk_v4_fp8_attention_fwd(**kwargs)
|
||||
|
||||
if backend == "triton":
|
||||
from sglang.srt.layers.attention.nsa.triton_decode import (
|
||||
from sglang.kernels.ops.attention.nsa_triton_decode import (
|
||||
triton_fp8_attention_fwd,
|
||||
)
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.attention.utils import canonicalize_stride
|
||||
from sglang.kernels.ops.kvcache.trtllm_fp8_kv_kernel import (
|
||||
fused_fp8_set_kv_buffer,
|
||||
)
|
||||
@@ -27,7 +28,6 @@ from sglang.srt.layers.attention.flashinfer_backend import (
|
||||
FlashInferAttnBackend,
|
||||
FlashInferMultiStepDraftBackend,
|
||||
)
|
||||
from sglang.srt.layers.attention.utils import canonicalize_stride
|
||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
|
||||
@@ -19,6 +19,10 @@ from sglang.kernels.ops.attention.pad import (
|
||||
from sglang.kernels.ops.attention.pad import (
|
||||
unpad_draft_extend_output as unpad_draft_extend_output_triton,
|
||||
)
|
||||
from sglang.kernels.ops.attention.utils import (
|
||||
concat_mla_absorb_q_general,
|
||||
mla_quantize_and_rope_for_fp8,
|
||||
)
|
||||
from sglang.kernels.ops.kvcache.kv_indices import (
|
||||
create_flashmla_kv_indices_triton,
|
||||
get_num_kv_index_blocks_flashmla,
|
||||
@@ -30,10 +34,6 @@ from sglang.srt.layers.attention.flashinfer_mla_backend import (
|
||||
FlashInferMLAAttnBackend,
|
||||
FlashInferMLAMultiStepDraftBackend,
|
||||
)
|
||||
from sglang.srt.layers.attention.utils import (
|
||||
concat_mla_absorb_q_general,
|
||||
mla_quantize_and_rope_for_fp8,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||
is_in_tc_piecewise_cuda_graph,
|
||||
|
||||
@@ -29,6 +29,9 @@ Package-internal helpers (the @triton.jit kernels, ``CPTritonContext``,
|
||||
out-of-tree callers; in-tree code should use ``get_parallel().dcp_enabled`` and
|
||||
``get_parallel().attn_dcp_*``."""
|
||||
|
||||
from sglang.kernels.ops.attention.dcp_kernels import (
|
||||
create_triton_kv_indices_for_dcp_triton,
|
||||
)
|
||||
from sglang.srt.layers.dcp.comm import (
|
||||
all_gather_kv_cache_for_dcp,
|
||||
all_gather_kv_cache_for_mha_chunk_extend,
|
||||
@@ -41,7 +44,6 @@ from sglang.srt.layers.dcp.comm import (
|
||||
get_attention_dcp_rank,
|
||||
get_attention_dcp_world_size,
|
||||
)
|
||||
from sglang.srt.layers.dcp.kernels import create_triton_kv_indices_for_dcp_triton
|
||||
from sglang.srt.layers.dcp.layout import (
|
||||
filter_dcp_local_kv_indices,
|
||||
get_dcp_lens,
|
||||
|
||||
@@ -25,11 +25,11 @@ from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.attention.dcp_kernels import CPTritonContext, correct_attn_out
|
||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
use_symmetric_memory,
|
||||
)
|
||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||
from sglang.srt.layers.dcp.kernels import CPTritonContext, correct_attn_out
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
|
||||
|
||||
@@ -20,7 +20,7 @@ from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.dcp.kernels import (
|
||||
from sglang.kernels.ops.attention.dcp_kernels import (
|
||||
create_dcp_kv_indices,
|
||||
update_kv_lens_and_indices,
|
||||
)
|
||||
|
||||
@@ -65,7 +65,7 @@ if _is_npu:
|
||||
)
|
||||
|
||||
if _is_hip:
|
||||
from sglang.srt.layers.attention.utils import (
|
||||
from sglang.kernels.ops.attention.utils import (
|
||||
fused_qk_rope_reshape_and_cache,
|
||||
)
|
||||
|
||||
|
||||
@@ -1837,7 +1837,7 @@ class MHATokenToKVPool(KVCache):
|
||||
# and viewed as ``store_dtype`` by ``set_kv_buffer``.
|
||||
if self.kv_cache_layout == "vectorized_5d":
|
||||
# Late-import to keep the NHD path import-clean.
|
||||
from sglang.srt.layers.attention.utils import (
|
||||
from sglang.kernels.ops.attention.utils import (
|
||||
launch_reshape_and_cache_shuffle_5d,
|
||||
)
|
||||
|
||||
|
||||
@@ -4,10 +4,10 @@ from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.attention.utils import concat_and_cast_mha_k_triton
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.dsa.dequant_k_cache import dequantize_k_cache_paged
|
||||
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
|
||||
from sglang.srt.layers.attention.utils import concat_and_cast_mha_k_triton
|
||||
from sglang.srt.layers.communicator import get_attn_tp_context
|
||||
from sglang.srt.layers.dcp import (
|
||||
all_gather_kv_cache_for_mha_chunk_extend,
|
||||
|
||||
@@ -12,6 +12,7 @@ import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
from sglang.kernels.ops.attention.utils import concat_and_cast_mha_k_triton
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
tensor_model_parallel_all_reduce,
|
||||
@@ -19,7 +20,6 @@ from sglang.srt.distributed import (
|
||||
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
||||
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
|
||||
from sglang.srt.layers.activation import SiluAndMul
|
||||
from sglang.srt.layers.attention.utils import concat_and_cast_mha_k_triton
|
||||
from sglang.srt.layers.communicator import (
|
||||
LayerCommunicator,
|
||||
LayerScatterModes,
|
||||
|
||||
@@ -5,7 +5,7 @@ from typing import TYPE_CHECKING, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
|
||||
from sglang.kernels.ops.attention.utils import create_flashinfer_kv_indices_triton
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
from sglang.srt.model_executor.forward_batch_info import (
|
||||
CaptureHiddenMode,
|
||||
|
||||
@@ -4,9 +4,9 @@ from typing import List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.attention.utils import create_flashinfer_kv_indices_triton
|
||||
from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
|
||||
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
|
||||
|
||||
@@ -4,8 +4,8 @@ from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.attention.utils import create_flashinfer_kv_indices_triton
|
||||
from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject
|
||||
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
|
||||
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user