[Kernel] Sweep missed dedicated kernels into kernels.ops (moe/quant siblings + dspark) (RFC #29630) (#32160)
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
f35411ee81
commit
62aa85d9aa
+5
-5
@@ -13,15 +13,15 @@ import triton
|
|||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
from sglang.kernels.jit.utils import is_arch_support_pdl
|
from sglang.kernels.jit.utils import is_arch_support_pdl
|
||||||
from sglang.kernels.ops.moe.inkling_gate_topk_renorm import (
|
from sglang.kernels.ops.moe.gate_topk import (
|
||||||
inkling_gate_topk_renorm_v2,
|
|
||||||
)
|
|
||||||
from sglang.srt.environ import envs
|
|
||||||
from sglang.srt.layers.moe.moe_runner.triton_utils.gate_topk import (
|
|
||||||
fpval_to_key,
|
fpval_to_key,
|
||||||
indx_to_key,
|
indx_to_key,
|
||||||
key_to_indx,
|
key_to_indx,
|
||||||
)
|
)
|
||||||
|
from sglang.kernels.ops.moe.inkling_gate_topk_renorm import (
|
||||||
|
inkling_gate_topk_renorm_v2,
|
||||||
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
+1
-1
@@ -7,6 +7,7 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.kernels.ops.speculative.dspark.dispatch import inputs_on_cuda
|
||||||
from sglang.kernels.ops.speculative.reject_sampling import (
|
from sglang.kernels.ops.speculative.reject_sampling import (
|
||||||
chain_speculative_sampling_triton,
|
chain_speculative_sampling_triton,
|
||||||
)
|
)
|
||||||
@@ -16,7 +17,6 @@ from sglang.srt.speculative.dflash_utils import (
|
|||||||
build_dflash_verify_target_probs,
|
build_dflash_verify_target_probs,
|
||||||
compute_dflash_correct_drafts_and_bonus,
|
compute_dflash_correct_drafts_and_bonus,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dispatch import inputs_on_cuda
|
|
||||||
|
|
||||||
|
|
||||||
class AcceptSampling:
|
class AcceptSampling:
|
||||||
+1
-1
@@ -7,7 +7,7 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dispatch import inputs_on_cuda
|
from sglang.kernels.ops.speculative.dspark.dispatch import inputs_on_cuda
|
||||||
from sglang.srt.utils import ceil_align
|
from sglang.srt.utils import ceil_align
|
||||||
|
|
||||||
|
|
||||||
+1
-1
@@ -8,7 +8,7 @@ import torch.nn.functional as F
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dispatch import inputs_on_cuda
|
from sglang.kernels.ops.speculative.dspark.dispatch import inputs_on_cuda
|
||||||
|
|
||||||
_BLOCK_V = 1024
|
_BLOCK_V = 1024
|
||||||
_IDX_SENTINEL = tl.constexpr(2147483647)
|
_IDX_SENTINEL = tl.constexpr(2147483647)
|
||||||
+1
-1
@@ -6,7 +6,7 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dispatch import (
|
from sglang.kernels.ops.speculative.dspark.dispatch import (
|
||||||
inputs_on_cuda,
|
inputs_on_cuda,
|
||||||
)
|
)
|
||||||
|
|
||||||
+1
-1
@@ -6,8 +6,8 @@ import triton
|
|||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
from sglang.kernels.ops.speculative.cache_locs import assign_extend_cache_locs_func
|
from sglang.kernels.ops.speculative.cache_locs import assign_extend_cache_locs_func
|
||||||
|
from sglang.kernels.ops.speculative.dspark.dispatch import inputs_on_cuda
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dispatch import inputs_on_cuda
|
|
||||||
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
|
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
|
||||||
|
|
||||||
|
|
||||||
@@ -33,6 +33,11 @@ from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import (
|
|||||||
BuildPageTablePositions,
|
BuildPageTablePositions,
|
||||||
ExpandPrefillCausally,
|
ExpandPrefillCausally,
|
||||||
)
|
)
|
||||||
|
from sglang.kernels.ops.speculative.dspark.dspark_attn_metadata import (
|
||||||
|
BuildBlockSeqLensCausal,
|
||||||
|
BuildDsparkSwaPageIndices,
|
||||||
|
ComputeDsparkWindowGather,
|
||||||
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
from sglang.srt.layers.attention.dsv4.compressor_v2 import (
|
from sglang.srt.layers.attention.dsv4.compressor_v2 import (
|
||||||
@@ -54,11 +59,6 @@ from sglang.srt.layers.attention.dsv4.sparse_prefill_utils import (
|
|||||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dspark_attn_metadata import (
|
|
||||||
BuildBlockSeqLensCausal,
|
|
||||||
BuildDsparkSwaPageIndices,
|
|
||||||
ComputeDsparkWindowGather,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
|
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
|
||||||
from sglang.srt.speculative.ragged_verify import (
|
from sglang.srt.speculative.ragged_verify import (
|
||||||
RaggedVerifyMode,
|
RaggedVerifyMode,
|
||||||
|
|||||||
@@ -3361,7 +3361,7 @@ class MHATokenToKVPoolMXFP8(MHATokenToKVPool):
|
|||||||
# payload and the interleaved UE8M0 scales.
|
# payload and the interleaved UE8M0 scales.
|
||||||
if not self.mxfp8_sf_interleaved or cache_k.dtype == self.store_dtype:
|
if not self.mxfp8_sf_interleaved or cache_k.dtype == self.store_dtype:
|
||||||
raise ValueError("MXFP8 KV cache requires K and V scale tensors.")
|
raise ValueError("MXFP8 KV cache requires K and V scale tensors.")
|
||||||
from sglang.srt.layers.quantization.mxfp8_quant import quant_store_kv_mxfp8
|
from sglang.kernels.ops.quantization.mxfp8_quant import quant_store_kv_mxfp8
|
||||||
|
|
||||||
quant_store_kv_mxfp8(
|
quant_store_kv_mxfp8(
|
||||||
cache_k,
|
cache_k,
|
||||||
@@ -3394,7 +3394,7 @@ class MHATokenToKVPoolMXFP8(MHATokenToKVPool):
|
|||||||
"""Write per-token UE8M0 K/V scales — interleaved into the FA4
|
"""Write per-token UE8M0 K/V scales — interleaved into the FA4
|
||||||
BlockScaledBasicChunk layout for page_size==128, flat otherwise."""
|
BlockScaledBasicChunk layout for page_size==128, flat otherwise."""
|
||||||
if self.mxfp8_sf_interleaved:
|
if self.mxfp8_sf_interleaved:
|
||||||
from sglang.srt.layers.quantization.mxfp8_interleave_sf import (
|
from sglang.kernels.ops.quantization.mxfp8_interleave_sf import (
|
||||||
store_sf_interleaved,
|
store_sf_interleaved,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -3436,7 +3436,7 @@ class MHATokenToKVPoolMXFP8(MHATokenToKVPool):
|
|||||||
# scale rows must travel with their fp8 payload or dequant reads
|
# scale rows must travel with their fp8 payload or dequant reads
|
||||||
# mismatched exponents.
|
# mismatched exponents.
|
||||||
if self.mxfp8_sf_interleaved:
|
if self.mxfp8_sf_interleaved:
|
||||||
from sglang.srt.layers.quantization.mxfp8_interleave_sf import (
|
from sglang.kernels.ops.quantization.mxfp8_interleave_sf import (
|
||||||
store_sf_interleaved,
|
store_sf_interleaved,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -9,6 +9,10 @@ import torch.nn.functional as F
|
|||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
from sglang.kernels.ops.attention.dsv4 import fused_q_norm_rope, fused_rope_inplace
|
from sglang.kernels.ops.attention.dsv4 import fused_q_norm_rope, fused_rope_inplace
|
||||||
|
from sglang.kernels.ops.speculative.dspark.dspark_draft_model import (
|
||||||
|
BuildStepLocal,
|
||||||
|
CommitKvProj,
|
||||||
|
)
|
||||||
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
|
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.layernorm import RMSNorm
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
@@ -40,10 +44,6 @@ from sglang.srt.runtime_context import get_parallel
|
|||||||
from sglang.srt.speculative.dspark_components.dspark_config import (
|
from sglang.srt.speculative.dspark_components.dspark_config import (
|
||||||
parse_dspark_draft_config,
|
parse_dspark_draft_config,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dspark_draft_model import (
|
|
||||||
BuildStepLocal,
|
|
||||||
CommitKvProj,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.ragged_verify import (
|
from sglang.srt.speculative.ragged_verify import (
|
||||||
RaggedVerifyMode,
|
RaggedVerifyMode,
|
||||||
read_ragged_verify_mode,
|
read_ragged_verify_mode,
|
||||||
|
|||||||
@@ -883,7 +883,7 @@ class InklingAttention(nn.Module):
|
|||||||
if prologue_q_descale is not None:
|
if prologue_q_descale is not None:
|
||||||
extra_attn_kwargs["q_descale"] = prologue_q_descale
|
extra_attn_kwargs["q_descale"] = prologue_q_descale
|
||||||
else:
|
else:
|
||||||
from sglang.srt.layers.quantization.mxfp8_quant import to_mxfp8
|
from sglang.kernels.ops.quantization.mxfp8_quant import to_mxfp8
|
||||||
|
|
||||||
q_mxfp = to_mxfp8(q.view(num_tokens, self.num_tp_heads, self.head_dim))
|
q_mxfp = to_mxfp8(q.view(num_tokens, self.num_tp_heads, self.head_dim))
|
||||||
q = q_mxfp.data.view(num_tokens, -1)
|
q = q_mxfp.data.view(num_tokens, -1)
|
||||||
|
|||||||
@@ -382,7 +382,7 @@ class InklingBatchDenseMLP(nn.Module, FusedMoELoadingMixin):
|
|||||||
|
|
||||||
def _swiglu(self, y_st2f: torch.Tensor, gammas_st: torch.Tensor) -> torch.Tensor:
|
def _swiglu(self, y_st2f: torch.Tensor, gammas_st: torch.Tensor) -> torch.Tensor:
|
||||||
# Helion's kernel can produce NaNs for small shared-expert batches.
|
# Helion's kernel can produce NaNs for small shared-expert batches.
|
||||||
from sglang.srt.layers.moe.moe_runner.triton_utils.inkling_moe import (
|
from sglang.kernels.ops.moe.inkling_moe import (
|
||||||
silu_and_mul_triton,
|
silu_and_mul_triton,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -10,20 +10,13 @@ from torch import nn
|
|||||||
from triton.language.extra import libdevice
|
from triton.language.extra import libdevice
|
||||||
|
|
||||||
from sglang.kernels.jit.utils import is_arch_support_pdl
|
from sglang.kernels.jit.utils import is_arch_support_pdl
|
||||||
|
from sglang.kernels.ops.moe.gate_topk import gate_topk
|
||||||
from sglang.kernels.ops.moe.inkling_gate_topk_renorm import (
|
from sglang.kernels.ops.moe.inkling_gate_topk_renorm import (
|
||||||
ensure_gate_gemv_fused_scratch,
|
ensure_gate_gemv_fused_scratch,
|
||||||
inkling_gate_gemv,
|
inkling_gate_gemv,
|
||||||
inkling_gate_gemv_fused,
|
inkling_gate_gemv_fused,
|
||||||
)
|
)
|
||||||
from sglang.srt.configs.inkling import InklingModelConfig
|
from sglang.kernels.ops.moe.inkling_moe import (
|
||||||
from sglang.srt.distributed import (
|
|
||||||
get_tensor_model_parallel_group,
|
|
||||||
)
|
|
||||||
from sglang.srt.environ import GateGemvMode, envs
|
|
||||||
from sglang.srt.layers.moe import get_moe_runner_backend
|
|
||||||
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
|
||||||
from sglang.srt.layers.moe.moe_runner.triton_utils.gate_topk import gate_topk
|
|
||||||
from sglang.srt.layers.moe.moe_runner.triton_utils.inkling_moe import (
|
|
||||||
FUSED_PREPROCESS_WIN_TOKENS,
|
FUSED_PREPROCESS_WIN_TOKENS,
|
||||||
compute_grouped_gemm_metadata,
|
compute_grouped_gemm_metadata,
|
||||||
fused_moe_preprocess,
|
fused_moe_preprocess,
|
||||||
@@ -34,9 +27,16 @@ from sglang.srt.layers.moe.moe_runner.triton_utils.inkling_moe import (
|
|||||||
select_grouped_gemm_block_m,
|
select_grouped_gemm_block_m,
|
||||||
silu_and_mul_helion,
|
silu_and_mul_helion,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.moe_runner.triton_utils.sigmoid_gate_topk_renorm import (
|
from sglang.kernels.ops.moe.sigmoid_gate_topk_renorm import (
|
||||||
sigmoid_gate_topk_renorm,
|
sigmoid_gate_topk_renorm,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.configs.inkling import InklingModelConfig
|
||||||
|
from sglang.srt.distributed import (
|
||||||
|
get_tensor_model_parallel_group,
|
||||||
|
)
|
||||||
|
from sglang.srt.environ import GateGemvMode, envs
|
||||||
|
from sglang.srt.layers.moe import get_moe_runner_backend
|
||||||
|
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
||||||
from sglang.srt.layers.moe.topk import PackedTopKOutput, StandardTopKOutput
|
from sglang.srt.layers.moe.topk import PackedTopKOutput, StandardTopKOutput
|
||||||
from sglang.srt.layers.moe.utils import RoutingMethodType
|
from sglang.srt.layers.moe.utils import RoutingMethodType
|
||||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||||
|
|||||||
@@ -7,6 +7,9 @@ from typing import Optional
|
|||||||
import msgspec
|
import msgspec
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.kernels.ops.speculative.dspark.dspark_draft_model import (
|
||||||
|
SampleStepTokens,
|
||||||
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.model_executor.forward_batch_info import (
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
@@ -18,9 +21,6 @@ from sglang.srt.runtime_context import get_parallel
|
|||||||
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
|
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
|
||||||
from sglang.srt.speculative.draft_worker_common import make_draft_input_v2
|
from sglang.srt.speculative.draft_worker_common import make_draft_input_v2
|
||||||
from sglang.srt.speculative.dspark_components.dspark_planner import VerifyWindow
|
from sglang.srt.speculative.dspark_components.dspark_planner import VerifyWindow
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dspark_draft_model import (
|
|
||||||
SampleStepTokens,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.spec_info import (
|
from sglang.srt.speculative.spec_info import (
|
||||||
SpeculativeAlgorithm,
|
SpeculativeAlgorithm,
|
||||||
spec_scale_global_num_tokens,
|
spec_scale_global_num_tokens,
|
||||||
|
|||||||
@@ -3,10 +3,10 @@ from typing import Optional
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.kernels.ops.speculative.cache_locs import assign_extend_cache_locs_func
|
from sglang.kernels.ops.speculative.cache_locs import assign_extend_cache_locs_func
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.kernels.ops.speculative.dspark.dspark_verify_window import (
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dspark_verify_window import (
|
|
||||||
BuildCommitInjectLayout,
|
BuildCommitInjectLayout,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
|
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,10 @@ import msgspec
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.kernels.ops.speculative.cache_locs import assign_extend_cache_locs_func
|
from sglang.kernels.ops.speculative.cache_locs import assign_extend_cache_locs_func
|
||||||
|
from sglang.kernels.ops.speculative.dspark.dspark_schedule import (
|
||||||
|
ScheduleVerifyLensTopk,
|
||||||
|
compute_sort_survival,
|
||||||
|
)
|
||||||
from sglang.srt.distributed import get_tp_group
|
from sglang.srt.distributed import get_tp_group
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
|
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
|
||||||
@@ -31,10 +35,6 @@ from sglang.srt.speculative.dspark_components.dspark_sps import (
|
|||||||
from sglang.srt.speculative.dspark_components.dspark_sts import (
|
from sglang.srt.speculative.dspark_components.dspark_sts import (
|
||||||
load_sts_calibration_from_path,
|
load_sts_calibration_from_path,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dspark_schedule import (
|
|
||||||
ScheduleVerifyLensTopk,
|
|
||||||
compute_sort_survival,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.ragged_verify import (
|
from sglang.srt.speculative.ragged_verify import (
|
||||||
RaggedVerifyLayout,
|
RaggedVerifyLayout,
|
||||||
RaggedVerifyMode,
|
RaggedVerifyMode,
|
||||||
|
|||||||
@@ -5,6 +5,23 @@ from typing import Optional
|
|||||||
import msgspec
|
import msgspec
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.kernels.ops.speculative.dspark.dspark_accept import (
|
||||||
|
AcceptGreedy,
|
||||||
|
AcceptSampling,
|
||||||
|
FinalizeAcceptLens,
|
||||||
|
SelectMixedAccept,
|
||||||
|
SoftmaxTemp,
|
||||||
|
accept_greedy_triton,
|
||||||
|
finalize_accept_lens_triton,
|
||||||
|
)
|
||||||
|
from sglang.kernels.ops.speculative.dspark.dspark_verify_window import (
|
||||||
|
BuildCommitInjectLayout,
|
||||||
|
BuildOutTokens,
|
||||||
|
BuildRaggedVerifyWindow,
|
||||||
|
RaggedVerifyWindow,
|
||||||
|
ScatterCompactToStrided,
|
||||||
|
scatter_compact_to_strided_into,
|
||||||
|
)
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode, ForwardMode
|
||||||
@@ -19,23 +36,6 @@ from sglang.srt.speculative.dspark_components.dspark_planner import (
|
|||||||
VerifyWindow,
|
VerifyWindow,
|
||||||
apply_logits_adjustments_strided,
|
apply_logits_adjustments_strided,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dspark_accept import (
|
|
||||||
AcceptGreedy,
|
|
||||||
AcceptSampling,
|
|
||||||
FinalizeAcceptLens,
|
|
||||||
SelectMixedAccept,
|
|
||||||
SoftmaxTemp,
|
|
||||||
accept_greedy_triton,
|
|
||||||
finalize_accept_lens_triton,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dspark_verify_window import (
|
|
||||||
BuildCommitInjectLayout,
|
|
||||||
BuildOutTokens,
|
|
||||||
BuildRaggedVerifyWindow,
|
|
||||||
RaggedVerifyWindow,
|
|
||||||
ScatterCompactToStrided,
|
|
||||||
scatter_compact_to_strided_into,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
|
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ configs (the block schedule and kernel config are chosen together).
|
|||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.layers.moe.moe_runner.triton_utils.inkling_moe import (
|
from sglang.kernels.ops.moe.inkling_moe import (
|
||||||
SMALL_M_BLOCK_SIZE_M,
|
SMALL_M_BLOCK_SIZE_M,
|
||||||
compute_grouped_gemm_metadata,
|
compute_grouped_gemm_metadata,
|
||||||
fused_moe_preprocess,
|
fused_moe_preprocess,
|
||||||
|
|||||||
@@ -16,16 +16,16 @@ from sglang.kernels.ops.attention import (
|
|||||||
dsv4_attn_metadata_kernels as attn_metadata_kernels,
|
dsv4_attn_metadata_kernels as attn_metadata_kernels,
|
||||||
)
|
)
|
||||||
from sglang.kernels.ops.speculative import ragged_verify_kernels
|
from sglang.kernels.ops.speculative import ragged_verify_kernels
|
||||||
from sglang.srt.speculative.dspark_components.dspark_planner import (
|
from sglang.kernels.ops.speculative.dspark import (
|
||||||
DSparkScheduleConfig,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.dspark_components.kernels import (
|
|
||||||
dspark_accept,
|
dspark_accept,
|
||||||
dspark_attn_metadata,
|
dspark_attn_metadata,
|
||||||
dspark_draft_model,
|
dspark_draft_model,
|
||||||
dspark_schedule,
|
dspark_schedule,
|
||||||
dspark_verify_window,
|
dspark_verify_window,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.speculative.dspark_components.dspark_planner import (
|
||||||
|
DSparkScheduleConfig,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
|
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|||||||
@@ -4,6 +4,9 @@ import unittest
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.kernels.ops.speculative.dspark.dspark_schedule import (
|
||||||
|
schedule_verify_lens_topk_from_survival,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.dspark_components.dspark_planner import (
|
from sglang.srt.speculative.dspark_components.dspark_planner import (
|
||||||
DSparkScheduleConfig,
|
DSparkScheduleConfig,
|
||||||
HostConfidenceBudgetPlanner,
|
HostConfidenceBudgetPlanner,
|
||||||
@@ -15,9 +18,6 @@ from sglang.srt.speculative.dspark_components.dspark_sps import (
|
|||||||
SpsAdditiveCostTable,
|
SpsAdditiveCostTable,
|
||||||
SpsCostTable,
|
SpsCostTable,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.dspark_components.kernels.dspark_schedule import (
|
|
||||||
schedule_verify_lens_topk_from_survival,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
|
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|||||||
Reference in New Issue
Block a user