Clean deprecated DeepSeek V4 Environs (#34926)

This commit is contained in:
Baizhou Zhang
2026-08-17 16:07:00 -07:00
committed by GitHub
parent b42abbb1ba
commit bc312d185d
28 changed files with 187 additions and 639 deletions
@@ -15,10 +15,7 @@ def validate_deepseek_v4_mega_moe_token_budget(
server_args: ServerArgs,
) -> None:
"""Ensure the DSV4 prefill budget fits MegaMoE's per-rank buffer."""
mega_moe_enabled = (
server_args.moe_a2a_backend == "megamoe"
or envs.SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE.get()
)
mega_moe_enabled = server_args.moe_a2a_backend == "megamoe"
if not mega_moe_enabled or server_args.disaggregation_mode == "decode":
# decode node will skip the check because decode bs is not relevant with --chunk-prefill-size
return
@@ -2547,12 +2547,6 @@ def _a2a_backend_overrides(view: Any) -> dict:
"requires the DeepEP or MegaMOE backend."
)
moe_a2a_backend = "deepep"
if envs.SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE.get() and moe_a2a_backend != "megamoe":
moe_a2a_backend = "megamoe"
logger.info(
"SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE is set, "
"auto-configuring --moe-a2a-backend megamoe."
)
if moe_a2a_backend != view.moe_a2a_backend:
return {"moe_a2a_backend": moe_a2a_backend}
return {}
+15 -36
View File
@@ -594,6 +594,10 @@ class Envs:
SGLANG_ENABLE_UNIFIED_RADIX_TREE = EnvBool(False)
# Registered TreeCore backend serving the unified radix cache.
SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND = EnvStr("python")
# TODO(DSV4): @ispobock this has bug on main branch when retract
SGLANG_OPT_SWA_RADIX_CACHE_COMPACT = EnvBool(False)
SGLANG_OPT_SWA_SPLIT_LEAF_ON_INSERT = EnvBool(False)
SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW = EnvBool(False)
# ===================================================================
# PD disaggregation runtime
@@ -796,6 +800,11 @@ class Envs:
SGLANG_HACK_FLASHMLA_BACKEND = EnvStr("tilelang")
SGLANG_USE_AITER_FP8_PER_TOKEN = EnvBool(False)
# DSV4 Aiter flags
SGLANG_OPT_USE_AITER_SILU_MUL = EnvBool(False)
SGLANG_OPT_USE_FUSED_QK_NORM_ROPE = EnvBool(True)
SGLANG_OPT_USE_AITER_INDEXER = EnvBool(False)
# ===================================================================
# Apple Silicon and MLX
# ===================================================================
@@ -1015,7 +1024,6 @@ class Envs:
# ===================================================================
# DeepGEMM Mega MoE
# ===================================================================
SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE = EnvBool(False)
SGLANG_OPT_DEEPGEMM_MEGA_MOE_NUM_MAX_TOKENS_PER_RANK = EnvInt(8192)
# When set, the mega-MoE x slot is packed E2M1 (FP4) instead of FP8 E4M3.
# Halves symm-buffer footprint and unlocks the MXF4 mainloop downstream.
@@ -1027,13 +1035,11 @@ class Envs:
# SGLANG_OPT_DEEPGEMM_MEGA_MOE_USE_FP4_ACTS is also set; DeepGEMM asserts
# this combination on the host side.
SGLANG_OPT_DEEPGEMM_MEGA_MOE_USE_MXF4_KIND = EnvBool(False)
SGLANG_OPT_FIX_MEGA_MOE_MEMORY = EnvBool(False)
# ===================================================================
# Top-k kernels
# ===================================================================
SGLANG_OPT_USE_FUSED_HASH_TOPK = EnvBool(True)
SGLANG_OPT_USE_JIT_KERNEL_FUSED_TOPK = EnvBool(True)
# Opt-in: route DeepSeek-V3 grouped topk through the unified Triton router
# instead of the flashinfer/AOT grouped kernels. Off by default (flashinfer is
# the tuned production path); the Triton path is bit-exact on DeepSeek-V3.2 e2e
@@ -1237,22 +1243,13 @@ class Envs:
SGLANG_CRASH_ON_NUMA_BIND_FAILURE = EnvBool(False)
# ===================================================================
# DeepSeek V4 - model and quantization
# DeepSeek V4
# ===================================================================
SGLANG_OPT_DPSK_V4_RADIX = EnvBool(True)
SGLANG_OPT_USE_OLD_COMPRESSOR = EnvBool(False)
SGLANG_OPT_USE_TRITON_SWA_PREPARE = EnvBool(True)
SGLANG_OPT_USE_AITER_MHC_PRE = EnvBool(True)
SGLANG_OPT_USE_AITER_MHC_POST = EnvBool(True)
SGLANG_OPT_USE_AITER_SILU_MUL = EnvBool(False)
SGLANG_OPT_USE_FUSED_COMPRESS = EnvBool(False)
SGLANG_OPT_USE_FUSED_COMPRESS_TRITON = EnvBool(False)
SGLANG_OPT_USE_FUSED_QK_NORM_ROPE = EnvBool(True)
SGLANG_OPT_USE_FUSED_CLAMP_ACT_MUL = EnvBool(True)
SGLANG_ENABLE_NVFP4_GEMM_SWIGLU_FUSION = EnvBool(True)
SGLANG_FIX_MTP_HC_HIDDEN = EnvBool(False)
# Model and Quantization
# Set False when using FP4-to-FP8 converted DeepSeek V4 checkpoint.
SGLANG_DSV4_FP4_EXPERTS = EnvBool(True)
# Set True to dequantize the FP4 experts to FP8 at runtime
SGLANG_DSV4_FP4_DEQUANT = EnvBool(False)
# Flash-0731 also accepts "low"; the active profile is checkpoint-resolved.
SGLANG_DSV4_REASONING_EFFORT = EnvStr("")
@@ -1260,18 +1257,13 @@ class Envs:
# trainer-side QAT and the DSA-CP path) instead of fp32 registers.
SGLANG_DSV4_USE_BF16_KV_QUANT_SOURCE = EnvBool(False)
# ===================================================================
# DeepSeek V4 - kernels and indexer
# ===================================================================
# Kernels and indexer
SGLANG_OPT_DEEPGEMM_HC_PRENORM = EnvBool(True)
SGLANG_OPT_USE_TILELANG_MHC_PRE = EnvBool(True)
SGLANG_OPT_USE_TILELANG_MHC_POST = EnvBool(True)
SGLANG_OPT_USE_FLASHINFER_MHC = EnvBool(False)
SGLANG_DSV4_MHC_PREWARM = EnvBool(True)
SGLANG_OPT_USE_TRITON_FUSED_MHC = EnvBool(True)
SGLANG_OPT_FUSE_MHC_POST_PRE = EnvBool(False)
SGLANG_OPT_USE_TILELANG_INDEXER = EnvBool(False)
SGLANG_OPT_USE_AITER_INDEXER = EnvBool(False)
SGLANG_OPT_DSV4_NONPAGED_INDEXER = EnvBool(True)
# Per-rank local query rows (after DP-attention sharding when enabled),
# not request ISL.
@@ -1281,26 +1273,13 @@ class Envs:
SGLANG_EXPERIMENTAL_ONLINE_C128_MTP = EnvBool(False)
SGLANG_DSV4_COMPRESS_STATE_DTYPE = EnvStr("float32")
SGLANG_FP8_PAGED_MQA_LOGITS_TORCH = EnvBool(False)
SGLANG_TOPK_TRANSFORM_512_TORCH = EnvBool(False)
SGLANG_OPT_FLASHMLA_SPARSE_PREFILL = EnvBool(True)
# ===================================================================
# DeepSeek V4 - cache, GEMM, and distributed
# ===================================================================
# TODO(DSV4): @ispobock this has bug on main branch when retract
SGLANG_OPT_SWA_RADIX_CACHE_COMPACT = EnvBool(False)
SGLANG_OPT_SWA_SPLIT_LEAF_ON_INSERT = EnvBool(False)
SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW = EnvBool(False)
# cache, GEMM, and distributed
SGLANG_OPT_FP8_WO_A_GEMM = EnvBool(True)
SGLANG_OPT_BF16_FP32_GEMM_ALGO = EnvStr("cublas")
SGLANG_OPT_USE_JIT_EP_ACTIVATION = EnvBool(True)
SGLANG_OPT_FUSE_WQA_WKV = EnvBool(True)
SGLANG_OPT_SWIGLU_CLAMP_FUSION = EnvBool(True)
SGLANG_OPT_USE_FUSED_STORE_CACHE = EnvBool(True)
SGLANG_OPT_USE_JIT_NORM = EnvBool(True)
SGLANG_OPT_USE_MULTI_STREAM_OVERLAP = EnvBool(True)
SGLANG_PREP_IN_CUDA_GRAPH = EnvBool(True)
SGLANG_DSV4_FIX_TP_ATTN_A2A_SCATTER = EnvBool(True)
# ===================================================================
# Inkling
@@ -25,9 +25,6 @@ from sglang.kernels.ops.attention.dsv4.metadata_kernel import (
init_compression_metadata as _init_compression_metadata_triton,
)
from sglang.kernels.ops.attention.dsv4.online_c128_mtp import OnlineC128MTPController
from sglang.kernels.ops.attention.dsv4.quant_k_cache import (
quant_to_nope_fp8_rope_bf16_pack_triton,
)
from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import (
BuildCausalSwaPageIndices,
BuildPageTablePositions,
@@ -68,7 +65,6 @@ from sglang.srt.runtime_context import get_parallel, get_spec
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
from sglang.srt.speculative.ragged_verify import (
RaggedVerifyMode,
compute_ragged_extend_lengths,
compute_target_verify_graph_key,
compute_uniform_extend_lengths,
read_ragged_verify_mode,
@@ -586,9 +582,7 @@ class DeepseekV4AttnBackend(
self.sparse_prefill_workspace = SparsePrefillWorkspace(self.device)
spec_alg = model_runner.spec_algorithm
self.needs_cpu_seq_lens = not spec_alg.is_dspark() and (
not _is_cuda
or not envs.SGLANG_PREP_IN_CUDA_GRAPH.get()
or self.online_c128_mtp.enabled()
not _is_cuda or self.online_c128_mtp.enabled()
)
self.is_dspark_draft = model_runner.is_draft_worker and spec_alg.is_dspark()
@@ -695,38 +689,10 @@ class DeepseekV4AttnBackend(
req_pool_indices.shape[0] == seq_lens.shape[0] == out_cache_loc.shape[0]
), f"{req_pool_indices.shape=} {seq_lens.shape=} {out_cache_loc.shape=}"
if envs.SGLANG_PREP_IN_CUDA_GRAPH.get():
return DSV4RawDecodeMetadata(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc,
)
core_attn_metadata = self.make_core_attn_metadata(
req_to_token=self.req_to_token,
req_pool_indices_repeated=req_pool_indices,
seq_lens_casual=seq_lens,
max_seq_len=max_seq_len,
out_loc=out_cache_loc,
need_compress=True,
)
indexer_metadata = self.init_forward_metadata_indexer(core_attn_metadata)
create = functools.partial(
create_paged_compressor_data,
is_prefill=False,
token_to_kv_pool=self.token_to_kv_pool,
req_to_token=self.req_to_token,
return DSV4RawDecodeMetadata(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
)
return DSV4Metadata(
core_attn_metadata,
indexer_metadata,
c4_compress_metadata=create(compress_ratio=4),
c128_compress_metadata=create(compress_ratio=128),
out_cache_loc=out_cache_loc,
)
def init_forward_metadata_prefill(
@@ -845,108 +811,44 @@ class DeepseekV4AttnBackend(
online_c128_state_slot_offset: int = 0,
ragged_layout: Optional[RaggedVerifyLayout] = None,
) -> Union[DSV4Metadata, DSV4RawVerifyMetadata]:
if envs.SGLANG_PREP_IN_CUDA_GRAPH.get():
assert out_cache_loc is not None
bs = len(seq_lens)
if self.needs_cpu_seq_lens:
assert seq_lens_cpu is not None
seq_lens_cpu_list = seq_lens_cpu.tolist()
else:
seq_lens_cpu_list = None
if ragged_layout is None:
self.extend_seq_lens_buffer[:bs].fill_(
self.speculative_num_draft_tokens
)
extend_seq_lens = self.extend_seq_lens_buffer[:bs]
extend_start_loc = None
verify_lens = None
total_verify_tokens = self.speculative_num_draft_tokens * bs
else:
self.extend_seq_lens_buffer[:bs].copy_(ragged_layout.verify_lens)
self.extend_start_loc_buffer[:bs].copy_(ragged_layout.extend_start_loc)
extend_seq_lens = self.extend_seq_lens_buffer[:bs]
extend_start_loc = self.extend_start_loc_buffer[:bs]
verify_lens = self.extend_seq_lens_buffer[:bs]
total_verify_tokens = ragged_layout.graph_num_tokens
return DSV4RawVerifyMetadata(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc,
extend_seq_lens=extend_seq_lens,
seq_lens_cpu=seq_lens_cpu_list,
c128_compress_metadata=self._make_target_verify_c128_metadata(
req_pool_indices,
seq_lens,
seq_lens_cpu_list,
extend_seq_lens,
use_prefill_cuda_graph,
online_c128_state_slot_offset,
),
extend_start_loc=extend_start_loc,
verify_lens=verify_lens,
total_verify_tokens=total_verify_tokens,
)
assert out_cache_loc is not None
bs = len(seq_lens)
if self.needs_cpu_seq_lens:
assert seq_lens_cpu is not None
seq_lens_cpu_list = seq_lens_cpu.tolist()
else:
seq_lens_cpu_list = (
seq_lens_cpu.tolist() if seq_lens_cpu is not None else seq_lens.tolist()
)
return self.init_forward_metadata_target_verify_old(
max_seq_len=max_seq_len,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu_list,
out_cache_loc=out_cache_loc,
use_prefill_cuda_graph=use_prefill_cuda_graph,
online_c128_state_slot_offset=online_c128_state_slot_offset,
ragged_layout=ragged_layout,
)
def init_forward_metadata_target_verify_old(
self,
max_seq_len: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_cpu: Optional[List[int]] = None,
out_cache_loc: Optional[torch.Tensor] = None,
use_prefill_cuda_graph: bool = False,
online_c128_state_slot_offset: int = 0,
ragged_layout: Optional[RaggedVerifyLayout] = None,
) -> DSV4Metadata:
seq_lens_cpu_list = None
if ragged_layout is None:
lengths = compute_uniform_extend_lengths(
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
extend_len=self.speculative_num_draft_tokens,
)
extend_seq_lens = self._move_to_device(lengths.extend_seq_lens_cpu)
self.extend_seq_lens_buffer[:bs].fill_(self.speculative_num_draft_tokens)
extend_seq_lens = self.extend_seq_lens_buffer[:bs]
extend_start_loc = None
verify_lens = None
total_verify_tokens = self.speculative_num_draft_tokens * bs
else:
lengths = compute_ragged_extend_lengths(
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
ragged_layout=ragged_layout,
)
extend_seq_lens = ragged_layout.verify_lens
seq_lens = lengths.seq_lens_extended
seq_lens_cpu = lengths.seq_lens_cpu_extended
extend_seq_lens_cpu = lengths.extend_seq_lens_cpu
num_tokens = lengths.num_tokens
extend_start_loc = lengths.extend_start_loc
if out_cache_loc is None:
out_cache_loc = seq_lens.new_zeros(num_tokens)
return self.init_forward_metadata_prefill(
max_seq_len=max_seq_len,
self.extend_seq_lens_buffer[:bs].copy_(ragged_layout.verify_lens)
self.extend_start_loc_buffer[:bs].copy_(ragged_layout.extend_start_loc)
extend_seq_lens = self.extend_seq_lens_buffer[:bs]
extend_start_loc = self.extend_start_loc_buffer[:bs]
verify_lens = self.extend_seq_lens_buffer[:bs]
total_verify_tokens = ragged_layout.graph_num_tokens
return DSV4RawVerifyMetadata(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
out_cache_loc=out_cache_loc,
num_tokens=num_tokens,
extend_seq_lens=extend_seq_lens,
extend_seq_lens_cpu=extend_seq_lens_cpu,
seq_lens_cpu=seq_lens_cpu_list,
c128_compress_metadata=self._make_target_verify_c128_metadata(
req_pool_indices,
seq_lens,
seq_lens_cpu_list,
extend_seq_lens,
use_prefill_cuda_graph,
online_c128_state_slot_offset,
),
extend_start_loc=extend_start_loc,
need_compress=True,
use_prefill_cuda_graph=use_prefill_cuda_graph,
online_c128_state_slot_offset=online_c128_state_slot_offset,
verify_lens=verify_lens,
total_verify_tokens=total_verify_tokens,
)
def init_forward_metadata_dspark_draft_block(
@@ -1627,19 +1529,11 @@ class DeepseekV4AttnBackend(
self, layer_id: int, swa_k: torch.Tensor, forward_batch: ForwardBatch
) -> None:
swa_loc = self.get_swa_out_cache_loc(forward_batch)
if envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get():
self.token_to_kv_pool.set_swa_key_buffer_radix_fused(
layer_id=layer_id,
swa_loc=swa_loc,
cache_k=swa_k,
)
else:
swa_k_pack = quant_to_nope_fp8_rope_bf16_pack_triton(swa_k)
self.token_to_kv_pool.set_swa_key_buffer_radix(
layer_id=layer_id,
swa_loc=swa_loc,
cache_nope_fp8_rope_bf16_pack=swa_k_pack,
)
self.token_to_kv_pool.set_swa_key_buffer_radix_fused(
layer_id=layer_id,
swa_loc=swa_loc,
cache_k=swa_k,
)
def forward(
self,
@@ -21,9 +21,6 @@ import torch.nn.functional as F
from sglang.kernels.ops.attention.dsv4.metadata_kernel import (
init_compression_metadata as _init_compression_metadata_triton,
)
from sglang.kernels.ops.attention.dsv4.quant_k_cache import (
quant_to_nope_fp8_rope_bf16_pack_triton,
)
from sglang.srt.environ import envs
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.dsv4.compressor_v2 import (
@@ -505,39 +502,10 @@ class DeepseekV4HipRadixBackend(
req_pool_indices.shape[0] == seq_lens.shape[0] == out_cache_loc.shape[0]
), f"{req_pool_indices.shape=} {seq_lens.shape=} {out_cache_loc.shape=}"
if envs.SGLANG_PREP_IN_CUDA_GRAPH.get():
return DSV4RawDecodeMetadata(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc,
)
core_attn_metadata = self.make_core_attn_metadata(
req_to_token=self.req_to_token,
req_pool_indices_repeated=req_pool_indices,
seq_lens_casual=seq_lens,
max_seq_len=max_seq_len,
out_loc=out_cache_loc,
need_compress=True,
)
self._attach_unified_kv_decode_streams(core_attn_metadata, req_pool_indices)
indexer_metadata = self.init_forward_metadata_indexer(core_attn_metadata)
create = functools.partial(
create_paged_compressor_data,
is_prefill=False,
token_to_kv_pool=self.token_to_kv_pool,
req_to_token=self.req_to_token,
return DSV4RawDecodeMetadata(
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
)
return DSV4Metadata(
core_attn_metadata,
indexer_metadata,
c4_compress_metadata=create(compress_ratio=4),
c128_compress_metadata=create(compress_ratio=128),
out_cache_loc=out_cache_loc,
)
def init_forward_metadata_prefill(
@@ -658,8 +626,7 @@ class DeepseekV4HipRadixBackend(
seq_lens_cpu: Optional[List[int]] = None,
ragged_layout=None,
) -> Union[DSV4Metadata, DSV4RawVerifyMetadata]:
# HIP path: build target-verify metadata eagerly even when
# SGLANG_PREP_IN_CUDA_GRAPH is enabled. The raw/lazy-upgrade route can
# HIP path: build target-verify metadata eagerly. The raw/lazy-upgrade route can
# hit planner invariants during graph capture for DSV4+EAGLE.
if seq_lens_cpu is None:
seq_lens_cpu = seq_lens.tolist()
@@ -1495,19 +1462,11 @@ class DeepseekV4HipRadixBackend(
self, layer_id: int, swa_k: torch.Tensor, forward_batch: ForwardBatch
) -> None:
swa_loc = self.get_swa_out_cache_loc(forward_batch)
if envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get():
self.token_to_kv_pool.set_swa_key_buffer_radix_fused(
layer_id=layer_id,
swa_loc=swa_loc,
cache_k=swa_k,
)
else:
swa_k_pack = quant_to_nope_fp8_rope_bf16_pack_triton(swa_k)
self.token_to_kv_pool.set_swa_key_buffer_radix(
layer_id=layer_id,
swa_loc=swa_loc,
cache_nope_fp8_rope_bf16_pack=swa_k_pack,
)
self.token_to_kv_pool.set_swa_key_buffer_radix_fused(
layer_id=layer_id,
swa_loc=swa_loc,
cache_k=swa_k,
)
def forward(
self,
@@ -1,21 +1,15 @@
from __future__ import annotations
import os
from functools import cached_property
from typing import TYPE_CHECKING, Any
import torch
import torch.nn as nn
from sglang.kernels.ops.attention.deepseek_v4_rope import (
apply_rotary_emb_triton,
fused_norm_rope_inplace_triton,
fused_softmax_pool_triton,
)
from sglang.kernels.ops.attention.dsv4.fused_compress_triton import (
fused_ape_pool_norm_rope,
)
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
from sglang.srt.layers.attention.nsa.nsa_indexer import rotate_activation
@@ -59,22 +53,6 @@ class CompressorHip(_CompressorBase):
self.norm = DeepseekRefRMSNorm(self.head_dim, eps=self.norm.variance_epsilon)
self._freqs_cis_real: torch.Tensor | None = None
@cached_property
def use_fused_compress(self) -> bool:
return envs.SGLANG_OPT_USE_FUSED_COMPRESS.get()
@cached_property
def use_hip_fused_compress(self) -> bool:
return envs.SGLANG_OPT_USE_FUSED_COMPRESS.get()
@cached_property
def use_fused_compress_triton(self) -> bool:
# The fused Triton kernel only benefits non-overlap (HCA, ratio=128)
# but HCA's K=128 loop is too sequential to outperform batched ops.
# CSA (overlap=True) has a reshape/overlap-transform semantic mismatch.
# Disabled until a tiled kernel for CSA overlap is implemented.
return False
def _get_states(
self,
forward_batch: ForwardBatch,
@@ -237,16 +215,10 @@ class CompressorHip(_CompressorBase):
beg_idx = prefix_lens[i] // self.ratio * self.ratio
end_idx = (prefix_lens[i] + extend_lens[i]) // self.ratio * self.ratio
if self.use_hip_fused_compress:
kv_compressed = fused_softmax_pool_triton(
kv_and_score_to_compress.kv_score,
kv_and_score_to_compress._item_size,
)
else:
kv_compressed = (
kv_and_score_to_compress.kv
* kv_and_score_to_compress.score.softmax(dim=1)
).sum(dim=1)
kv_compressed = fused_softmax_pool_triton(
kv_and_score_to_compress.kv_score,
kv_and_score_to_compress._item_size,
)
assert kv_compressed.dtype == torch.float32
@@ -254,15 +226,9 @@ class CompressorHip(_CompressorBase):
assert freqs_cis.size(0) == kv_compressed.size(
0
), f"{freqs_cis.shape=} {kv_compressed.shape=}"
if self.use_hip_fused_compress:
fused_norm_rope_inplace_triton(
kv_compressed, self.norm.weight, self.norm.eps, freqs_cis
)
else:
kv_compressed = self.norm(kv_compressed)
apply_rotary_emb_triton(
kv_compressed[..., -self.rope_head_dim :], freqs_cis
)
fused_norm_rope_inplace_triton(
kv_compressed, self.norm.weight, self.norm.eps, freqs_cis
)
del beg_idx, end_idx
if self.rotate:
@@ -343,34 +309,6 @@ class CompressorHip(_CompressorBase):
).view(-1, self.ratio, self.coff * self.head_dim)
bs = seq_lens.size(0)
if self.use_fused_compress_triton and not self.overlap:
# Fused path for non-overlap (HCA, ratio=128, coff=1):
# APE + softmax-pool + norm + RoPE in one kernel.
# Overlap (CSA) is excluded because the overlap_transform_decode
# rearranges A/B halves across the coff dimension in a way
# that simple reshape cannot replicate correctly.
raw = kv_and_score_to_compress.kv_score
gathered = raw.reshape(bs, self.ratio, raw.shape[-1]).contiguous()
comp_positions = (seq_lens - 1) // self.ratio * self.ratio
freqs_real_table = self._get_freqs_cis_real()
freqs_batch = freqs_real_table[comp_positions]
kv_compressed = fused_ape_pool_norm_rope(
kv_score_gathered=gathered,
ape=self.ape,
rms_weight=self.norm.weight,
rms_eps=self.norm.eps,
freqs_cis_real=freqs_batch,
head_dim=self.head_dim,
rope_head_dim=self.rope_head_dim,
ratio=self.ratio,
overlap=self.overlap,
)
if self.rotate:
kv_compressed = rotate_activation(kv_compressed)
return kv_compressed
# Unfused reference path
kv_and_score_to_compress.score.add_(self.ape.unsqueeze(0))
@@ -387,27 +325,14 @@ class CompressorHip(_CompressorBase):
bs, self.ratio * self.coff, self.head_dim
)
if self.use_hip_fused_compress:
kv_compressed = fused_softmax_pool_triton(
kv_and_score_to_compress.kv_score,
kv_and_score_to_compress._item_size,
)
else:
kv_compressed = (
kv_and_score_to_compress.kv
* kv_and_score_to_compress.score.softmax(dim=1)
).sum(dim=1)
if self.use_hip_fused_compress:
freqs_cis = self._init_freqs_cis_per_decode_step(forward_batch, seq_lens)
fused_norm_rope_inplace_triton(
kv_compressed, self.norm.weight, self.norm.eps, freqs_cis
)
else:
kv_compressed = self.norm(kv_compressed)
freqs_cis = self.freqs_cis[(seq_lens - 1) // self.ratio * self.ratio]
apply_rotary_emb_triton(
kv_compressed[..., -self.rope_head_dim :], freqs_cis
)
kv_compressed = fused_softmax_pool_triton(
kv_and_score_to_compress.kv_score,
kv_and_score_to_compress._item_size,
)
freqs_cis = self._init_freqs_cis_per_decode_step(forward_batch, seq_lens)
fused_norm_rope_inplace_triton(
kv_compressed, self.norm.weight, self.norm.eps, freqs_cis
)
if self.rotate:
kv_compressed = rotate_activation(kv_compressed)
@@ -455,12 +380,9 @@ class CompressorHip(_CompressorBase):
forward_batch: ForwardBatch,
attn_backend: AttentionBackend,
) -> torch.Tensor:
if self.use_fused_compress and (
envs.SGLANG_OPT_DPSK_V4_RADIX.get()
and (
forward_batch.forward_mode.is_decode()
or forward_batch.forward_mode.is_extend_without_speculative()
)
if (
forward_batch.forward_mode.is_decode()
or forward_batch.forward_mode.is_extend_without_speculative()
):
return self.compress_fused(
kv_score, forward_batch, attn_backend=attn_backend
@@ -6,7 +6,6 @@ import torch
import torch.nn as nn
from sglang.kernels.fused_op import BaseFusedOp
from sglang.kernels.ops.attention.dsa.triton_kernel import act_quant
from sglang.kernels.ops.attention.dsv4 import (
linear_bf16_fp32,
triton_create_paged_compress_data,
@@ -17,9 +16,6 @@ from sglang.kernels.ops.attention.dsv4.compress_old import (
compress_forward,
compress_fused_norm_rope_inplace,
)
from sglang.kernels.ops.attention.dsv4.quant_k_cache import (
quant_to_nope_fp8_rope_bf16_pack_triton,
)
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
from sglang.srt.environ import envs
from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp
@@ -185,15 +181,12 @@ class CompressorBackendMixin:
)
if out_loc.shape[0] > new_compressed_kv.shape[0]:
out_loc = out_loc[: new_compressed_kv.shape[0]]
if envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get():
token_to_kv_pool.set_extra_key_buffer_fused(
layer_id=layer_id,
loc=out_loc,
cache_k=new_compressed_kv,
)
else:
pack = quant_to_nope_fp8_rope_bf16_pack_triton(new_compressed_kv.bfloat16())
token_to_kv_pool.set_extra_key_buffer(layer_id, out_loc, pack)
token_to_kv_pool.set_extra_key_buffer_fused(
layer_id=layer_id,
loc=out_loc,
cache_k=new_compressed_kv,
)
def forward_indexer_compressor(
self,
@@ -217,22 +210,12 @@ class CompressorBackendMixin:
loc=out_loc,
cache_k=new_compressed_kv,
)
elif envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get():
else:
token_to_kv_pool.set_index_k_fused(
layer_id=layer_id,
loc=out_loc,
cache_k=new_compressed_kv,
)
else:
new_compressed_kv_fp8, new_compressed_kv_scale = act_quant(
new_compressed_kv
)
token_to_kv_pool.set_index_k_scale_buffer(
layer_id=layer_id,
loc=out_loc,
index_k=new_compressed_kv_fp8,
index_k_scale=new_compressed_kv_scale,
)
def is_overlap_compress(compress_ratio: int) -> bool:
@@ -218,56 +218,45 @@ class CompressorBackendMixin:
is_unified_kv_triton,
)
if _is_hip and not envs.SGLANG_OPT_USE_JIT_NORM.get():
self._forward_unified_hip(
token_to_kv_pool=token_to_kv_pool,
kv_score_input=kv_score_input,
state_pool=state_pool,
compressor=compressor,
layer_id=layer_id,
out_loc = self._get_out_loc(compressor.ratio)
use_fp4_indexer = (
compressor.is_in_indexer and self.enable_deepseek_v4_fp4_indexer
)
bf16_store = False
if compressor.is_in_indexer:
kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer(layer_id)
page_size = token_to_kv_pool.get_index_k_page_size()
elif is_unified_kv_triton():
kv_cache = token_to_kv_pool.get_unified_kv(layer_id)
page_size = 1
out_loc = getattr(
self.forward_metadata.core_metadata.unified,
f"c{compressor.ratio}_out_loc",
)
bf16_store = True
else:
out_loc = self._get_out_loc(compressor.ratio)
use_fp4_indexer = (
compressor.is_in_indexer and self.enable_deepseek_v4_fp4_indexer
)
bf16_store = False
if compressor.is_in_indexer:
kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer(layer_id)
page_size = token_to_kv_pool.get_index_k_page_size()
elif is_unified_kv_triton():
kv_cache = token_to_kv_pool.get_unified_kv(layer_id)
page_size = 1
out_loc = getattr(
self.forward_metadata.core_metadata.unified,
f"c{compressor.ratio}_out_loc",
)
bf16_store = True
else:
_, _, compress_kv_pool = token_to_kv_pool.layer_mapping[layer_id]
assert compress_kv_pool is not None
kv_cache = token_to_kv_pool.get_extra_key_buffer(layer_id)
page_size = token_to_kv_pool.get_extra_key_page_size(layer_id)
if hasattr(compress_kv_pool, "translate_loc_to_hisparse_device"):
out_loc = compress_kv_pool._translate_loc_to_hisparse_device(
out_loc
)
self._forward_compress_all_in_one(
kv_score_buffer=state_pool.kv_score_buffer.kv_score,
kv_score_input=kv_score_input,
ape=compressor.ape,
head_dim=compressor.head_dim,
norm=compressor.norm,
freqs_cis_cache=compressor.freqs_cis,
kv_cache=kv_cache.view(dtype=torch.uint8),
is_indexer=compressor.is_in_indexer,
rotate=compressor.rotate,
compress_ratio=compressor.ratio,
page_size=page_size,
out_loc=out_loc,
use_fp4_indexer=use_fp4_indexer,
bf16_store=bf16_store,
)
_, _, compress_kv_pool = token_to_kv_pool.layer_mapping[layer_id]
assert compress_kv_pool is not None
kv_cache = token_to_kv_pool.get_extra_key_buffer(layer_id)
page_size = token_to_kv_pool.get_extra_key_page_size(layer_id)
if hasattr(compress_kv_pool, "translate_loc_to_hisparse_device"):
out_loc = compress_kv_pool._translate_loc_to_hisparse_device(out_loc)
self._forward_compress_all_in_one(
kv_score_buffer=state_pool.kv_score_buffer.kv_score,
kv_score_input=kv_score_input,
ape=compressor.ape,
head_dim=compressor.head_dim,
norm=compressor.norm,
freqs_cis_cache=compressor.freqs_cis,
kv_cache=kv_cache.view(dtype=torch.uint8),
is_indexer=compressor.is_in_indexer,
rotate=compressor.rotate,
compress_ratio=compressor.ratio,
page_size=page_size,
out_loc=out_loc,
use_fp4_indexer=use_fp4_indexer,
bf16_store=bf16_store,
)
online_c128_mtp = getattr(self, "online_c128_mtp", None)
if online_c128_mtp is not None:
online_c128_mtp.write_prefix_states(
@@ -292,11 +281,7 @@ class CompressorBackendMixin:
from sglang.kernels.ops.attention.deepseek_v4_rope import (
fused_norm_rope_inplace_triton,
)
from sglang.kernels.ops.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
compress_ratio = compressor.ratio
head_dim = compressor.head_dim
@@ -364,32 +349,18 @@ class CompressorBackendMixin:
if kv_to_store.shape[0] == 0:
return
if envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get():
# fused kernel: BF16 in -> FP8 quant + paged scatter in one launch
if is_indexer:
token_to_kv_pool.set_index_k_fused(
layer_id=layer_id,
loc=out_loc_to_store,
cache_k=kv_to_store,
)
else:
token_to_kv_pool.set_extra_key_buffer_fused(
layer_id=layer_id,
loc=out_loc_to_store,
cache_k=kv_to_store,
)
if is_indexer:
token_to_kv_pool.set_index_k_fused(
layer_id=layer_id,
loc=out_loc_to_store,
cache_k=kv_to_store,
)
else:
if is_indexer:
kv_fp8, kv_scale = act_quant(kv_to_store)
token_to_kv_pool.set_index_k_scale_buffer(
layer_id=layer_id,
loc=out_loc_to_store,
index_k=kv_fp8,
index_k_scale=kv_scale,
)
else:
pack = quant_to_nope_fp8_rope_bf16_pack_triton(kv_to_store.bfloat16())
token_to_kv_pool.set_extra_key_buffer(layer_id, out_loc_to_store, pack)
token_to_kv_pool.set_extra_key_buffer_fused(
layer_id=layer_id,
loc=out_loc_to_store,
cache_k=kv_to_store,
)
# NOTE: alias for backward compatibility
forward_indexer_compressor = forward_unified
@@ -587,10 +587,7 @@ class C4IndexerBackendMixin:
ks = torch.zeros_like(ke)
# SGL Top-K synthesizes sequential indices for trivial rows without
# reading logits, so DeepGEMM can receive an empty range for them.
if (
self.dsa_topk_backend.is_sgl_kernel()
and not envs.SGLANG_TOPK_TRANSFORM_512_TORCH.get()
):
if self.dsa_topk_backend.is_sgl_kernel():
ke = torch.where(ke - ks > c4_indexer.index_topk, ke, ks)
c4_page_size = indexer_metadata.c4_page_size
max_seqlen_k = (final_c4_len + c4_page_size - 1) // c4_page_size * c4_page_size
@@ -811,10 +808,7 @@ class C4IndexerBackendMixin:
elif core_metadata.c4_sparse_raw_indices is not None:
raw_indices = core_metadata.c4_sparse_raw_indices
if (
envs.SGLANG_TOPK_TRANSFORM_512_TORCH.get()
or self.dsa_topk_backend.is_torch()
):
if self.dsa_topk_backend.is_torch():
topk_transform_512_pytorch_vectorized(
logits,
c4_seq_lens,
+16 -25
View File
@@ -320,7 +320,6 @@ def _transpose_mega_moe_sf_for_utccp(sf: torch.Tensor) -> torch.Tensor:
def build_mega_moe_experts_weights(experts) -> None:
from deep_gemm import (
transform_sf_into_required_layout,
transform_weights_for_mega_moe,
)
if getattr(experts, "_mega_moe_weights_built", False):
@@ -353,31 +352,23 @@ def build_mega_moe_experts_weights(experts) -> None:
disable_ue8m0_cast=False,
)
if envs.SGLANG_OPT_FIX_MEGA_MOE_MEMORY.get():
# Build the interleaved L1 weight + scale once; share the weight buffer
# between `w13_weight.data` (normal deep-ep path) and `mega_l1_weights[0]`
# (mega moe path). Mega moe additionally needs a UTCCP-transposed scale;
# the deep-ep path consumes the non-transposed interleaved scale and a
# swizzle-aware activation kernel. L2 weight is untouched by the mega
# transform, so the existing `w2_weight.data` is shared directly.
w13_interleaved, w13_sf_interleaved = _interleave_mega_moe_l1_weights(
(w13, w13_sf)
)
w13_sf_utccp = _transpose_mega_moe_sf_for_utccp(w13_sf_interleaved)
w2_sf_utccp = _transpose_mega_moe_sf_for_utccp(w2_sf)
# Build the interleaved L1 weight + scale once; share the weight buffer
# between `w13_weight.data` (normal deep-ep path) and `mega_l1_weights[0]`
# (mega moe path). Mega moe additionally needs a UTCCP-transposed scale;
# the deep-ep path consumes the non-transposed interleaved scale and a
# swizzle-aware activation kernel. L2 weight is untouched by the mega
# transform, so the existing `w2_weight.data` is shared directly.
w13_interleaved, w13_sf_interleaved = _interleave_mega_moe_l1_weights((w13, w13_sf))
w13_sf_utccp = _transpose_mega_moe_sf_for_utccp(w13_sf_interleaved)
w2_sf_utccp = _transpose_mega_moe_sf_for_utccp(w2_sf)
experts.w13_weight.data = w13_interleaved
experts.w13_weight_scale_inv.data = w13_sf_interleaved
experts.w2_weight_scale_inv.data = w2_sf
experts.w13_weight_scale_inv.format_ue8m0 = True
experts.w2_weight_scale_inv.format_ue8m0 = True
experts.w13_weight.data = w13_interleaved
experts.w13_weight_scale_inv.data = w13_sf_interleaved
experts.w2_weight_scale_inv.data = w2_sf
experts.w13_weight_scale_inv.format_ue8m0 = True
experts.w2_weight_scale_inv.format_ue8m0 = True
experts.mega_l1_weights = (experts.w13_weight.data, w13_sf_utccp)
experts.mega_l2_weights = (experts.w2_weight.data, w2_sf_utccp)
else:
l1_pair, l2_pair = transform_weights_for_mega_moe((w13, w13_sf), (w2, w2_sf))
experts.mega_l1_weights = l1_pair
experts.mega_l2_weights = l2_pair
experts.mega_l1_weights = (experts.w13_weight.data, w13_sf_utccp)
experts.mega_l2_weights = (experts.w2_weight.data, w2_sf_utccp)
experts._mega_moe_weights_built = True
+10 -36
View File
@@ -19,7 +19,6 @@ from typing import TYPE_CHECKING
import torch
from sglang.srt.environ import envs
from sglang.srt.models.deepseek_common.utils import _device_sm
if TYPE_CHECKING:
@@ -139,41 +138,16 @@ def build_sm90_mega_moe_experts_weights(experts) -> None:
f"expected {expected_k_groups_2} (k2={k2}, group_k={scale_group_k})"
)
if envs.SGLANG_OPT_FIX_MEGA_MOE_MEMORY.get():
w13_interleaved = _interleave_l1_weight_only(w13)
experts.w13_weight.data = w13_interleaved
experts.mega_l1_weights = (
experts.w13_weight.data,
experts.w13_weight_scale_inv.data,
)
experts.mega_l2_weights = (
experts.w2_weight.data,
experts.w2_weight_scale_inv.data,
)
else:
import deep_gemm
w13_sf = deep_gemm.transform_sf_into_required_layout(
w13_sf_fp32,
mn=n1,
k=k1,
recipe=(128, 128),
num_groups=num_groups,
disable_ue8m0_cast=True,
)
w2_sf = deep_gemm.transform_sf_into_required_layout(
w2_sf_fp32,
mn=n2,
k=k2,
recipe=(128, 128),
num_groups=num_groups,
disable_ue8m0_cast=True,
)
l1_pair, l2_pair = deep_gemm.transform_weights_for_mega_moe_sm90(
(w13, w13_sf), (w2, w2_sf)
)
experts.mega_l1_weights = l1_pair
experts.mega_l2_weights = l2_pair
w13_interleaved = _interleave_l1_weight_only(w13)
experts.w13_weight.data = w13_interleaved
experts.mega_l1_weights = (
experts.w13_weight.data,
experts.w13_weight_scale_inv.data,
)
experts.mega_l2_weights = (
experts.w2_weight.data,
experts.w2_weight_scale_inv.data,
)
experts._mega_moe_sm90_fp8_weights = True
experts._mega_moe_weights_built = True
@@ -4,7 +4,6 @@ import logging
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, List, Optional, Tuple
import einops
import torch
import triton
import triton.language as tl
@@ -30,7 +29,7 @@ from sglang.srt.layers.moe.moe_runner.base import (
register_post_permute,
register_pre_permute,
)
from sglang.srt.layers.moe.utils import MoeRunnerBackend
from sglang.srt.layers.moe.utils import MoeRunnerBackend, get_moe_a2a_backend
from sglang.srt.runtime_context import get_exec
from sglang.srt.utils import (
ceil_div,
@@ -61,7 +60,7 @@ _is_cuda = is_cuda()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
_is_musa = is_musa()
# Imported only for the SGLANG_OPT_FIX_MEGA_MOE_MEMORY=False fallback path.
if not (_is_npu or _is_hip) and _is_cuda:
from sglang.kernels.ops.activation.activation import (
silu_and_mul as _legacy_silu_and_mul,
@@ -253,11 +252,7 @@ class DeepGemmRunnerCore(MoeRunnerCore):
assert self.config.activation in ("silu", "situ")
assert self.config.is_gated
self.swiglu_limit = self.config.swiglu_limit
self.use_swizzle = False
if envs.SGLANG_OPT_FIX_MEGA_MOE_MEMORY.get():
assert envs.SGLANG_OPT_SWIGLU_CLAMP_FUSION.get()
assert envs.SGLANG_OPT_USE_JIT_EP_ACTIVATION.get()
self.use_swizzle = True
self.use_swizzle = get_moe_a2a_backend().is_megamoe()
def run(
self,
@@ -395,7 +390,7 @@ class DeepGemmRunnerCore(MoeRunnerCore):
scale_ue8m0=False,
)
del down_input
elif envs.SGLANG_OPT_FIX_MEGA_MOE_MEMORY.get():
elif self.use_swizzle:
swiglu_limit_arg: Optional[float] = self.swiglu_limit
down_input_fp8 = torch.empty(
@@ -423,9 +418,6 @@ class DeepGemmRunnerCore(MoeRunnerCore):
)
del gateup_output
else:
# Hacky byte-equal fallback that reproduces the optimize-branch
# code path exactly: bf16 silu_and_mul then a separate per-token
# group fp8 quant. Kept behind the mega-moe-memory flag.
from sglang.kernels.ops.quantization.fp8_kernel import (
sglang_per_token_group_quant_fp8,
)
@@ -637,25 +629,7 @@ class DeepGemmRunnerCore(MoeRunnerCore):
swiglu_limit_arg: Optional[float] = None
if self.swiglu_limit is not None:
# DeepSeek V4: clamped swiglu requires the DSV4 JIT EP activation.
assert (
envs.SGLANG_OPT_USE_JIT_EP_ACTIVATION.get()
), "DeepSeek V4 requires SGLANG_OPT_USE_JIT_EP_ACTIVATION=True"
if envs.SGLANG_OPT_SWIGLU_CLAMP_FUSION.get():
swiglu_limit_arg = self.swiglu_limit
else:
gateup_output = einops.rearrange(
gateup_output, "grp tok hidden -> (grp tok) hidden"
)
gateup_output = _apply_swiglu_limit(
gateup_output, swiglu_limit=self.swiglu_limit
)
gateup_output = einops.rearrange(
gateup_output,
"(grp tok) hidden -> grp tok hidden",
grp=num_groups,
)
swiglu_limit_arg = self.swiglu_limit
# Act.
if self.config.activation == "situ":
@@ -1351,9 +1325,6 @@ def _varlen_deep_gemm_silu_mul_quant(
# DSV4-specific activations (clamped swiglu, swizzled gate|up layout) stay
# on the DSV4 JIT kernel; it is the only implementation carrying them.
if swiglu_limit is not None or swizzle:
assert (
envs.SGLANG_OPT_USE_JIT_EP_ACTIVATION.get()
), "swiglu_limit / swizzle require SGLANG_OPT_USE_JIT_EP_ACTIVATION=True"
assert N % 4 == 0 and G % 4 == 0 and D // 8 >= E, (
"DSV4 JIT activation requires N % 4 == 0, G % 4 == 0 and "
f"D // 8 >= num_experts, got N={N} G={G} D={D} E={E}"
@@ -25,7 +25,6 @@ from sglang.srt.distributed import get_tp_group
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory,
)
from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
from sglang.srt.layers.moe.utils import get_moe_padding_size
@@ -689,20 +688,13 @@ def _fused_moe_kernel_sequence(
swiglu_limit_for_triton: Optional[float] = None
swiglu_limit_for_silu_and_mul_clamp: Optional[float] = None
if envs.SGLANG_OPT_SWIGLU_CLAMP_FUSION.get():
if filter_expert:
swiglu_limit_for_triton = swiglu_limit
else:
assert (
_is_cuda or _is_xpu
), "fused silu_and_mul_clamp kernel is CUDA/XPU only; HIP must disable SWIGLU_CLAMP_FUSION"
swiglu_limit_for_silu_and_mul_clamp = swiglu_limit
if filter_expert:
swiglu_limit_for_triton = swiglu_limit
else:
half = N // 2
intermediate_cache1[:, :half].clamp_(max=swiglu_limit)
intermediate_cache1[:, half:].clamp_(
min=-swiglu_limit, max=swiglu_limit
)
assert (
_is_cuda or _is_xpu
), "fused silu_and_mul_clamp kernel is CUDA/XPU only; HIP must disable SWIGLU_CLAMP_FUSION"
swiglu_limit_for_silu_and_mul_clamp = swiglu_limit
if not filter_expert:
if swiglu_limit_for_silu_and_mul_clamp is not None:
+4 -12
View File
@@ -875,7 +875,7 @@ def fused_topk(
num_token_non_padded=num_token_non_padded,
)
# ===== END TO BE REFACTORED ====
elif _is_cuda and envs.SGLANG_OPT_USE_JIT_KERNEL_FUSED_TOPK.get():
elif _is_cuda:
# Unified Triton router (subsumes the AOT topk_softmax CUDA kernel).
from sglang.kernels.ops.moe.moe_fused_gate import (
moe_fused_gate as _jit_moe_fused_gate,
@@ -915,7 +915,7 @@ def fused_topk(
topk_weights *= (
routed_scaling_factor if routed_scaling_factor is not None else 1.0
)
elif _is_cuda and envs.SGLANG_OPT_USE_JIT_KERNEL_FUSED_TOPK.get():
elif _is_cuda:
# Unified Triton router (subsumes the AOT topk_sigmoid CUDA kernel).
from sglang.kernels.ops.moe.moe_fused_gate import (
moe_fused_gate as _jit_moe_fused_gate,
@@ -2168,16 +2168,8 @@ def select_experts(
if scoring_func not in ("sqrtsoftplus", "sigmoid"):
assert not apply_routed_scaling_factor_on_output, "Not implemented"
# Keep sigmoid flag-off byte-identical: only use the JIT gate when the flag is on.
use_jit_fused_gate = envs.SGLANG_OPT_USE_JIT_KERNEL_FUSED_TOPK.get()
if scoring_func == "sqrtsoftplus" or (
scoring_func == "sigmoid" and use_jit_fused_gate
):
_biased_topk = (
biased_topk_jit_kernel_impl if use_jit_fused_gate else biased_topk_impl
)
topk_weights, topk_ids = _biased_topk(
if scoring_func == "sqrtsoftplus" or scoring_func == "sigmoid":
topk_weights, topk_ids = biased_topk_jit_kernel_impl(
hidden_states=hidden_states,
gating_output=router_logits,
correction_bias=correction_bias,
@@ -4,8 +4,6 @@ from typing import Optional, Tuple
import torch
import triton
from sglang.srt.environ import envs
logger = logging.getLogger(__name__)
_FUSED_HC_POST_PRE_M_THRESHOLD = 64
@@ -104,7 +102,6 @@ def try_fused_hc_post_pre(
if (
_TRITON_MHC_POST_PRE_RUNTIME_DISABLED
or not envs.SGLANG_OPT_USE_TRITON_FUSED_MHC.get()
or not is_gfx95_supported
or x.shape[0] == 0
or x.shape[0] > _FUSED_HC_POST_PRE_M_THRESHOLD
+2 -5
View File
@@ -305,9 +305,7 @@ class DeepseekV2MLP(nn.Module):
"Only silu is supported for now."
)
self.act_fn = SiluAndMul()
self.use_fused_clamp_act_mul = (
_is_hip and envs.SGLANG_OPT_USE_FUSED_CLAMP_ACT_MUL.get()
)
self.use_fused_clamp_act_mul = _is_hip
self._fused_clamp_fp8_checked = False
self._fused_clamp_use_fp8 = False
@@ -758,8 +756,7 @@ class DeepseekV2MoE(nn.Module):
fc1_n = self.shared_experts.gate_up_proj.output_size_per_partition
if (
envs.SGLANG_ENABLE_NVFP4_GEMM_SWIGLU_FUSION.get()
and is_sm100_supported()
is_sm100_supported()
and isinstance(
self.shared_experts.gate_up_proj.quant_method,
ModelOptFp4LinearMethod,
+3 -7
View File
@@ -1779,7 +1779,7 @@ class DeepseekV4DecoderLayer(nn.Module):
)
return y, post.squeeze(-1), comb, norm is not None
if _is_hip and envs.SGLANG_OPT_USE_AITER_MHC_PRE.get():
if _is_hip:
from aiter.ops.mhc import mhc_pre
post, comb, y = mhc_pre(
@@ -1860,7 +1860,7 @@ class DeepseekV4DecoderLayer(nn.Module):
return mhc_post(x, residual, post, comb)
elif _is_hip and envs.SGLANG_OPT_USE_AITER_MHC_POST.get():
elif _is_hip:
from aiter.ops.mhc import mhc_post
result = torch.empty_like(residual)
@@ -2036,7 +2036,6 @@ class DeepseekV4DecoderLayer(nn.Module):
)
_use_tp_attn_a2a_scatter = (
not _use_cp
and envs.SGLANG_DSV4_FIX_TP_ATTN_A2A_SCATTER.get()
and get_parallel().attn_tp_size > 1
and not get_moe_a2a_backend().is_none()
)
@@ -3211,10 +3210,7 @@ class DeepseekV4ForCausalLM(nn.Module):
if self._mhc_prewarmed_at_load:
return
self._mhc_prewarmed_at_load = True
if _is_npu or not (
envs.SGLANG_DSV4_MHC_PREWARM.get()
and envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get()
):
if _is_npu or not envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get():
return
layer = next(
(m for m in self.model.layers if isinstance(m, DeepseekV4DecoderLayer)),
-5
View File
@@ -5510,7 +5510,6 @@ class ServerArgs:
envs.SGLANG_OPT_USE_TILELANG_INDEXER.set(True)
elif is_hip():
envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.set(False)
envs.SGLANG_OPT_USE_FUSED_COMPRESS.set(True)
envs.SGLANG_OPT_FP8_WO_A_GEMM.set(False)
envs.SGLANG_OPT_USE_JIT_INDEXER_METADATA.set(False)
envs.SGLANG_OPT_USE_TOPK_V2.set(False)
@@ -6886,10 +6885,6 @@ class ServerArgs:
self.enforce_shared_experts_fusion = True
logger.info(f"Waterfill is enabled with moe_a2a_backend='{a2a_backend}'.")
if a2a_backend == "megamoe":
if not envs.SGLANG_OPT_FIX_MEGA_MOE_MEMORY.is_set():
envs.SGLANG_OPT_FIX_MEGA_MOE_MEMORY.set(True)
if a2a_backend == "deepep":
if self.moe_runner_backend == "flashinfer_cutedsl":
if self.deepep_mode == "auto":
@@ -147,7 +147,6 @@ class DSparkVerifyPlanner:
f"draft checkpoint that includes the confidence head, or run "
f"SGLANG_RAGGED_VERIFY_MODE=static."
)
self._require_prep_in_cuda_graph()
sps_table = build_sps_cost_table(
server_args=self.server_args,
verify_num_draft_tokens=self.verify_num_draft_tokens,
@@ -208,16 +207,6 @@ class DSparkVerifyPlanner:
"Pass a profiled --speculative-dspark-sps-table-path."
)
def _require_prep_in_cuda_graph(self) -> None:
if not envs.SGLANG_PREP_IN_CUDA_GRAPH.get():
raise ValueError(
f"DSpark ragged-verify mode {self._ragged_verify_mode.value!r} "
f"requires SGLANG_PREP_IN_CUDA_GRAPH=1 (the captured-graph prepare "
f"path). It is currently disabled, which would put per-step "
f"verify_lens_cpu host reads on the critical path. Set "
f"SGLANG_PREP_IN_CUDA_GRAPH=1 or run SGLANG_RAGGED_VERIFY_MODE=static."
)
@property
def carries_confidence(self) -> bool:
return self._confidence_head is not None