Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
25 commits
Select commit Hold shift + click to select a range
7c9cf23
[Perf] Tune H200 group32 FP8 dense GEMMs for TP8 decode
BBuf Sep 25, 2026
e75f407
Keep the original short-K GEMM configuration above batch 32
BBuf Sep 25, 2026
a49faa1
[Perf] Bound Hopper FP4 indexer work by live context lengths
BBuf Sep 25, 2026
5496c9e
Prepare the indexer top-k plan before scoring and test large graph ba…
BBuf Sep 25, 2026
63d795e
[Perf] Cache exact group32 weight expansions for Hopper prefill GEMMs
BBuf Sep 25, 2026
c15ba96
[Fix] Guard live updates of cached Hopper FP8 weight expansions
BBuf Sep 25, 2026
a0d87fb
[Perf] Use compensated BF16 mHC projections for Hopper prefill
BBuf Sep 25, 2026
a99be3e
[Fix] Reject inexact Hopper BF16 weight expansions
BBuf Sep 25, 2026
f8e12f6
[Perf] Fuse Hopper Marlin SwiGLU without changing activation rounding
BBuf Sep 25, 2026
2b4827d
[Perf] Fuse Hopper indexer top-k masking and slot mapping
BBuf Sep 25, 2026
0d3c45d
[Perf] Extend measured Hopper mHC fusion and compensated projection p…
BBuf Sep 25, 2026
2a7aff9
[Perf] Specialize Hopper Marlin TP8 single-token MXFP4 gate tile
BBuf Sep 25, 2026
c810bea
[Perf] Overlap Hopper decode mHC statistics without changing arithmetic
BBuf Sep 25, 2026
d7b4b66
[Perf] Defer Hopper MXFP4 expert padding until Marlin repacking
BBuf Sep 25, 2026
361f10d
[Test] Supply quantization metadata in MXFP4 shard fixture
BBuf Sep 25, 2026
60a9fa1
[Test] Model Marlin GEMM2 rounding before route weighting
BBuf Sep 25, 2026
114ab11
Adapt Hopper paths to the kernel split and simplify test helpers
BBuf Sep 26, 2026
b96859b
Preserve Hopper paged indexer optimization across the backend split
BBuf Sep 26, 2026
f3476d4
Check candidate mask reuse at the scoring call
BBuf Sep 26, 2026
50e39bd
Remove the Hopper group32 linear test file
BBuf Sep 26, 2026
2cd22c3
Merge latest main and preserve Hopper and gfx950 mHC dispatch
BBuf Oct 2, 2026
79ca023
Clarify Hopper dispatch and indexer cache contracts
BBuf Oct 2, 2026
22e352d
test: adapt MXFP4 TP loading coverage to current MoE loader API
BBuf Oct 2, 2026
94d21e3
test: align sparse prefill mock with SWA page-size API
BBuf Oct 2, 2026
179b1ab
perf: fuse Blackwell flattened prefill index selection
BBuf Oct 2, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 28 additions & 9 deletions python/sglang/kernels/jit/csrc/elementwise/activation.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@ struct ActivationParams {
// for per-token routing and BLOCK_SIZE_M for sorted/TMA routing.
const int32_t* __restrict__ expert_ids;
uint32_t expert_step;
float clamp_limit;
};

