[Refactor] Deduplicate kernel helpers and remove unused code (#40197)

This commit is contained in:
Xiaoyu Zhang
2026-09-19 09:27:30 +08:00
committed by GitHub
parent 10b0bcfd18
commit 986959e3c4
28 changed files with 191 additions and 1093 deletions
@@ -21,6 +21,7 @@
#include <sgl_kernel/type.cuh> // For bf16_t, fp32_t, device::cast
#include <sgl_kernel/utils.cuh> // For LaunchKernel
#include <sgl_kernel/warp.cuh>
#include <tvm/ffi/container/tensor.h>
@@ -48,19 +49,6 @@ struct KdaPackedDecodeParams {
int32_t use_lower_bound;
};
__device__ __forceinline__ float warp_allreduce_sum(float v) {
#if defined(__HIP_PLATFORM_AMD__)
constexpr uint64_t kFullMask = 0xffffffffffffffffull;
#else
constexpr uint32_t kFullMask = 0xffffffffu;
#endif
#pragma unroll
for (int off = 16; off > 0; off >>= 1) {
v += __shfl_xor_sync(kFullMask, v, off);
}
return v;
}
// K = V = 128 specialization: one lane owns 4 consecutive K-elements (16B).
template <int kWarps, bool kUsePDL>
__global__
@@ -104,8 +92,8 @@ __launch_bounds__(kWarps * 32) void kda_packed_decode_kernel(const KdaPackedDeco
k_sq += k[e] * k[e];
}
// tl: q / sqrt(sum(q*q) + 1e-6), then * scale
const float q_inv = 1.0f / sqrtf(warp_allreduce_sum(q_sq) + 1e-6f);
const float k_inv = 1.0f / sqrtf(warp_allreduce_sum(k_sq) + 1e-6f);
const float q_inv = 1.0f / sqrtf(warp::reduce_sum<32>(q_sq) + 1e-6f);
const float k_inv = 1.0f / sqrtf(warp::reduce_sum<32>(k_sq) + 1e-6f);
#pragma unroll
for (int e = 0; e < kElems; ++e) {
q[e] = q[e] * q_inv * params.scale;
@@ -143,7 +131,7 @@ __launch_bounds__(kWarps * 32) void kda_packed_decode_kernel(const KdaPackedDeco
h[e] *= decay[e];
t += h[e] * k[e];
}
t = warp_allreduce_sum(t);
t = warp::reduce_sum<32>(t);
const float v_new = (cast<fp32_t>(v_ptr[r]) - t) * beta;
float o_acc = 0.0f;
#pragma unroll
@@ -151,7 +139,7 @@ __launch_bounds__(kWarps * 32) void kda_packed_decode_kernel(const KdaPackedDeco
h[e] += v_new * k[e];
o_acc += h[e] * q[e];
}
o_acc = warp_allreduce_sum(o_acc);
o_acc = warp::reduce_sum<32>(o_acc);
*reinterpret_cast<float4*>(h_base + r * K + e0) = make_float4(h[0], h[1], h[2], h[3]);
if (lane == 0) {
o_ptr[r] = cast<bf16_t>(o_acc);
@@ -161,7 +161,7 @@ __global__ __launch_bounds__(1024, 1) void inkling_ar_sconv_norm_kernel(const __
// ---- 1. push: wait for the producer's output (PDL; no-op without a PDL
// launch or an early-triggering producer), multicast-store this rank's
// partial row, and issue the residual load (it lands under the barrier). ----
asm volatile("griddepcontrol.wait;" ::: "memory");
device::PDLWaitPrimary<true>();
const auto* in_row = static_cast<const __nv_bfloat16*>(p.in) + t * p.in_stride_t;
const auto* sh_row =
p.shared == nullptr ? nullptr : static_cast<const __nv_bfloat16*>(p.shared) + t * p.shared_stride_t;
@@ -393,7 +393,7 @@ __launch_bounds__(1024, 1) void inkling_ar_sconv_norm_verify_kernel(const __grid
// GPU, so Phase 2's cross-token (neighbor) staging reads are race-free -- the
// per-block barrier only synchronized the same blockIdx across ranks and did
// NOT order block t-j's push before block t's read.
asm volatile("griddepcontrol.wait;" ::: "memory");
device::PDLWaitPrimary<true>();
auto* mc = static_cast<__nv_bfloat16*>(p.mc_stage);
const auto* in = static_cast<const __nv_bfloat16*>(p.in);
const auto* sh = static_cast<const __nv_bfloat16*>(p.shared);
@@ -1028,7 +1028,7 @@ __global__ __launch_bounds__(1024, 1) void inkling_ar_col_decode_kernel(const __
static_cast<const __nv_bfloat16*>(p.residual_in) + static_cast<int64_t>(t) * p.H + c0[i]);
}
}
asm volatile("griddepcontrol.wait;" ::: "memory");
device::PDLWaitPrimary<true>();
// ---- 1. entry: peers' producer partials visible (block t <-> peer t) ----
inkling_ar::block_system_barrier<kNumGPU>(p.state, p.flag_ptrs, p.rank);
@@ -1764,7 +1764,7 @@ __launch_bounds__(1024, 1) void inkling_ar_ssconv_norm_decode_kernel(const __gri
}
// ---- 1. push: this rank's partial row + its window shard ----
asm volatile("griddepcontrol.wait;" ::: "memory");
device::PDLWaitPrimary<true>();
const auto* in_row = static_cast<const __nv_bfloat16*>(p.in) + t * p.in_stride_t;
auto* slot = static_cast<__nv_bfloat16*>(p.mc_stage) + (static_cast<uint64_t>(p.rank) * p.T + t) * p.D;
auto* wslot = static_cast<__nv_bfloat16*>(p.mc_wstage) + static_cast<uint64_t>(t) * W1 * p.D;
@@ -13,10 +13,12 @@ See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#include <sgl_kernel/bits.h>
#include <sgl_kernel/tensor.h>
#include <sgl_kernel/utils.h>
#include <sgl_kernel/utils.cuh>
#include <sgl_kernel/warp.cuh>
#include <tvm/ffi/container/tensor.h>
@@ -33,28 +35,8 @@ namespace sglang {
using Vec = int4;
inline uint32_t next_pow2(uint32_t x) noexcept {
--x;
x |= x >> 1;
x |= x >> 2;
x |= x >> 4;
x |= x >> 8;
x |= x >> 16;
return x + 1;
}
namespace moe {
__device__ __forceinline__ int warp_exclusive_scan(int v, unsigned mask = 0xffffffffu) {
int original = v;
#pragma unroll
for (int offset = 1; offset < WARP_SIZE; offset <<= 1) {
int n = __shfl_up_sync(mask, v, offset);
if ((threadIdx.x & (WARP_SIZE - 1)) >= offset) v += n;
}
return v - original;
}
template <typename scalar_t>
__global__ void count_and_sort_expert_tokens_kernel(
const scalar_t* __restrict__ topk_ids,
@@ -187,14 +169,14 @@ __global__ void moe_align_block_size_kernel(
const int warp_id = tid / WARP_SIZE;
const int lane_id = tid & (WARP_SIZE - 1);
const int num_warps_for_scan = (scan_size + WARP_SIZE - 1) / WARP_SIZE;
const int warp_sum = warp_exclusive_scan(padded_count) + padded_count;
const int warp_sum = device::warp::inclusive_sum<32>(padded_count);
if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = warp_sum;
__syncthreads();
// warp0 accumulate all the block's prefix sum
if (tid < WARP_SIZE) {
int val = (tid < num_warps_for_scan) ? warp_sums[tid] : 0;
int incl = warp_exclusive_scan(val) + val;
int incl = device::warp::inclusive_sum<32>(val);
warp_sums[tid] = incl;
}
__syncthreads();
@@ -213,13 +195,13 @@ __global__ void moe_align_block_size_kernel(
// Perform 2 level exclusive-prefix-sum to scan_buf
int v = (tid < scan_size) ? scan_buf[tid] : 0;
int pre = warp_exclusive_scan(v);
int pre = device::warp::inclusive_sum<32>(v) - v;
if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = pre + v;
__syncthreads();
if (warp_id == 0) {
int val = (lane_id < num_warps_for_scan) ? warp_sums[lane_id] : 0;
warp_sums[lane_id] = warp_exclusive_scan(val);
warp_sums[lane_id] = device::warp::inclusive_sum<32>(val) - val;
}
__syncthreads();
@@ -409,7 +391,7 @@ __global__ void moe_align_block_size_kernel_v2(
}
// Level 1: intra-warp exclusive scan on thread_sum
int32_t warp_prefix = warp_exclusive_scan(thread_sum);
int32_t warp_prefix = device::warp::inclusive_sum<32>(thread_sum) - thread_sum;
int32_t warp_total = warp_prefix + thread_sum;
if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = warp_total;
__syncthreads();
@@ -418,7 +400,7 @@ __global__ void moe_align_block_size_kernel_v2(
const int num_warps = (blockDim.x + WARP_SIZE - 1) / WARP_SIZE;
if (tid < WARP_SIZE) {
int val = (tid < num_warps) ? warp_sums[tid] : 0;
warp_sums[tid] = warp_exclusive_scan(val);
warp_sums[tid] = device::warp::inclusive_sum<32>(val) - val;
}
__syncthreads();
@@ -510,7 +492,7 @@ struct MoeAlignBlockSizeKernel {
pad_sorted_token_ids,
(int32_t)max_num_tokens_padded);
} else if (num_experts <= 1024) {
const size_t scan_size = next_pow2(num_experts);
const size_t scan_size = host::round_up_pow2(static_cast<uint32_t>(num_experts));
const size_t shared_mem_size = (num_experts + (num_experts + 1) + scan_size + WARP_SIZE) * sizeof(int32_t);
auto align_kernel = moe::moe_align_block_size_kernel<scalar_t>;
@@ -24,10 +24,12 @@ limitations under the License.
// using fused scatter for eligible shapes and two kernels otherwise. Larger
// domains keep the old path through the Python dispatcher.
#include <sgl_kernel/bits.h>
#include <sgl_kernel/tensor.h>
#include <sgl_kernel/utils.h>
#include <sgl_kernel/utils.cuh>
#include <sgl_kernel/warp.cuh>
#include <tvm/ffi/container/tensor.h>
@@ -44,28 +46,8 @@ namespace sglang {
using Vec = int4;
inline uint32_t next_pow2(uint32_t x) noexcept {
--x;
x |= x >> 1;
x |= x >> 2;
x |= x >> 4;
x |= x >> 8;
x |= x >> 16;
return x + 1;
}
namespace moe_lora_merged {
__device__ __forceinline__ int warp_exclusive_scan(int v, unsigned mask = 0xffffffffu) {
int original = v;
#pragma unroll
for (int offset = 1; offset < WARP_SIZE; offset <<= 1) {
int n = __shfl_up_sync(mask, v, offset);
if ((threadIdx.x & (WARP_SIZE - 1)) >= offset) v += n;
}
return v - original;
}
// Inline mirror of _fused_virtual_topk_ids_kernel (virtual_experts.py). Returns
// the merged virtual expert id for flat slot `i` (range [-1, virtual_num_experts);
// -1 is the dropped/masked sentinel). The caller adds +1 to get the histogram
@@ -232,14 +214,14 @@ __global__ void moe_align_block_size_kernel(
const int warp_id = tid / WARP_SIZE;
const int lane_id = tid & (WARP_SIZE - 1);
const int num_warps_for_scan = (scan_size + WARP_SIZE - 1) / WARP_SIZE;
const int warp_sum = warp_exclusive_scan(padded_count) + padded_count;
const int warp_sum = device::warp::inclusive_sum<32>(padded_count);
if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = warp_sum;
__syncthreads();
// warp0 accumulate all the block's prefix sum
if (tid < WARP_SIZE) {
int val = (tid < num_warps_for_scan) ? warp_sums[tid] : 0;
int incl = warp_exclusive_scan(val) + val;
int incl = device::warp::inclusive_sum<32>(val);
warp_sums[tid] = incl;
}
__syncthreads();
@@ -258,13 +240,13 @@ __global__ void moe_align_block_size_kernel(
// Perform 2 level exclusive-prefix-sum to scan_buf
int v = (tid < scan_size) ? scan_buf[tid] : 0;
int pre = warp_exclusive_scan(v);
int pre = device::warp::inclusive_sum<32>(v) - v;
if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = pre + v;
__syncthreads();
if (warp_id == 0) {
int val = (lane_id < num_warps_for_scan) ? warp_sums[lane_id] : 0;
warp_sums[lane_id] = warp_exclusive_scan(val);
warp_sums[lane_id] = device::warp::inclusive_sum<32>(val) - val;
}
__syncthreads();
@@ -384,12 +366,12 @@ __global__ void fused_align_scatter_kernel(
padded_count = (count + block_size - 1) / block_size * block_size;
scan_buf[tid] = padded_count;
}
const int warp_sum = warp_exclusive_scan(padded_count) + padded_count;
const int warp_sum = device::warp::inclusive_sum<32>(padded_count);
if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = warp_sum;
__syncthreads();
if (tid < WARP_SIZE) {
int val = (tid < num_warps_for_scan) ? warp_sums[tid] : 0;
int incl = warp_exclusive_scan(val) + val;
int incl = device::warp::inclusive_sum<32>(val);
warp_sums[tid] = incl;
}
__syncthreads();
@@ -402,12 +384,12 @@ __global__ void fused_align_scatter_kernel(
if (tid >= num_experts && tid < scan_size) scan_buf[tid] = 0;
__syncthreads();
int v = (tid < scan_size) ? scan_buf[tid] : 0;
int pre = warp_exclusive_scan(v);
int pre = device::warp::inclusive_sum<32>(v) - v;
if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = pre + v;
__syncthreads();
if (warp_id == 0) {
int val = (lane_id < num_warps_for_scan) ? warp_sums[lane_id] : 0;
warp_sums[lane_id] = warp_exclusive_scan(val);
warp_sums[lane_id] = device::warp::inclusive_sum<32>(val) - val;
}
__syncthreads();
int off = warp_sums[warp_id];
@@ -500,7 +482,7 @@ struct MoeLoraMergedAlignKernel {
int32_t* cumsum_buffer_ptr = static_cast<int32_t*>(cumsum_buffer.data_ptr());
size_t numel = topk_ids.numel();
const size_t scan_size = next_pow2(num_experts);
const size_t scan_size = host::round_up_pow2(static_cast<uint32_t>(num_experts));
if (fuse_scatter) {
// One block does fill + histogram + scan + expert_ids + scatter. Extra
@@ -1,6 +1,7 @@
"""Public interface of sglang.kernels.jit.utils."""
from sglang.kernels.jit.utils.arch import (
get_activation_cuda_cflags,
get_jit_cuda_arch,
is_arch_support_pdl,
override_jit_cuda_arch,
@@ -27,6 +28,7 @@ __all__ = [
"make_cpp_args",
"load_jit",
"override_jit_cuda_arch",
"get_activation_cuda_cflags",
"get_jit_cuda_arch",
"is_arch_support_pdl",
"KERNEL_PATH",
+8
View File
@@ -179,3 +179,11 @@ def is_arch_support_pdl() -> bool:
if is_hip_runtime() or is_musa_runtime():
return False
return get_jit_cuda_arch().major >= 9
def get_activation_cuda_cflags() -> list[str]:
"""Match the AOT activation fast-math policy without changing other kernels."""
# Blackwell needs precise expf; HIP clang rejects --use_fast_math.
if is_hip_runtime() or get_jit_cuda_arch().major >= 10:
return []
return ["--use_fast_math"]
@@ -6,9 +6,8 @@ import torch
from sglang.kernels.jit.utils import (
cache_once,
get_jit_cuda_arch,
get_activation_cuda_cflags,
is_arch_support_pdl,
is_hip_runtime,
load_jit,
make_cpp_args,
)
@@ -18,19 +17,9 @@ if TYPE_CHECKING:
from tvm_ffi.module import Module
def _fast_math_flags() -> list[str]:
# Mirrors sgl-kernel's CMake policy: fast-math on SM90, precise on
# SM100+ (Blackwell needs bit-exact expf), off on HIP (clang rejects).
if is_hip_runtime():
return []
if get_jit_cuda_arch().major >= 10:
return []
return ["--use_fast_math"]
@cache_once
def activation_module(dtype: torch.dtype, *, fast_math: bool = True) -> Module:
fast_math_flags = _fast_math_flags()
fast_math_flags = get_activation_cuda_cflags()
if not fast_math and not fast_math_flags:
return activation_module(dtype)
args = make_cpp_args(dtype, is_arch_support_pdl())
@@ -9,49 +9,6 @@ def dequantize_k_cache(quant_k_cache):
return _dequantize_k_cache_fast_wrapped(quant_k_cache)
def _dequantize_k_cache_ref(
quant_k_cache: torch.Tensor, # (num_blocks, block_size, 1, bytes_per_token)
dv: int = 512,
tile_size: int = 128,
d: int = 576,
) -> torch.Tensor:
"""
De-quantize the k-cache
"""
assert dv % tile_size == 0
original_ndim = quant_k_cache.ndim
if original_ndim == 3:
# set block_size = 1
quant_k_cache = quant_k_cache.unsqueeze(1)
num_tiles = dv // tile_size
num_blocks, block_size, h_k, _ = quant_k_cache.shape
assert h_k == 1
result = torch.empty(
(num_blocks, block_size, d), dtype=torch.bfloat16, device=quant_k_cache.device
)
quant_k_cache = quant_k_cache.view(num_blocks, block_size, -1)
input_nope = quant_k_cache[..., :dv]
input_scale = quant_k_cache[..., dv : dv + num_tiles * 4].view(torch.float32)
input_rope = quant_k_cache[..., dv + num_tiles * 4 :].view(torch.bfloat16)
result[..., dv:] = input_rope
for tile_idx in range(0, num_tiles):
cur_nope = input_nope[
..., tile_idx * tile_size : (tile_idx + 1) * tile_size
].to(torch.float32)
cur_scales = input_scale[..., tile_idx].unsqueeze(-1)
result[..., tile_idx * tile_size : (tile_idx + 1) * tile_size] = (
cur_nope * cur_scales
)
if original_ndim == 3:
return result.view(num_blocks, 1, -1)
else:
return result.view(num_blocks, block_size, 1, -1)
def _dequantize_k_cache_fast_wrapped(
quant_k_cache: torch.Tensor,
dv: int = 512,
@@ -11,17 +11,19 @@ Two variants:
2. Split-K: adaptive split-K with fused fast path (adapted from DSv4)
"""
import functools
import torch
import triton
import triton.language as tl
from sglang.kernels.ops.attention.dsa.triton_sparse_mla import (
_PREFERRED_BLOCK_K,
_cu_count,
_kv_splits_heuristic,
_next_pow2,
_no_async_copy,
_row_strides,
_sparse_mla_block_k,
_sparse_mla_reduce_kernel,
_validate_input_dtypes,
)
from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz
@@ -70,44 +72,6 @@ def _get_splitk_bufs(
LOG2E = 1.4426950408889634
@functools.lru_cache(maxsize=1)
def _cu_count() -> int:
return torch.cuda.get_device_properties(
torch.cuda.current_device()
).multi_processor_count
def _prev_pow2(n: int) -> int:
if n < 1:
return 1
return 1 << (n.bit_length() - 1)
def _next_pow2(n: int) -> int:
if n < 1:
return 1
return 1 << (n - 1).bit_length()
def _kv_splits_heuristic(
T: int,
H: int,
block_h: int,
num_cu: int | None = None,
target_wg_per_cu: float = 2.0,
max_kv_splits: int = 64,
) -> int:
if num_cu is None:
num_cu = _cu_count()
target_wg = max(1, int(target_wg_per_cu * num_cu))
head_blocks = max(1, (H + block_h - 1) // block_h)
base_ctas = max(1, T * head_blocks)
if base_ctas >= target_wg:
return 1
splits_to_fill = max(1, target_wg // base_ctas)
return _prev_pow2(min(splits_to_fill, max_kv_splits))
@triton.jit
def _sparse_mla_decode_fused_kernel(
q_nope_ptr, # [N, H, D_V]
@@ -512,64 +476,6 @@ def _sparse_mla_decode_split_kernel(
)
@triton.jit
def _sparse_mla_decode_reduce_kernel(
lse_partial_ptr, # [N, KV_SPLITS, H_padded] fp32
acc_partial_ptr, # [N, KV_SPLITS, H_padded, D_V] bf16
out_ptr, # [N, H, D_V]
H: tl.constexpr,
D_V: tl.constexpr,
KV_SPLITS: tl.constexpr,
ACTIVE_SPLITS: tl.constexpr,
ACTIVE_SPLITS_POW2: tl.constexpr,
D_CHUNK: tl.constexpr,
BLOCK_K: tl.constexpr,
):
t = tl.program_id(0)
h = tl.program_id(1)
dc = tl.program_id(2)
d_offs = dc * D_CHUNK + tl.arange(0, D_CHUNK)
# tl.arange needs a power-of-two extent, but ACTIVE_SPLITS is only a power
# of two when topk // BLOCK_K is. Iterate over the padded range and mask the
# tail: -3.4e38 drives exp2() to 0 without the NaN an -inf would produce.
k_offs = tl.arange(0, ACTIVE_SPLITS_POW2)
k_mask = k_offs < ACTIVE_SPLITS
d_mask = d_offs < D_V
H_padded = tl.cdiv(H, 16) * 16
lse_base = t * KV_SPLITS * H_padded
lse_p = tl.load(
lse_partial_ptr + lse_base + k_offs * H_padded + h,
mask=k_mask,
other=-3.4e38,
)
ap_base = t * KV_SPLITS * H_padded * D_V
a_p = tl.load(
acc_partial_ptr
+ ap_base
+ k_offs[:, None] * H_padded * D_V
+ h * D_V
+ d_offs[None, :],
mask=k_mask[:, None] & d_mask[None, :],
other=0.0,
).to(tl.float32)
lse_max = tl.max(lse_p, axis=0)
weights = tl.exp2(lse_p - lse_max)
w_sum = tl.sum(weights, axis=0)
scale = tl.exp2(lse_p - lse_max - tl.log2(tl.maximum(w_sum, 1.0e-30)))
out = tl.sum(a_p * scale[:, None], axis=0)
tl.store(
out_ptr + t * H * D_V + h * D_V + d_offs,
out.to(tl.bfloat16),
mask=d_mask,
)
def triton_sparse_mla_decode_splitk(
q_nope: torch.Tensor,
q_rope: torch.Tensor,
@@ -710,7 +616,7 @@ def triton_sparse_mla_decode_splitk(
D_CHUNK = 64
grid_reduce = (bs, H, (d_v + D_CHUNK - 1) // D_CHUNK)
_sparse_mla_decode_reduce_kernel[grid_reduce](
_sparse_mla_reduce_kernel[grid_reduce](
lse_partial,
acc_partial,
out,
@@ -17,6 +17,10 @@ from sglang.kernels.ops.attention.cute_utils import (
fence_before_tma_store,
simple_tma_copy,
)
from sglang.kernels.ops.attention.linear.tma import (
make_chunk_tma_args,
make_recurrent_state_tma_args,
)
class Sm100ChunkHKernel:
@@ -53,46 +57,6 @@ class Sm100ChunkHKernel:
self.num_stages = num_stages
self.num_warps = 10
@cute.jit
def _make_bf16_tma_args(
self,
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
):
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(self.BT, 1, (64, dim // 64), stages),
stride=(64, 0, (1, self.BT * 64), self.BT * dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(self.BT, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def _make_h_tma_args(self, tensor: cute.Tensor, op: cpasync.TmaCopyOp):
# number of elements to fill 128B
num_elems = 128 // (tensor.element_type.width // 8)
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(1, 1, self.V_dim, (num_elems, self.K_dim // num_elems)),
stride=(0, 0, num_elems, (1, self.V_dim * num_elems)),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, None, num_elems)),
slayout,
cta_tiler=(1, 1, self.V_dim, self.K_dim),
)
return atom, tma_tensor, slayout
@cute.jit
def __call__(
self,
@@ -112,13 +76,13 @@ class Sm100ChunkHKernel:
tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
K_args = self._make_bf16_tma_args(K, self.K_dim, tma_g2s, self.num_stages)
V_args = self._make_bf16_tma_args(V, self.V_dim, tma_g2s, self.num_stages)
W_args = self._make_bf16_tma_args(W, self.K_dim, tma_g2s, self.num_stages)
V_new_args = self._make_bf16_tma_args(V_new, self.V_dim, tma_s2g, 1)
H0_args = self._make_h_tma_args(h0, tma_g2s)
HT_args = self._make_h_tma_args(ht, tma_s2g)
H_args = self._make_h_tma_args(h, tma_s2g)
K_args = make_chunk_tma_args(K, self.K_dim, tma_g2s, self.num_stages, self.BT)
V_args = make_chunk_tma_args(V, self.V_dim, tma_g2s, self.num_stages, self.BT)
W_args = make_chunk_tma_args(W, self.K_dim, tma_g2s, self.num_stages, self.BT)
V_new_args = make_chunk_tma_args(V_new, self.V_dim, tma_s2g, 1, self.BT)
H0_args = make_recurrent_state_tma_args(h0, tma_g2s, self.K_dim, self.V_dim)
HT_args = make_recurrent_state_tma_args(ht, tma_s2g, self.K_dim, self.V_dim)
H_args = make_recurrent_state_tma_args(h, tma_s2g, self.K_dim, self.V_dim)
# h0/ht may be the full state pool ([num_slots, ...]) rather than a
# per-sequence gather, so the sequence count comes from cu_seqlens and
@@ -18,6 +18,9 @@ from sglang.kernels.ops.attention.cute_utils import (
mma_bf16,
simple_tma_copy,
)
from sglang.kernels.ops.attention.linear.tma import (
make_chunk_tma_args,
)
class Sm100ChunkUWKernel:
@@ -50,33 +53,6 @@ class Sm100ChunkUWKernel:
self.BT = 64
self.num_warps = 2 + 4 + 4
@cute.jit
def _make_tma_args(
self,
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
num_stages: int,
op: cpasync.TmaCopyOp,
):
# logical layout: [BT, dim]
# permute for TMA: [dim/64, BT, 64] with swizzling
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(self.BT, 1, (64, dim // 64), num_stages),
stride=(64, 0, (1, self.BT * 64), self.BT * dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
# we need to convert gmem layout to (T, H, (64, D/64)) for make_tiled_tma_atom()
# to emit a single 4D TMA. otherwise, it will emit (D/64)x 3D TMA.
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(self.BT, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def __call__(
self,
@@ -96,10 +72,10 @@ class Sm100ChunkUWKernel:
tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
K_args = self._make_tma_args(K, self.K_dim, self.num_stages, tma_g2s)
V_args = self._make_tma_args(V, self.V_dim, self.num_stages, tma_g2s)
U_args = self._make_tma_args(U, self.V_dim, 1, tma_s2g)
W_args = self._make_tma_args(W, self.K_dim, 1, tma_s2g)
K_args = make_chunk_tma_args(K, self.K_dim, tma_g2s, self.num_stages, self.BT)
V_args = make_chunk_tma_args(V, self.V_dim, tma_g2s, self.num_stages, self.BT)
U_args = make_chunk_tma_args(U, self.V_dim, tma_s2g, 1, self.BT)
W_args = make_chunk_tma_args(W, self.K_dim, tma_s2g, 1, self.BT)
grid = (num_sms // self.Hv, self.Hv, 1)
block = (self.num_warps * 32, 1, 1)
@@ -17,6 +17,10 @@ from sglang.kernels.ops.attention.cute_utils import (
fence_before_tma_store,
simple_tma_copy,
)
from sglang.kernels.ops.attention.linear.tma import (
make_chunk_tma_args,
make_output_state_tma_args,
)
class Sm100ChunkOKernel:
@@ -48,50 +52,6 @@ class Sm100ChunkOKernel:
self.num_stages = num_stages
self.num_warps = 10
@cute.jit
def _make_bf16_tma_args(
self,
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
):
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(self.BT, 1, (64, dim // 64), stages),
stride=(64, 0, (1, self.BT * 64), self.BT * dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(self.BT, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def _make_h_tma_args(
self,
tensor: cute.Tensor,
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
):
num_elems = 128 // (tensor.element_type.width // 8)
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(1, self.V_dim, (num_elems, self.K_dim // num_elems), stages),
stride=(0, num_elems, (1, self.V_dim * num_elems), self.V_dim * self.K_dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, num_elems)),
slayout,
cta_tiler=(1, self.V_dim, self.K_dim),
)
return atom, tma_tensor, slayout
@cute.jit
def __call__(
self,
@@ -112,13 +72,15 @@ class Sm100ChunkOKernel:
block = (self.num_warps * 32, 1, 1)
tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
Q_args = self._make_bf16_tma_args(q, self.K_dim, tma_g2s, self.num_stages)
K_args = self._make_bf16_tma_args(k, self.K_dim, tma_g2s, self.num_stages)
V_args = self._make_bf16_tma_args(
v_new_chunks, self.V_dim, tma_g2s, self.num_stages
Q_args = make_chunk_tma_args(q, self.K_dim, tma_g2s, self.num_stages, self.BT)
K_args = make_chunk_tma_args(k, self.K_dim, tma_g2s, self.num_stages, self.BT)
V_args = make_chunk_tma_args(
v_new_chunks, self.V_dim, tma_g2s, self.num_stages, self.BT
)
H_args = self._make_h_tma_args(h, tma_g2s, self.num_stages)
O_args = self._make_bf16_tma_args(o, self.V_dim, tma_s2g, 1)
H_args = make_output_state_tma_args(
h, tma_g2s, self.num_stages, self.K_dim, self.V_dim
)
O_args = make_chunk_tma_args(o, self.V_dim, tma_s2g, 1, self.BT)
self.kernel(
Q_args,
K_args,
@@ -33,6 +33,10 @@ from sglang.kernels.ops.attention.cute_utils import (
fence_before_tma_store,
simple_tma_copy,
)
from sglang.kernels.ops.attention.linear.tma import (
make_chunk_tma_args,
make_recurrent_state_tma_args,
)
class Sm100KdaChunkHKernel:
@@ -60,45 +64,6 @@ class Sm100KdaChunkHKernel:
self.num_stages = num_stages
self.num_warps = 10
@cute.jit
def _make_bf16_tma_args(
self,
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
):
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(self.BT, 1, (64, dim // 64), stages),
stride=(64, 0, (1, self.BT * 64), self.BT * dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(self.BT, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def _make_h_tma_args(self, tensor: cute.Tensor, op: cpasync.TmaCopyOp):
num_elems = 128 // (tensor.element_type.width // 8)
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(1, 1, self.V_dim, (num_elems, self.K_dim // num_elems)),
stride=(0, 0, num_elems, (1, self.V_dim * num_elems)),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, None, num_elems)),
slayout,
cta_tiler=(1, 1, self.V_dim, self.K_dim),
)
return atom, tma_tensor, slayout
@cute.jit
def __call__(
self,
@@ -118,13 +83,13 @@ class Sm100KdaChunkHKernel:
tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
K_args = self._make_bf16_tma_args(K, self.K_dim, tma_g2s, self.num_stages)
V_args = self._make_bf16_tma_args(V, self.V_dim, tma_g2s, self.num_stages)
W_args = self._make_bf16_tma_args(W, self.K_dim, tma_g2s, self.num_stages)
V_new_args = self._make_bf16_tma_args(V_new, self.V_dim, tma_s2g, 1)
H0_args = self._make_h_tma_args(h0, tma_g2s)
HT_args = self._make_h_tma_args(ht, tma_s2g)
H_args = self._make_h_tma_args(h, tma_s2g)
K_args = make_chunk_tma_args(K, self.K_dim, tma_g2s, self.num_stages, self.BT)
V_args = make_chunk_tma_args(V, self.V_dim, tma_g2s, self.num_stages, self.BT)
W_args = make_chunk_tma_args(W, self.K_dim, tma_g2s, self.num_stages, self.BT)
V_new_args = make_chunk_tma_args(V_new, self.V_dim, tma_s2g, 1, self.BT)
H0_args = make_recurrent_state_tma_args(h0, tma_g2s, self.K_dim, self.V_dim)
HT_args = make_recurrent_state_tma_args(ht, tma_s2g, self.K_dim, self.V_dim)
H_args = make_recurrent_state_tma_args(h, tma_s2g, self.K_dim, self.V_dim)
# h0/ht may be the full state pool ([num_slots, ...]) rather than a
# per-sequence gather, so the sequence count comes from cu_seqlens and
@@ -33,6 +33,9 @@ from sglang.kernels.ops.attention.cute_utils import (
mma_bf16,
simple_tma_copy,
)
from sglang.kernels.ops.attention.linear.tma import (
make_chunk_tma_args,
)
class Sm100KdaChunkUWKernel:
@@ -57,28 +60,6 @@ class Sm100KdaChunkUWKernel:
self.BT = 64
self.num_warps = 2 + 4 + 4
@cute.jit
def _make_tma_args(
self,
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
num_stages: int,
op: cpasync.TmaCopyOp,
):
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(self.BT, 1, (64, dim // 64), num_stages),
stride=(64, 0, (1, self.BT * 64), self.BT * dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(self.BT, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def __call__(
self,
@@ -98,12 +79,12 @@ class Sm100KdaChunkUWKernel:
tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
KL_args = self._make_tma_args(KL, self.K_dim, self.num_stages, tma_g2s)
KR_args = self._make_tma_args(KR, self.K_dim, self.num_stages, tma_g2s)
KG_args = self._make_tma_args(KG, self.K_dim, self.num_stages, tma_g2s)
V_args = self._make_tma_args(V, self.V_dim, self.num_stages, tma_g2s)
U_args = self._make_tma_args(U, self.V_dim, 1, tma_s2g)
W_args = self._make_tma_args(W, self.K_dim, 1, tma_s2g)
KL_args = make_chunk_tma_args(KL, self.K_dim, tma_g2s, self.num_stages, self.BT)
KR_args = make_chunk_tma_args(KR, self.K_dim, tma_g2s, self.num_stages, self.BT)
KG_args = make_chunk_tma_args(KG, self.K_dim, tma_g2s, self.num_stages, self.BT)
V_args = make_chunk_tma_args(V, self.V_dim, tma_g2s, self.num_stages, self.BT)
U_args = make_chunk_tma_args(U, self.V_dim, tma_s2g, 1, self.BT)
W_args = make_chunk_tma_args(W, self.K_dim, tma_s2g, 1, self.BT)
grid = (num_sms // self.Hv, self.Hv, 1)
block = (self.num_warps * 32, 1, 1)
@@ -32,6 +32,10 @@ from sglang.kernels.ops.attention.cute_utils import (
fence_before_tma_store,
simple_tma_copy,
)
from sglang.kernels.ops.attention.linear.tma import (
make_chunk_tma_args,
make_output_state_tma_args,
)
class Sm100KdaChunkOKernel:
@@ -58,50 +62,6 @@ class Sm100KdaChunkOKernel:
self.num_stages = num_stages
self.num_warps = 10
@cute.jit
def _make_bf16_tma_args(
self,
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
):
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(self.BT, 1, (64, dim // 64), stages),
stride=(64, 0, (1, self.BT * 64), self.BT * dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(self.BT, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def _make_h_tma_args(
self,
tensor: cute.Tensor,
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
):
num_elems = 128 // (tensor.element_type.width // 8)
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(1, self.V_dim, (num_elems, self.K_dim // num_elems), stages),
stride=(0, num_elems, (1, self.V_dim * num_elems), self.V_dim * self.K_dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, num_elems)),
slayout,
cta_tiler=(1, self.V_dim, self.K_dim),
)
return atom, tma_tensor, slayout
@cute.jit
def __call__(
self,
@@ -121,14 +81,16 @@ class Sm100KdaChunkOKernel:
block = (self.num_warps * 32, 1, 1)
tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
Q_args = self._make_bf16_tma_args(qg2, self.K_dim, tma_g2s, self.num_stages)
Q2_args = self._make_bf16_tma_args(qg, self.K_dim, tma_g2s, self.num_stages)
K_args = self._make_bf16_tma_args(kg, self.K_dim, tma_g2s, self.num_stages)
V_args = self._make_bf16_tma_args(
v_new_chunks, self.V_dim, tma_g2s, self.num_stages
Q_args = make_chunk_tma_args(qg2, self.K_dim, tma_g2s, self.num_stages, self.BT)
Q2_args = make_chunk_tma_args(qg, self.K_dim, tma_g2s, self.num_stages, self.BT)
K_args = make_chunk_tma_args(kg, self.K_dim, tma_g2s, self.num_stages, self.BT)
V_args = make_chunk_tma_args(
v_new_chunks, self.V_dim, tma_g2s, self.num_stages, self.BT
)
H_args = self._make_h_tma_args(h, tma_g2s, self.num_stages)
O_args = self._make_bf16_tma_args(o, self.V_dim, tma_s2g, 1)
H_args = make_output_state_tma_args(
h, tma_g2s, self.num_stages, self.K_dim, self.V_dim
)
O_args = make_chunk_tma_args(o, self.V_dim, tma_s2g, 1, self.BT)
self.kernel(
Q_args,
Q2_args,
@@ -341,111 +341,6 @@ def _get_padded_input_buffers(
return e
# Multi-seq varlen repack cache for the Phase 2.2 path. Keyed by id(cu_seqlens).
# Stores: (orig_seq_lens, padded_seq_lens, new_cu_seqlens_tensor,
# new_chunk_indices_tensor, new_T_total, padded_input_buffers).
# All tensors are GPU-side and pre-allocated at cache build time. Per-call the
# kernel reads from / writes to these buffers; we copy caller's input slices in
# and output slices back (only the valid prefix of each seq).
_multiseq_repack_cache = {}
_caller_layout_O_cache = {}
def _get_caller_layout_O_buffer(multiseq_info, dtype, V_dim, device):
"""Per-shape cached output buffer at caller's contiguous layout (sum of
seq_lens, no padding gaps). Filled by per-seq copies from K4's padded O."""
caller_T = multiseq_info["caller_T"]
H_x_V = multiseq_info.get("_H_V") # not strictly needed since we fix B=1 H known
key = (caller_T, V_dim, dtype, device.index if device.index is not None else 0)
e = _caller_layout_O_cache.get(key)
if e is None:
# B=1 enforced upstream when multiseq_info is built.
e = torch.empty(
1,
caller_T,
multiseq_info["q_pad"].shape[2],
V_dim,
dtype=dtype,
device=device,
)
_caller_layout_O_cache[key] = e
return e
def _get_multiseq_repack_info(cu_seqlens, q, k, v, g, beta, BT, device):
"""Build (and cache) the padded layout for multi-seq varlen with non-aligned
seqs. Returns None if all seqs are already 64-aligned (caller can use the
existing varlen_pure path)."""
import weakref
key = id(cu_seqlens)
cached = _multiseq_repack_cache.get(key)
if cached is not None:
wref, e = cached
if wref() is cu_seqlens:
return e
# id collision after GC: rebuild
del _multiseq_repack_cache[key]
cu_cpu = cu_seqlens.cpu().tolist()
seq_lens = [cu_cpu[i + 1] - cu_cpu[i] for i in range(len(cu_cpu) - 1)]
if all(sl % BT == 0 for sl in seq_lens):
_multiseq_repack_cache[key] = (weakref.ref(cu_seqlens), None)
return None
padded_lens = [((sl + BT - 1) // BT) * BT for sl in seq_lens]
new_cu = [0]
for pl in padded_lens:
new_cu.append(new_cu[-1] + pl)
new_T_total = new_cu[-1]
B = q.shape[0]
H = q.shape[2]
K = q.shape[3]
# Pre-allocated padded input buffers. q/k/v/beta tail = 0 (zero MMA), g
# tail = -1e3 sentinel (zero gate activation). The PER-SEQ tail regions
# are between (new_cu[i] + seq_lens[i], new_cu[i+1]) — pre-fill once.
q_pad = torch.zeros(B, new_T_total, H, K, dtype=q.dtype, device=device)
k_pad = torch.zeros_like(q_pad)
v_pad = torch.zeros(B, new_T_total, H, v.shape[3], dtype=v.dtype, device=device)
beta_pad = torch.zeros(B, new_T_total, H, dtype=beta.dtype, device=device)
g_pad = torch.zeros(B, new_T_total, H, K, dtype=g.dtype, device=device)
for i, (sl, pl) in enumerate(zip(seq_lens, padded_lens)):
if sl < pl:
tail_start = new_cu[i] + sl
tail_end = new_cu[i + 1]
g_pad[:, tail_start:tail_end] = -1000.0
new_cu_tensor = torch.tensor(new_cu, dtype=cu_seqlens.dtype, device=device)
new_chunk_indices = prepare_chunk_indices(new_cu_tensor, BT)
# Build index map: dst_indices[i] = position in padded layout where orig
# row i lives. Used by index_copy_ to do the scatter in one op (instead of
# N_seqs × 5 separate slice copies, which cost ~5us each in Python).
T_total_orig = cu_cpu[-1]
dst_indices_list = []
for i, sl in enumerate(seq_lens):
for j in range(sl):
dst_indices_list.append(new_cu[i] + j)
dst_indices = torch.tensor(dst_indices_list, dtype=torch.long, device=device)
# Mark this cu_seqlens as VARLEN_PURE eligible — every seq in the new
# layout is 64-aligned by construction.
_varlen_pure_cache[id(new_cu_tensor)] = True
e = {
"seq_lens": seq_lens,
"padded_lens": padded_lens,
"new_cu": new_cu,
"new_T_total": new_T_total,
"new_cu_tensor": new_cu_tensor,
"new_chunk_indices": new_chunk_indices,
"q_pad": q_pad,
"k_pad": k_pad,
"v_pad": v_pad,
"g_pad": g_pad,
"beta_pad": beta_pad,
"dst_indices": dst_indices,
"T_total_orig": T_total_orig,
}
_multiseq_repack_cache[key] = (weakref.ref(cu_seqlens), e)
return e
def _get_buffers(dev, dtype_k, B, T, H, K_dim, V_dim, NT, N_seqs, BT):
"""All beta fusion lives in akk_inv kernel epilogue (post-inv column-scale)."""
key = (dev.index or 0, B, T, H, K_dim, V_dim, NT, N_seqs)
@@ -0,0 +1,75 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Shared TMA layouts for the Blackwell GDN and KDA chunk kernels."""
import cutlass
from cutlass import cute
from cutlass.cute.nvgpu import cpasync
@cute.jit
def make_chunk_tma_args(
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
chunk_size: cutlass.Constexpr[int],
):
# Divide the contiguous dimension so the descriptor issues one 4D TMA.
slayout = cute.make_layout(
(chunk_size, 1, (64, dim // 64), stages),
stride=(64, 0, (1, chunk_size * 64), chunk_size * dim),
)
slayout = cute.make_composed_layout(cute.make_swizzle(3, 4, 3), 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(chunk_size, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def make_recurrent_state_tma_args(
tensor: cute.Tensor,
op: cpasync.TmaCopyOp,
key_dim: cutlass.Constexpr[int],
value_dim: cutlass.Constexpr[int],
):
num_elems = 128 // (tensor.element_type.width // 8)
slayout = cute.make_layout(
(1, 1, value_dim, (num_elems, key_dim // num_elems)),
stride=(0, 0, num_elems, (1, value_dim * num_elems)),
)
slayout = cute.make_composed_layout(cute.make_swizzle(3, 4, 3), 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, None, num_elems)),
slayout,
cta_tiler=(1, 1, value_dim, key_dim),
)
return atom, tma_tensor, slayout
@cute.jit
def make_output_state_tma_args(
tensor: cute.Tensor,
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
key_dim: cutlass.Constexpr[int],
value_dim: cutlass.Constexpr[int],
):
num_elems = 128 // (tensor.element_type.width // 8)
slayout = cute.make_layout(
(1, value_dim, (num_elems, key_dim // num_elems), stages),
stride=(0, num_elems, (1, value_dim * num_elems), value_dim * key_dim),
)
slayout = cute.make_composed_layout(cute.make_swizzle(3, 4, 3), 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, num_elems)),
slayout,
cta_tiler=(1, value_dim, key_dim),
)
return atom, tma_tensor, slayout
@@ -1,7 +1,7 @@
import triton
import triton.language as tl
from sglang.kernels.jit.utils import is_arch_support_pdl
from sglang.kernels.ops.gemm.kernel_utils import (
_resolve_token_positions as _resolve_token_positions,
)
def get_pdl_launch_metadata() -> tuple[bool, dict]:
@@ -12,20 +12,3 @@ def get_pdl_launch_metadata() -> tuple[bool, dict]:
"""
enable_pdl = is_arch_support_pdl()
return enable_pdl, ({"launch_pdl": True} if enable_pdl else {})
@triton.jit
def _resolve_token_positions(
sorted_token_ids, seg_start, s_offset, seg_len, SORTED_BY_ADAPTER: tl.constexpr
):
"""Map logical segment offsets to physical token positions.
When SORTED_BY_ADAPTER is True, segments are grouped by adapter and
sorted_token_ids provides the indirection to the original token rows.
When False, tokens are already contiguous starting at seg_start.
"""
if SORTED_BY_ADAPTER:
return tl.load(
sorted_token_ids + seg_start + s_offset, mask=s_offset < seg_len
).to(tl.int64)
return (seg_start + s_offset).to(tl.int64)
@@ -6,9 +6,8 @@ import torch
from sglang.kernels.jit.utils import (
cache_once,
get_jit_cuda_arch,
get_activation_cuda_cflags,
is_arch_support_pdl,
is_hip_runtime,
load_jit,
make_cpp_args,
)
@@ -21,16 +20,6 @@ def _make_name(*args):
return "kimi_k3_" + "_".join(str(a) for a in args)
def _fast_math_flags() -> list[str]:
# Mirrors sgl-kernel's CMake policy: fast-math on SM90, precise on
# SM100+ (Blackwell needs bit-exact expf), off on HIP (clang rejects).
if is_hip_runtime():
return []
if get_jit_cuda_arch().major >= 10:
return []
return ["--use_fast_math"]
@cache_once
def _jit_situ_and_mul_module(in_dtype: torch.dtype, out_dtype: torch.dtype) -> Module:
"""Compile and cache the JIT SiTU-and-mul module for an (in, out) dtype pair."""
@@ -40,7 +29,7 @@ def _jit_situ_and_mul_module(in_dtype: torch.dtype, out_dtype: torch.dtype) -> M
*args,
cuda_files=["kimi_k3/situ_and_mul.cuh"],
cuda_wrappers=[("run", f"SituAndMulKernel<{args}>::run")],
extra_cuda_cflags=_fast_math_flags(),
extra_cuda_cflags=get_activation_cuda_cflags(),
)
@@ -12,9 +12,7 @@ import triton.language as tl
from sglang.kernels.ops.gemm.trtllm_lora_temp.kernel_utils import (
get_pdl_launch_metadata,
)
from sglang.kernels.ops.moe.moe_align import (
moe_align_block_size as jit_moe_align_block_size,
)
from sglang.kernels.ops.moe.virtual_experts import _align_block_size_large
from sglang.srt.lora.trtllm_lora_temp.environ import lora_envs
@@ -404,195 +402,6 @@ from sglang.srt.lora.trtllm_lora_temp.specialized_expand import ( # noqa: E402,
)
def _align_block_size_jit(
topk_ids: torch.Tensor,
block_size: int,
num_experts: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""CUDA JIT alignment for up to 8191 experts.
Expert IDs are shifted by one so ``-1`` maps to a sentinel bucket. The
fused allocation stays int4-aligned for the kernel's vectorized clear.
"""
assert num_experts <= 8191, (
f"_align_block_size_jit supports at most 8191 experts "
f"(num_moe_experts * max_loras), got {num_experts}"
)
device = topk_ids.device
flat_topk_ids = topk_ids.reshape(-1)
if flat_topk_ids.dtype == torch.int64:
flat_topk_ids = flat_topk_ids.to(torch.int32)
num_total_tokens = flat_topk_ids.numel()
if num_total_tokens == 0:
empty = torch.empty(0, dtype=torch.int32, device=device)
return empty, empty, torch.zeros(1, dtype=torch.int32, device=device)
# JIT kernel uses +1 offset convention: -1 -> bucket 0 (sentinel),
# expert i -> bucket i+1. So pass num_experts + 1 as the bucket count.
jit_num_experts = num_experts + 1
if num_total_tokens < jit_num_experts:
max_num_tokens_padded = num_total_tokens * block_size
else:
max_num_tokens_padded = num_total_tokens + jit_num_experts * (block_size - 1)
# Align every sub-buffer offset to a multiple of 4 (VEC_SIZE). The CUDA
# kernel fills sorted_token_ids with vectorized int4 writes whose last
# store can spill up to 3 int32s past the logical end. With a fused
# allocation the spill would corrupt the adjacent sub-buffer.
_A4 = lambda n: (n + 3) & ~3 # noqa: E731
max_num_tokens_padded = _A4(max_num_tokens_padded)
max_num_m_blocks = (max_num_tokens_padded + block_size - 1) // block_size
max_num_m_blocks_padded = _A4(max_num_m_blocks)
num_post_pad_size = _A4(1) # 1 element, padded to 4
cumsum_size = _A4(jit_num_experts + 1)
# Single allocation sliced into 4 views (zero-copy) to avoid
# per-call Python overhead of 4 separate torch.empty calls.
total_buf = (
max_num_tokens_padded
+ max_num_m_blocks_padded
+ num_post_pad_size
+ cumsum_size
)
buf = torch.empty(total_buf, dtype=torch.int32, device=device)
off = 0
sorted_token_ids = buf[off : off + max_num_tokens_padded]
off += max_num_tokens_padded
expert_ids = buf[off : off + max_num_m_blocks]
off += max_num_m_blocks_padded
num_tokens_post_padded = buf[off : off + 1]
off += num_post_pad_size
cumsum_buffer = buf[off : off + jit_num_experts + 1]
jit_moe_align_block_size(
flat_topk_ids,
jit_num_experts,
block_size,
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
cumsum_buffer,
True, # pad_sorted_token_ids
)
return sorted_token_ids, expert_ids, num_tokens_post_padded
@torch.compile(dynamic=True)
def _align_block_size_torch(
topk_ids: torch.Tensor,
block_size: int,
num_experts: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Pure-PyTorch align_block_size for num_experts > 1024, compiled via torch.compile.
Fallback for platforms where the CUDA JIT kernel is unavailable (e.g. AMD/ROCm).
Out-of-range topk_ids (negative sentinels left by EP dispatch, or virtual-
expert IDs >= num_experts produced when those sentinels are combined with
a per-adapter offset) are routed into a dedicated sentinel bucket. Without
this, indexing ``padded_offsets[sorted_expert_ids]`` would wrap (-1) or
OOB-read, and the bad expert ids would propagate into the downstream LoRA
GEMM as real expert slots.
"""
device = topk_ids.device
flat_topk_ids = topk_ids.reshape(-1).to(torch.int64)
num_total_tokens = flat_topk_ids.numel()
sentinel = num_experts
valid_mask = (flat_topk_ids >= 0) & (flat_topk_ids < num_experts)
safe_topk_ids = torch.where(
valid_mask,
flat_topk_ids,
torch.full_like(flat_topk_ids, sentinel),
)
bucket_count = num_experts + 1
max_total_padded_tokens = (
(num_total_tokens + bucket_count * (block_size - 1) + block_size - 1)
// block_size
) * block_size
max_num_blocks = max_total_padded_tokens // block_size
sorted_token_ids = torch.full(
(max_total_padded_tokens,),
num_total_tokens,
dtype=torch.int32,
device=device,
)
expert_ids = torch.full(
(max_num_blocks,),
-1,
dtype=torch.int32,
device=device,
)
if num_total_tokens == 0:
num_tokens_post_padded = torch.zeros((1,), dtype=torch.int32, device=device)
return sorted_token_ids, expert_ids, num_tokens_post_padded
sorted_order = torch.argsort(safe_topk_ids)
sorted_expert_ids = safe_topk_ids[sorted_order]
expert_range = torch.arange(bucket_count, device=device, dtype=torch.int64)
counts_offsets = torch.searchsorted(sorted_expert_ids, expert_range, right=False)
counts_end = torch.searchsorted(sorted_expert_ids, expert_range, right=True)
counts = counts_end - counts_offsets
padded_counts = ((counts + block_size - 1) // block_size) * block_size
total_padded_tokens = padded_counts.sum().to(torch.int32).reshape(1)
padded_offsets = torch.cumsum(padded_counts, dim=0) - padded_counts
token_ranks = (
torch.arange(num_total_tokens, device=device, dtype=torch.int64)
- counts_offsets[sorted_expert_ids]
)
output_positions = padded_offsets[sorted_expert_ids] + token_ranks
sorted_token_ids.scatter_(
0,
output_positions.to(torch.int64),
sorted_order.to(torch.int32),
)
block_counts = padded_counts // block_size
real_block_counts = block_counts.clone()
real_block_counts[sentinel] = 0
actual_num_blocks = real_block_counts.sum()
if max_num_blocks <= 0:
return sorted_token_ids, expert_ids, total_padded_tokens
block_offsets = torch.cumsum(real_block_counts, dim=0)
all_block_positions = torch.arange(max_num_blocks, device=device, dtype=torch.int64)
assigned_experts = torch.searchsorted(
block_offsets, all_block_positions, right=True
).to(torch.int32)
expert_ids.copy_(
torch.where(
all_block_positions < actual_num_blocks,
assigned_experts,
torch.full_like(assigned_experts, -1),
)
)
return sorted_token_ids, expert_ids, total_padded_tokens
def _align_block_size_large(
topk_ids: torch.Tensor,
block_size: int,
num_experts: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Dispatch to the CUDA JIT kernel when available, otherwise fall back to
the pure-PyTorch torch.compile path (needed on AMD/ROCm or when the JIT
module fails to load)."""
try:
return _align_block_size_jit(topk_ids, block_size, num_experts)
except Exception:
return _align_block_size_torch(topk_ids, block_size, num_experts)
def _merged_experts_fused_moe_lora_add_fake(
output: torch.Tensor,
hidden_states: torch.Tensor,
@@ -16,20 +16,6 @@ if TYPE_CHECKING:
from sglang.srt.layers.moe.topk import TopKConfig, TopKOutput
def _apply_routed_scaling_after_renorm(
topk_weights: torch.Tensor,
topk_config: "TopKConfig",
) -> torch.Tensor:
"""Mirror GPU post-renorm scaling when apply_routed_scaling_factor_on_output is set."""
if (
topk_config.renormalize
and topk_config.apply_routed_scaling_factor_on_output
and topk_config.routed_scaling_factor is not None
):
return topk_weights * topk_config.routed_scaling_factor
return topk_weights
def fused_topk_npu(
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
@@ -38,7 +24,6 @@ def fused_topk_npu(
expert_location_dispatch_info: Optional["ExpertLocationDispatchInfo"] = None,
layer_id: Optional[int] = None,
) -> "TopKOutput":
use_grouped_topk = topk_config.use_grouped_topk
renormalize = topk_config.renormalize
correction_bias = topk_config.correction_bias
@@ -50,83 +50,6 @@ def _extract_positions_from_plan(
return positions
def _compress_forward_c128_fallback(
kv_score_buffer: torch.Tensor,
kv_score_input: torch.Tensor,
ape: torch.Tensor,
plan: Union[CompressorDecodePlan, CompressorPrefillPlan],
head_dim: int,
) -> torch.Tensor:
"""PyTorch fallback for C128 compress_forward on HIP (wave64).
Fully vectorized, compatible with CUDA graph capture.
kv_score_buffer: [num_pages, 128, head_dim * 2]
ape: [128, head_dim]
IMPORTANT: This also performs the write to state buffer (like the JIT kernel).
The JIT kernel does: (1) write kv_score_input to buffer, (2) compress from buffer.
"""
num_total_slots = kv_score_buffer.shape[0] * kv_score_buffer.shape[1]
num_pages = kv_score_buffer.shape[0]
last_dim = kv_score_buffer.shape[-1]
# Step 1: WRITE kv_score_input to state buffer
if num_total_slots > 0:
buf_flat = kv_score_buffer.view(-1, last_dim)
if plan.is_decode:
# Decode: plan_d has write_loc per batch item
plan_raw = plan[1].view(torch.int32) # [bs, 4]
write_locs = plan_raw[:, 1].long()
# Only write valid locations (>= 0 and < buffer size)
valid_write = (write_locs >= 0) & (write_locs < num_total_slots)
if valid_write.any():
buf_flat[write_locs[valid_write]] = kv_score_input[valid_write]
else:
# Prefill: plan_w has {ragged_id, write_loc} per write entry
plan_w = plan[2] # [num_w, 8] uint8 = WritePlan
if plan_w.shape[0] > 0:
plan_w_raw = plan_w.view(torch.int32) # [num_w, 2]
ragged_ids = plan_w_raw[:, 0].long() & 0xFFFF
write_locs = plan_w_raw[:, 1].long()
valid_write = (write_locs >= 0) & (write_locs < num_total_slots)
ragged_ids_safe = ragged_ids.clamp(
min=0, max=kv_score_input.shape[0] - 1
)
if valid_write.any():
buf_flat[write_locs[valid_write]] = kv_score_input[
ragged_ids_safe[valid_write]
]
# Step 2: COMPRESS (read from buffer page and do softmax-pool)
plan_c = plan[1] # plan_d for decode, plan_c for prefill
num_tokens = plan_c.shape[0]
if num_pages == 0 or num_tokens == 0:
return kv_score_input.new_zeros(num_tokens, head_dim)
plan_c_raw = plan_c.view(torch.int32) # [N, 4]
read_page_0 = plan_c_raw[:, 2].long()
# Use torch.where instead of clamp to handle -1 (invalid) gracefully
valid_read = (read_page_0 >= 0) & (read_page_0 < num_pages)
read_page_0_safe = torch.where(
valid_read, read_page_0, torch.zeros_like(read_page_0)
)
gathered = kv_score_buffer[read_page_0_safe] # [N, 128, head_dim*2]
kv = gathered[:, :, :head_dim].float()
score = gathered[:, :, head_dim:].float() + ape.float().unsqueeze(0)
weights = score.softmax(dim=1)
out = (weights * kv).sum(dim=1)
# For decode: zero out non-boundary tokens (seq_len % 128 != 0)
# so they don't corrupt kvcache location 0 when stored.
if plan.is_decode:
seq_lens = plan_c_raw[:, 0].to(torch.int32)
is_boundary = (seq_lens % 128 == 0).unsqueeze(-1) # [N, 1]
out = torch.where(is_boundary, out, torch.zeros_like(out))
return out.to(kv_score_input.dtype)
class CompressorBackendMixin:
def __init__(self):
super().__init__()
@@ -28,7 +28,7 @@ if _is_cuda:
moe_sum_reduce_triton,
)
from sglang.kernels.ops.moe.moe_wna16_marlin import moe_wna16_marlin_gemm
from sglang.kernels.ops.moe.trtllm_lora_temp.virtual_experts import (
from sglang.kernels.ops.moe.virtual_experts import (
_align_block_size_jit as moe_align_block_size,
)
from sglang.srt.layers.moe.fused_moe_triton.fused_marlin_moe import (
@@ -3,11 +3,9 @@
import asyncio
import base64
import copy
import json
import math
import os
import re
import subprocess
from concurrent.futures import ThreadPoolExecutor, as_completed
from dataclasses import dataclass, field
from io import BytesIO
@@ -231,39 +229,6 @@ def _decode_frames_and_timestamps(vdw, ele):
return video_tensor, timestamps
def _ffprobe_has_audio(src, stdin=None, label=None) -> bool:
# Header-only audio-stream probe for HTTP URLs; avoids full download.
try:
r = subprocess.run(
[
"ffprobe",
"-v",
"quiet",
"-print_format",
"json",
"-show_streams",
"-select_streams",
"a",
src,
],
input=stdin,
capture_output=True,
timeout=30,
)
if r.returncode != 0:
stderr = r.stderr.decode("utf-8", errors="replace")
raise RuntimeError(f"ffprobe failed for {label}: {stderr}")
return bool(json.loads(r.stdout).get("streams"))
except subprocess.TimeoutExpired:
logger.error("ffprobe timed out for %s", label)
raise
except FileNotFoundError as e:
raise RuntimeError("ffprobe not found; install ffmpeg") from e
except json.JSONDecodeError:
logger.error("ffprobe returned invalid JSON for %s", label)
raise
class MiMoProcessor:
def __init__(
self,
@@ -690,7 +655,6 @@ class MiMoProcessor:
def process_video(
self, video_input: VideoInput | VideoAudioInput, temporal_padding_factor=None
):
def smart_resize_video(
num_total_frames, min_pixels, max_pixels, total_max_pixels, **kwargs
):