template <
Expand All @@ -61,7 +62,8 @@ template <
bool kUsePDL,
bool kFilterExpert,
bool kRoundActivation = false,
bool kReuseInput = false>
bool kReuseInput = false,
bool kClamp = false>
__global__ void act_and_mul_kernel(const __grid_constant__ ActivationParams params) {
using namespace device;
constexpr auto kVecSize = kMaxVecBytes / sizeof(T);
Expand All @@ -80,11 +82,17 @@ __global__ void act_and_mul_kernel(const __grid_constant__ ActivationParams para
PDLWaitPrimary<kUsePDL>();
const auto gate = device::load_as<vec_t>(params.input, input_offset);
const auto up = device::load_as<vec_t>(params.input, input_offset + num_vecs);
const float limit = device::cast<fp32_t>(device::cast<T>(params.clamp_limit));
vec_t out;
#pragma unroll
for (int i = 0; i < kVecSize; ++i) {
const float gate_f32 = device::cast<fp32_t>(gate[i]);
const float up_f32 = device::cast<fp32_t>(up[i]);
float gate_f32 = device::cast<fp32_t>(gate[i]);
float up_f32 = device::cast<fp32_t>(up[i]);
if constexpr (kClamp) {
static_assert(kAct == ActivationKind::kSiLU);
gate_f32 = gate_f32 > limit ? limit : gate_f32;
up_f32 = up_f32 > limit ? limit : (up_f32 < -limit ? -limit : up_f32);
}
if constexpr (kRoundActivation) {
const T activated = device::cast<T>(apply_activation_f32<kAct>(gate_f32));
out[i] = device::cast<T>(device::cast<fp32_t>(activated) * up_f32);
Expand Down Expand Up @@ -153,13 +161,14 @@ struct ActivationKernel {
return nullptr;
}

template <bool kRoundActivation = false, bool kReuseInput = false>
template <bool kRoundActivation = false, bool kReuseInput = false, bool kClamp = false>
static void launch(
const tvm::ffi::TensorView& input,
const tvm::ffi::TensorView& out,
const std::string& type,
const int32_t* expert_ids,
uint32_t expert_step) {
uint32_t expert_step,
float clamp_limit = 0.0f) {
using namespace host;

auto N = SymbolicSize{"num_tokens"};
Expand Down Expand Up @@ -198,8 +207,14 @@ struct ActivationKernel {
.num_tokens = num_tokens,
.expert_ids = expert_ids,
.expert_step = expert_step,
.clamp_limit = clamp_limit,
};
if (expert_ids != nullptr) {
if constexpr (kClamp) {
RuntimeCheck(type == "silu" && expert_ids == nullptr, "clamping requires unfiltered SiLU");
const auto kernel =
act_and_mul_kernel<T, ActivationKind::kSiLU, kUsePDL, false, kRoundActivation, kReuseInput, true>;
LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(kernel, params);
} else if (expert_ids != nullptr) {
RuntimeCheck(expert_step > 0, "expert_step must be positive");
const auto kernel = select_kernel<true, kRoundActivation, kReuseInput>(type);
LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(kernel, params);
Expand All @@ -213,9 +228,13 @@ struct ActivationKernel {
launch(input, out, type, /*expert_ids=*/nullptr, /*expert_step=*/1);
}

static void
run_activation_with_rounding(const tvm::ffi::TensorView input, const tvm::ffi::TensorView out, std::string type) {
launch<true>(input, out, type, /*expert_ids=*/nullptr, /*expert_step=*/1);
static void run_activation_with_rounding(
const tvm::ffi::TensorView input, const tvm::ffi::TensorView out, std::string type, double clamp_limit) {
if (clamp_limit > 0) {
launch<true, false, true>(input, out, type, /*expert_ids=*/nullptr, /*expert_step=*/1, clamp_limit);
} else {
launch<true>(input, out, type, /*expert_ids=*/nullptr, /*expert_step=*/1);
}
}

static void run_activation_with_rounding_input_inplace(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -404,6 +404,7 @@ bool is_valid_config(
#define MXFP4_GET_IF(W_TYPE) \
MXFP4_GET_IF_M1(W_TYPE, 8, 8, 256) \
MXFP4_GET_IF_M1(W_TYPE, 8, 4, 128) \
MXFP4_GET_IF_M1(W_TYPE, 4, 8, 128) \
MXFP4_GET_IF_M234(W_TYPE, 16, 4, 256) \
MXFP4_GET_IF_M234(W_TYPE, 8, 4, 128)

Expand Down Expand Up @@ -492,6 +493,16 @@ exec_config_t determine_exec_config(
bool is_zp_float,
int max_shared_mem,
int sms) {
#if SGL_CUDA_ARCH == 900
if constexpr (std::is_same_v<scalar_t, __nv_bfloat16> && !kIsEP && !kHasBias) {
// H200 TP8 DeepSeek-V4.1 gate/up: six sparse expert rows need a narrower
// N tile and fewer persistent blocks than the occupancy-only heuristic.
if (q_type == host::kFE2M1f && group_size == 32 && prob_m == 1 && prob_n == 640 && prob_k == 5120 && top_k == 6 &&
thread_m_blocks == 1 && m_block_size_8 && !has_act_order && !has_zp) {
return exec_config_t{2, thread_config_t{128, 64, 128}};
}
}
#endif
exec_config_t exec_cfg = exec_config_t{1, thread_config_t{-1, -1, -1}};
thread_config_t* thread_configs = thread_m_blocks > 1 ? large_batch_thread_configs : small_batch_thread_configs;
int thread_configs_size = thread_m_blocks > 1 ? sizeof(large_batch_thread_configs) / sizeof(thread_config_t)
Expand Down
8 changes: 5 additions & 3 deletions python/sglang/kernels/ops/activation/activation.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,14 +67,14 @@ def _run_activation_inplace(

@register_custom_op(mutates_args=["out"])
def _run_activation_with_rounding_inplace(
op_name: str, input: torch.Tensor, out: torch.Tensor
op_name: str, input: torch.Tensor, out: torch.Tensor, clamp_limit: float = 0.0
) -> None:
hidden_size = input.shape[-1] // 2
# Fast-math changes FP16 SiLU at eager rounding boundaries on SM90.
module = activation_module(input.dtype, fast_math=False)
input_2d = input.view(-1, hidden_size * 2)
out_2d = out.view(-1, hidden_size)
module.run_activation_with_rounding(input_2d, out_2d, op_name)
module.run_activation_with_rounding(input_2d, out_2d, op_name, clamp_limit)


@register_custom_op(mutates_args=["input"])
Expand Down Expand Up @@ -184,11 +184,13 @@ def silu_and_mul(
def silu_and_mul_with_activation_rounding(
input: torch.Tensor,
out: Optional[torch.Tensor] = None,
*,
clamp_limit: float = 0.0,
) -> torch.Tensor:
hidden_size = input.shape[-1] // 2
if out is None:
out = input.new_empty(*input.shape[:-1], hidden_size)
_run_activation_with_rounding_inplace("silu", input, out)
_run_activation_with_rounding_inplace("silu", input, out, clamp_limit)
return out


Expand Down
11 changes: 11 additions & 0 deletions python/sglang/kernels/ops/attention/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -359,3 +359,14 @@ def qwen38_qsa_sm121_varlen(
capabilities=frozenset({CapabilityRequirement.CUDA}),
)
)

for _fn in ("fp4_index_logits_paged", "finish_paged_indexer_topk"):
register_kernel(
KernelSpec(
op=f"attention.{_fn}",
backend=KernelBackend.TRITON,
target=f"sglang.kernels.ops.attention.dsv4.fp4_indexer:{_fn}",
capabilities=frozenset({CapabilityRequirement.CUDA}),
)
)
del _fn
12 changes: 12 additions & 0 deletions python/sglang/kernels/ops/attention/dsv4/candidate_blocks.py
Original file line number Diff line number Diff line change
Expand Up @@ -143,6 +143,18 @@ def amax_topk_blocks(
return blocks


def candidate_block_mask(
blocks: torch.Tensor, width: int, block_size: int
) -> torch.Tensor:
"""Materialize a source's block IDs once for Hopper paged-score consumers."""
num_blocks = (width + block_size - 1) // block_size
keep = torch.zeros(
(*blocks.shape[:-1], num_blocks + 1), dtype=torch.bool, device=blocks.device
)
keep.scatter_(-1, blocks.to(torch.int64).masked_fill(blocks < 0, num_blocks), True)
return keep[..., :num_blocks].repeat_interleave(block_size, dim=-1)[..., :width]


def select_candidate_block_ids(
logits: torch.Tensor,
compress_lens: Union[torch.Tensor, int],
Expand Down
Loading
Loading