diff --git a/python/sglang/kernels/jit/csrc/dsa/kpool_topk_transform.cuh b/python/sglang/kernels/jit/csrc/dsa/kpool_topk_transform.cuh deleted file mode 100644 index 308ec1ced..000000000 --- a/python/sglang/kernels/jit/csrc/dsa/kpool_topk_transform.cuh +++ /dev/null @@ -1,440 +0,0 @@ -/** - * @NOTE: The radix top-k core (fast_topk_cuda_tl_impl) is adapted from - * https://github.com/tile-ai/tilelang/blob/main/examples/deepseek_v32/topk_selector.py - * and was previously shipped as an AOT sgl-kernel op (fast_kpool_topk_transform_fused). - * It is re-implemented here as a lightweight JIT kernel for the NSA kpool indexer: - * select pool groups at pool granularity, expand each group to `pool_size` token - * indices, and optionally transform those indices through a page table or ragged offset. - * - * The pool-level top-k value is a compile-time constant injected via -DSGL_GROUP_TOPK. - */ -#include // For TensorMatcher, SymbolicSize, SymbolicDevice, is_type -#include // For RuntimeCheck, RuntimeDeviceCheck - -#include // For LaunchKernel, type aliases - -#include -#include - -#include -#include -#include -#include - -namespace { - -#ifndef C10_LIKELY -#define C10_LIKELY(expr) (__builtin_expect(static_cast(expr), 1)) -#endif - -#ifndef SGL_GROUP_TOPK -#define SGL_GROUP_TOPK 256 -#endif - -// Compile-time pool-level top-k (number of groups selected per row). -inline constexpr int kGroupTopK = SGL_GROUP_TOPK; -inline constexpr int kThreadsPerBlock = 1024; - -// Reduced from 128KB to 32KB to improve occupancy. -// Each radix pass needs at most ~K candidates in the threshold bin, -// so 4K entries per round (2 rounds = 8K entries = 32KB) is sufficient. -inline constexpr std::size_t kSmem = 8 * 1024 * sizeof(uint32_t); // 32KB (bytes) - -struct FastTopKParams { - const float* __restrict__ input; // [B, input_stride] - const int32_t* __restrict__ row_starts; // [B] or nullptr - int32_t* __restrict__ indices; // unused here (kept for layout parity) - const int32_t* __restrict__ lengths; // [B] - int64_t input_stride; -}; - -__device__ __forceinline__ auto convert_to_uint8(float x) -> uint8_t { - __half h = __float2half_rn(x); - uint16_t bits = __half_as_ushort(h); - uint16_t key = (bits & 0x8000) ? static_cast(~bits) : static_cast(bits | 0x8000); - return static_cast(key >> 8); -} - -__device__ __forceinline__ auto convert_to_uint32(float x) -> uint32_t { - uint32_t bits = __float_as_uint(x); - return (bits & 0x80000000u) ? ~bits : (bits | 0x80000000u); -} - -template -__device__ void -fast_topk_cuda_tl_impl(const float* __restrict__ input, int* __restrict__ index, int row_start, int length) { - // An optimized topk kernel copied from tilelang kernel - // We assume length > K here, or it will crash - int topk = K; - constexpr auto BLOCK_SIZE = 1024; - constexpr auto RADIX = 256; - constexpr auto SMEM_INPUT_SIZE = kSmem / (2 * sizeof(int)); - - alignas(128) __shared__ int s_histogram_buf[2][RADIX + 128]; - alignas(128) __shared__ int s_counter; - alignas(128) __shared__ int s_threshold_bin_id; - alignas(128) __shared__ int s_num_input[2]; - - auto& s_histogram = s_histogram_buf[0]; - // allocate for two rounds - extern __shared__ int s_input_idx[][SMEM_INPUT_SIZE]; - - const int tx = threadIdx.x; - - // stage 1: 8bit coarse histogram - if (tx < RADIX + 1) s_histogram[tx] = 0; - __syncthreads(); - - for (int idx = tx; idx < length; idx += BLOCK_SIZE) { - const auto bin = convert_to_uint8(input[idx + row_start]); - ::atomicAdd(&s_histogram[bin], 1); - } - __syncthreads(); - - const auto run_cumsum = [&] { -#pragma unroll 8 - for (int i = 0; i < 8; ++i) { - static_assert(1 << 8 == RADIX); - if (C10_LIKELY(tx < RADIX)) { - const auto j = 1 << i; - const auto k = i & 1; - auto value = s_histogram_buf[k][tx]; - if (tx < RADIX - j) { - value += s_histogram_buf[k][tx + j]; - } - s_histogram_buf[k ^ 1][tx] = value; - } - __syncthreads(); - } - }; - - run_cumsum(); - if (tx < RADIX && s_histogram[tx] > topk && s_histogram[tx + 1] <= topk) { - s_threshold_bin_id = tx; - s_num_input[0] = 0; - s_counter = 0; - } - __syncthreads(); - - const auto threshold_bin = s_threshold_bin_id; - topk -= s_histogram[threshold_bin + 1]; - - if (topk == 0) { - for (int idx = tx; idx < length; idx += BLOCK_SIZE) { - const auto bin = static_cast(convert_to_uint8(input[idx + row_start])); - if (bin > threshold_bin) { - const auto pos = ::atomicAdd(&s_counter, 1); - index[pos] = idx; - } - } - __syncthreads(); - return; - } else { - __syncthreads(); - if (tx < RADIX + 1) { - s_histogram[tx] = 0; - } - __syncthreads(); - - for (int idx = tx; idx < length; idx += BLOCK_SIZE) { - const auto raw_input = input[idx + row_start]; - const auto bin = static_cast(convert_to_uint8(raw_input)); - if (bin > threshold_bin) { - const auto pos = ::atomicAdd(&s_counter, 1); - index[pos] = idx; - } else if (bin == threshold_bin) { - const auto pos = ::atomicAdd(&s_num_input[0], 1); - /// NOTE: (dark) fuse the histogram computation here - if (C10_LIKELY(pos < SMEM_INPUT_SIZE)) { - s_input_idx[0][pos] = idx; - const auto bin = convert_to_uint32(raw_input); - const auto sub_bin = (bin >> 24) & 0xFF; - ::atomicAdd(&s_histogram[sub_bin], 1); - } - } - } - __syncthreads(); - } - - // stage 2: refine with 8bit radix passes -#pragma unroll 4 - for (int round = 0; round < 4; ++round) { - __shared__ int s_last_remain; - const auto r_idx = round % 2; - - // clip here to prevent overflow - const auto _raw_num_input = s_num_input[r_idx]; - const auto num_input = (_raw_num_input < int(SMEM_INPUT_SIZE)) ? _raw_num_input : int(SMEM_INPUT_SIZE); - - run_cumsum(); - if (tx < RADIX && s_histogram[tx] > topk && s_histogram[tx + 1] <= topk) { - s_threshold_bin_id = tx; - s_num_input[r_idx ^ 1] = 0; - s_last_remain = topk - s_histogram[tx + 1]; - } - __syncthreads(); - - const auto threshold_bin = s_threshold_bin_id; - topk -= s_histogram[threshold_bin + 1]; - - if (topk == 0) { - for (int i = tx; i < num_input; i += BLOCK_SIZE) { - const auto idx = s_input_idx[r_idx][i]; - const auto offset = 24 - round * 8; - const auto bin = (convert_to_uint32(input[idx + row_start]) >> offset) & 0xFF; - if (bin > threshold_bin) { - const auto pos = ::atomicAdd(&s_counter, 1); - index[pos] = idx; - } - } - __syncthreads(); - break; - } else { - __syncthreads(); - if (tx < RADIX + 1) { - s_histogram[tx] = 0; - } - __syncthreads(); - for (int i = tx; i < num_input; i += BLOCK_SIZE) { - const auto idx = s_input_idx[r_idx][i]; - const auto raw_input = input[idx + row_start]; - const auto offset = 24 - round * 8; - const auto bin = (convert_to_uint32(raw_input) >> offset) & 0xFF; - if (bin > threshold_bin) { - const auto pos = ::atomicAdd(&s_counter, 1); - index[pos] = idx; - } else if (bin == threshold_bin) { - if (round == 3) { - const auto pos = ::atomicAdd(&s_last_remain, -1); - if (pos > 0) { - index[K - pos] = idx; - } - } else { - const auto pos = ::atomicAdd(&s_num_input[r_idx ^ 1], 1); - if (C10_LIKELY(pos < SMEM_INPUT_SIZE)) { - /// NOTE: (dark) fuse the histogram computation here - s_input_idx[r_idx ^ 1][pos] = idx; - const auto bin = convert_to_uint32(raw_input); - const auto sub_bin = (bin >> (offset - 8)) & 0xFF; - ::atomicAdd(&s_histogram[sub_bin], 1); - } - } - } - } - __syncthreads(); - } - } -} - -__device__ __forceinline__ int32_t transform_kpool_token( - int32_t raw_token, - const int32_t* __restrict__ page_table_entry, - const int32_t* __restrict__ topk_indices_offset, - int32_t offset) { - if (page_table_entry != nullptr) { - return page_table_entry[raw_token]; - } - if (topk_indices_offset != nullptr) { - return raw_token + offset; - } - return raw_token; -} - -template -__global__ __launch_bounds__(kThreadsPerBlock) void kpool_topk_transform_kernel( - const __grid_constant__ FastTopKParams params, - int32_t* __restrict__ dst_token_indices, - const int64_t dst_stride, - const int32_t pool_size, - const int32_t token_topk, - const int32_t out_cols, - const int32_t* __restrict__ page_table, - const int64_t page_table_stride, - const int32_t* __restrict__ topk_indices_offset, - const int32_t* __restrict__ seq_lens) { - const auto& [input, row_starts, _, lengths, input_stride] = params; - const auto bid = static_cast(blockIdx.x); - const auto tid = threadIdx.x; - const auto row_start = row_starts == nullptr ? 0 : row_starts[bid]; - const auto length = lengths[bid]; - const auto score = input + bid * input_stride; - const auto dst = dst_token_indices + bid * dst_stride; - const auto page_table_entry = page_table == nullptr ? nullptr : page_table + bid * page_table_stride; - const auto offset = topk_indices_offset == nullptr ? 0 : topk_indices_offset[bid]; - const bool append_tail = seq_lens != nullptr; - const auto full_pool_token_len = length * pool_size; - const auto history_len = full_pool_token_len < token_topk ? full_pool_token_len : token_topk; - const auto tail_count = append_tail ? seq_lens[bid] % pool_size : 0; - - if (length <= K) { - for (int col = tid; col < out_cols; col += kThreadsPerBlock) { - if (col < history_len) { - const auto group_rank = col / pool_size; - const auto slot = col % pool_size; - const auto raw_token = group_rank * pool_size + slot; - dst[col] = transform_kpool_token(raw_token, page_table_entry, topk_indices_offset, offset); - } else if (append_tail && col < history_len + tail_count) { - const auto raw_token = length * pool_size + (col - history_len); - dst[col] = transform_kpool_token(raw_token, page_table_entry, topk_indices_offset, offset); - } else { - dst[col] = -1; - } - } - return; - } - - __shared__ int s_indices[K]; - fast_topk_cuda_tl_impl(score, s_indices, row_start, length); - for (int col = tid; col < out_cols; col += kThreadsPerBlock) { - if (col < history_len) { - const auto group_rank = col / pool_size; - const auto group_id = s_indices[group_rank]; - const auto slot = col % pool_size; - const auto raw_token = group_id * pool_size + slot; - dst[col] = transform_kpool_token(raw_token, page_table_entry, topk_indices_offset, offset); - } else if (append_tail && col < history_len + tail_count) { - const auto raw_token = length * pool_size + (col - history_len); - dst[col] = transform_kpool_token(raw_token, page_table_entry, topk_indices_offset, offset); - } else { - dst[col] = -1; - } - } -} - -template -void setup_kernel_smem_once(host::DebugInfo where = {}) { - [[maybe_unused]] - static const auto result = [] { - const auto fptr = std::bit_cast(f); - return ::cudaFuncSetAttribute(fptr, ::cudaFuncAttributeMaxDynamicSharedMemorySize, kMaxDynamicSMEM); - }(); - host::RuntimeDeviceCheck(result, where); -} - -template -const T* optional_data_ptr(const tvm::ffi::Optional& opt) { - if (!opt.has_value()) { - return nullptr; - } - return static_cast(opt.value().data_ptr()); -} - -struct KpoolTopKTransformKernel { - static constexpr auto kernel = kpool_topk_transform_kernel; - - // Pool-level radix top-k for the NSA kpool indexer. - // score : [B, S] strided float32 scores (one score per pool group) - // lengths : [B] int32 valid group count per row - // dst_token_indices : [B, out_cols] int32 output token indices (contiguous) - // pool_size : tokens per pool group - // page_table (opt) : [B, P] strided int32 raw-token -> real-token map - // topk_indices_offset : [B] int32 per-row offset added to raw tokens (ragged) - // row_starts (opt) : [B] int32 score row start offsets - // seq_lens (opt) : [B] int32 sequence lengths; enables tail append - static void transform( - const tvm::ffi::TensorView score, - const tvm::ffi::TensorView lengths, - const tvm::ffi::TensorView dst_token_indices, - const int64_t pool_size, - const tvm::ffi::Optional page_table_opt, - const tvm::ffi::Optional topk_indices_offset_opt, - const tvm::ffi::Optional row_starts_opt, - const tvm::ffi::Optional seq_lens_opt) { - using namespace host; - - auto B = SymbolicSize{"batch_size"}; - auto S = SymbolicSize{"score_stride"}; - auto out_cols_sym = SymbolicSize{"out_cols"}; - auto device = SymbolicDevice{}; - device.set_options(); - - TensorMatcher({B, -1}) // strided scores - .with_strides({S, 1}) - .with_dtype() - .with_device(device) - .verify(score); - TensorMatcher({B}) // lengths, contiguous int32 - .with_dtype() - .with_device(device) - .verify(lengths); - TensorMatcher({B, out_cols_sym}) // output, contiguous int32 - .with_dtype() - .with_device(device) - .verify(dst_token_indices); - - RuntimeCheck(pool_size > 1, "pool_size must be > 1, got ", pool_size); - RuntimeCheck( - !(page_table_opt.has_value() && topk_indices_offset_opt.has_value()), - "page_table and topk_indices_offset are mutually exclusive"); - - const auto out_cols = static_cast(out_cols_sym.unwrap()); - const auto tail_cols = seq_lens_opt.has_value() ? static_cast(pool_size) - 1 : 0; - RuntimeCheck(out_cols > tail_cols, "dst_token_indices columns ", out_cols, " must exceed tail ", tail_cols); - const auto token_topk = out_cols - tail_cols; - RuntimeCheck(token_topk % static_cast(pool_size) == 0, "token_topk must be a multiple of pool_size"); - RuntimeCheck( - token_topk / static_cast(pool_size) == kGroupTopK, - "this module is built for group_topk=", - kGroupTopK, - " but got ", - token_topk / static_cast(pool_size)); - - const auto batch_size = static_cast(B.unwrap()); - - int64_t page_table_stride = 0; - const int32_t* page_table_ptr = nullptr; - if (page_table_opt.has_value()) { - auto P = SymbolicSize{"page_table_stride"}; - TensorMatcher({B, -1}) // strided page table - .with_strides({P, 1}) - .with_dtype() - .with_device(device) - .verify(page_table_opt.value()); - page_table_ptr = static_cast(page_table_opt.value().data_ptr()); - page_table_stride = static_cast(P.unwrap()); - } - - if (topk_indices_offset_opt.has_value()) { - TensorMatcher({B}) // - .with_dtype() - .with_device(device) - .verify(topk_indices_offset_opt.value()); - } - if (row_starts_opt.has_value()) { - TensorMatcher({B}) // - .with_dtype() - .with_device(device) - .verify(row_starts_opt.value()); - } - if (seq_lens_opt.has_value()) { - TensorMatcher({B}) // - .with_dtype() - .with_device(device) - .verify(seq_lens_opt.value()); - } - - const auto params = FastTopKParams{ - .input = static_cast(score.data_ptr()), - .row_starts = optional_data_ptr(row_starts_opt), - .indices = nullptr, - .lengths = static_cast(lengths.data_ptr()), - .input_stride = static_cast(S.unwrap()), - }; - - setup_kernel_smem_once(); - LaunchKernel(batch_size, kThreadsPerBlock, device.unwrap(), kSmem)( - kernel, - params, - static_cast(dst_token_indices.data_ptr()), - static_cast(dst_token_indices.strides()[0]), - static_cast(pool_size), - token_topk, - out_cols, - page_table_ptr, - page_table_stride, - optional_data_ptr(topk_indices_offset_opt), - optional_data_ptr(seq_lens_opt)); - } -}; - -} // namespace diff --git a/python/sglang/kernels/jit/csrc/elementwise/resolve_future_token_ids.cuh b/python/sglang/kernels/jit/csrc/elementwise/resolve_future_token_ids.cuh deleted file mode 100644 index ef178ba63..000000000 --- a/python/sglang/kernels/jit/csrc/elementwise/resolve_future_token_ids.cuh +++ /dev/null @@ -1,57 +0,0 @@ -#include // For TensorMatcher, SymbolicSize, SymbolicDevice -#include // For RuntimeCheck, div_ceil - -#include // For LaunchKernel - -#include -#include - -#include -#include - -namespace { - -template -__global__ void resolve_future_token_ids_kernel(T* __restrict__ input_ids, const T* __restrict__ future_map, size_t n) { - size_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) { - T val = input_ids[idx]; - if (val < 0) { - T key = -val; - if (key < 0) key = 0; // clamp for overflow - input_ids[idx] = future_map[key]; - } - } -} - -constexpr size_t kBlockSize = 256; - -template -struct ResolveFutureTokenIds { - static void run(tvm::ffi::TensorView input_ids, tvm::ffi::TensorView future_map) { - using namespace host; - - SymbolicSize N = {"num_tokens"}; - SymbolicSize M = {"map_size"}; - SymbolicDevice device_; - device_.set_options(); - - TensorMatcher({N}).with_dtype().with_device(device_).verify(input_ids); - - TensorMatcher({M}).with_dtype().with_device(device_).verify(future_map); - - const size_t num_tokens = N.unwrap(); - if (num_tokens == 0) return; - - const size_t grid_size = div_ceil(num_tokens, kBlockSize); - const DLDevice device = device_.unwrap(); - - LaunchKernel(grid_size, kBlockSize, device)( - resolve_future_token_ids_kernel, - static_cast(input_ids.data_ptr()), - static_cast(future_map.data_ptr()), - num_tokens); - } -}; - -} // namespace diff --git a/python/sglang/kernels/ops/attention/__init__.py b/python/sglang/kernels/ops/attention/__init__.py index 781522e36..4d77dc1bc 100644 --- a/python/sglang/kernels/ops/attention/__init__.py +++ b/python/sglang/kernels/ops/attention/__init__.py @@ -97,9 +97,7 @@ for _mod, _fn in [ ("dsa.cp_split", "dsa_cp_round_robin_split_q_seqs_kernel"), ("dsv4.fp4_indexer", "quantize_fp4_indexer_tensor"), ("dsv4.fp4_indexer", "store_fp4_index_k_cache"), - ("dsv4.fused_scale", "fused_scale"), ("dsv4.rms_normalize_hip", "rms_normalize_triton"), - ("dsv4.compress_c128_hip", "_compress_forward_c128_triton"), ]: register_kernel( KernelSpec( @@ -136,7 +134,6 @@ for _mod, _fn in [ ("deepseek_v4_rope", "precompute_freqs_cis"), ("fused_qk_norm_rope_store", "fused_qk_norm_rope_swa_store"), ("fused_qk_rmsnorm_rope_gate", "fused_qk_gemma_rmsnorm_rope_gate"), - ("fused_qk_norm", "fused_qk_norm"), ("rotary_triton", "triton_mrope_fused"), ("rotary_triton", "triton_ernie45_rope_fused_inplace"), ("mrope", "apply_interleaved_rope_triton"), diff --git a/python/sglang/kernels/ops/attention/dsv4/compress_c128_hip.py b/python/sglang/kernels/ops/attention/dsv4/compress_c128_hip.py deleted file mode 100644 index 84653d518..000000000 --- a/python/sglang/kernels/ops/attention/dsv4/compress_c128_hip.py +++ /dev/null @@ -1,292 +0,0 @@ -"""HIP c128 compression kernels for DSV4 (RFC #29630, Phase 2.5). - -Migrated from ``sglang.srt.layers.attention.dsv4.compressor_v2``; the -kernels are defined under an ``is_hip`` guard exactly as before. -""" - -from __future__ import annotations - -from typing import Union - -import torch - -from sglang.kernels.jit.utils import is_hip_runtime -from sglang.kernels.ops.attention.dsv4 import ( - CompressorDecodePlan, - CompressorPrefillPlan, -) - -_is_hip = is_hip_runtime() - -if _is_hip: - import triton - import triton.language as tl - - @triton.jit - def _c128_compress_decode_kernel( - buf_ptr, - input_ptr, - ape_ptr, - out_ptr, - plan_ptr, - buf_stride_slot, - input_stride_b, - ape_stride_r, - out_stride_b, - bs, - HEAD_DIM: tl.constexpr, - BLOCK_D: tl.constexpr, - COMPRESS_RATIO: tl.constexpr, - ): - """Fused C128 decode: write to state buffer + online softmax-pool. - - plan_ptr points to int32 view: [bs, 4] where each row is - {seq_len, write_loc, read_page_0, read_page_1}. - """ - bid = tl.program_id(0) - if bid >= bs: - return - - # Parse plan - plan_base = plan_ptr + bid * 4 - seq_len = tl.load(plan_base).to(tl.int32) - write_loc = tl.load(plan_base + 1).to(tl.int32) - read_page_0 = tl.load(plan_base + 2).to(tl.int32) - - d = tl.arange(0, BLOCK_D) - last_dim: tl.constexpr = HEAD_DIM * 2 - - # Step 1: Write kv_score_input to state buffer at write_loc - d_mask_full = d < last_dim - input_val = tl.load( - input_ptr + bid * input_stride_b + d, mask=d_mask_full, other=0.0 - ) - tl.store(buf_ptr + write_loc * buf_stride_slot + d, input_val, mask=d_mask_full) - - # Step 2: Check boundary condition - d_mask_hd = d < HEAD_DIM - if seq_len % COMPRESS_RATIO != 0: - tl.store( - out_ptr + bid * out_stride_b + d, - tl.zeros([BLOCK_D], tl.float32), - mask=d_mask_hd, - ) - return - - # Step 3: Online softmax-pool over 128 slots in the page - page_base = read_page_0 * COMPRESS_RATIO * buf_stride_slot - m_prev = tl.full([BLOCK_D], float("-inf"), tl.float32) - kv_acc = tl.zeros([BLOCK_D], tl.float32) - w_acc = tl.zeros([BLOCK_D], tl.float32) - - for k in tl.static_range(COMPRESS_RATIO): - slot_addr = page_base + k * buf_stride_slot - kv_val = tl.load(buf_ptr + slot_addr + d, mask=d_mask_hd, other=0.0).to( - tl.float32 - ) - sc_val = tl.load( - buf_ptr + slot_addr + HEAD_DIM + d, mask=d_mask_hd, other=0.0 - ).to(tl.float32) - ape_val = tl.load( - ape_ptr + k * ape_stride_r + d, mask=d_mask_hd, other=0.0 - ).to(tl.float32) - score_k = sc_val + ape_val - - m_new = tl.maximum(m_prev, score_k) - exp_old = tl.where(m_prev == float("-inf"), 0.0, tl.exp(m_prev - m_new)) - exp_cur = tl.where(score_k == float("-inf"), 0.0, tl.exp(score_k - m_new)) - kv_acc = kv_acc * exp_old + exp_cur * kv_val - w_acc = w_acc * exp_old + exp_cur - m_prev = m_new - - compressed = kv_acc / w_acc - tl.store(out_ptr + bid * out_stride_b + d, compressed, mask=d_mask_hd) - - @triton.jit - def _c128_compress_prefill_write_kernel( - buf_ptr, - input_ptr, - plan_w_ptr, - buf_stride_slot, - input_stride_b, - num_w, - BLOCK_D: tl.constexpr, - LAST_DIM: tl.constexpr, - ): - """Prefill write phase: scatter kv_score_input tokens into state buffer.""" - wid = tl.program_id(0) - if wid >= num_w: - return - - # WritePlan: {ragged_id(u32), write_loc(i32)} = 8 bytes = 2 int32s - plan_base = plan_w_ptr + wid * 2 - ragged_id = (tl.load(plan_base).to(tl.int32)) & 0xFFFF - write_loc = tl.load(plan_base + 1).to(tl.int32) - - d = tl.arange(0, BLOCK_D) - d_mask = d < LAST_DIM - - if write_loc >= 0: - input_val = tl.load( - input_ptr + ragged_id * input_stride_b + d, mask=d_mask, other=0.0 - ) - tl.store(buf_ptr + write_loc * buf_stride_slot + d, input_val, mask=d_mask) - - @triton.jit - def _c128_compress_prefill_compress_kernel( - buf_ptr, - ape_ptr, - out_ptr, - plan_c_ptr, - buf_stride_slot, - ape_stride_r, - out_stride_b, - num_c, - HEAD_DIM: tl.constexpr, - BLOCK_D: tl.constexpr, - COMPRESS_RATIO: tl.constexpr, - ): - """Prefill compress phase: online softmax-pool for each compress plan entry.""" - cid = tl.program_id(0) - if cid >= num_c: - return - - # CompressPlan: {seq_len(u32), ragged_id(u16)|buffer_len(u16), read_page_0(i32), read_page_1(i32)} - plan_base = plan_c_ptr + cid * 4 - read_page_0 = tl.load(plan_base + 2).to(tl.int32) - - d = tl.arange(0, BLOCK_D) - d_mask_hd = d < HEAD_DIM - - if read_page_0 < 0: - tl.store( - out_ptr + cid * out_stride_b + d, - tl.zeros([BLOCK_D], tl.float32), - mask=d_mask_hd, - ) - return - - page_base = read_page_0 * COMPRESS_RATIO * buf_stride_slot - m_prev = tl.full([BLOCK_D], float("-inf"), tl.float32) - kv_acc = tl.zeros([BLOCK_D], tl.float32) - w_acc = tl.zeros([BLOCK_D], tl.float32) - - for k in tl.static_range(COMPRESS_RATIO): - slot_addr = page_base + k * buf_stride_slot - kv_val = tl.load(buf_ptr + slot_addr + d, mask=d_mask_hd, other=0.0).to( - tl.float32 - ) - sc_val = tl.load( - buf_ptr + slot_addr + HEAD_DIM + d, mask=d_mask_hd, other=0.0 - ).to(tl.float32) - ape_val = tl.load( - ape_ptr + k * ape_stride_r + d, mask=d_mask_hd, other=0.0 - ).to(tl.float32) - score_k = sc_val + ape_val - - m_new = tl.maximum(m_prev, score_k) - exp_old = tl.where(m_prev == float("-inf"), 0.0, tl.exp(m_prev - m_new)) - exp_cur = tl.where(score_k == float("-inf"), 0.0, tl.exp(score_k - m_new)) - kv_acc = kv_acc * exp_old + exp_cur * kv_val - w_acc = w_acc * exp_old + exp_cur - m_prev = m_new - - compressed = kv_acc / w_acc - tl.store(out_ptr + cid * out_stride_b + d, compressed, mask=d_mask_hd) - - -def _compress_forward_c128_triton( - kv_score_buffer: torch.Tensor, - kv_score_input: torch.Tensor, - ape: torch.Tensor, - plan: Union[CompressorDecodePlan, CompressorPrefillPlan], - head_dim: int, -) -> torch.Tensor: - """Triton C128 compress_forward for HIP (wave64). - - Fuses write + online-softmax-pool into Triton kernels. - CUDA graph compatible. - """ - 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] - compress_ratio = 128 - - buf_flat = kv_score_buffer.view(-1, last_dim) - buf_stride_slot = last_dim # elements per slot - - BLOCK_D = triton.next_power_of_2(last_dim) - - if plan.is_decode: - # Decode path: single kernel does write + compress - plan_raw = plan[1].view(torch.int32) # [bs, 4] - bs = plan_raw.shape[0] - out = torch.empty( - bs, head_dim, dtype=torch.float32, device=kv_score_input.device - ) - - if bs > 0 and num_total_slots > 0: - grid = (bs,) - _c128_compress_decode_kernel[grid]( - buf_flat, - kv_score_input, - ape, - out, - plan_raw, - buf_stride_slot, - kv_score_input.stride(0), - ape.stride(0), - out.stride(0), - bs, - HEAD_DIM=head_dim, - BLOCK_D=triton.next_power_of_2(head_dim), - COMPRESS_RATIO=compress_ratio, - num_warps=8, - ) - return out - else: - # Prefill path: separate write kernel + compress kernel - plan_c_raw = plan[1].view(torch.int32) # [num_c, 4] - plan_w = plan[2] # [num_w, 8] uint8 - plan_w_raw = plan_w.view(torch.int32) # [num_w, 2] - num_c = plan_c_raw.shape[0] - num_w = plan_w_raw.shape[0] - - out = torch.empty( - num_c, head_dim, dtype=torch.float32, device=kv_score_input.device - ) - - # Phase 1: Write - if num_w > 0 and num_total_slots > 0: - grid_w = (num_w,) - _c128_compress_prefill_write_kernel[grid_w]( - buf_flat, - kv_score_input, - plan_w_raw, - buf_stride_slot, - kv_score_input.stride(0), - num_w, - BLOCK_D=BLOCK_D, - LAST_DIM=last_dim, - num_warps=4, - ) - - # Phase 2: Compress - if num_c > 0 and num_pages > 0: - grid_c = (num_c,) - _c128_compress_prefill_compress_kernel[grid_c]( - buf_flat, - ape, - out, - plan_c_raw, - buf_stride_slot, - ape.stride(0), - out.stride(0), - num_c, - HEAD_DIM=head_dim, - BLOCK_D=triton.next_power_of_2(head_dim), - COMPRESS_RATIO=compress_ratio, - num_warps=8, - ) - - return out diff --git a/python/sglang/kernels/ops/attention/dsv4/fused_scale.py b/python/sglang/kernels/ops/attention/dsv4/fused_scale.py deleted file mode 100644 index 74ba6144f..000000000 --- a/python/sglang/kernels/ops/attention/dsv4/fused_scale.py +++ /dev/null @@ -1,51 +0,0 @@ -"""Fused scale kernel for the DSV4 indexer. - -Migrated from ``sglang.srt.layers.attention.dsv4.indexer`` (RFC #29630, Phase 2.5). -""" - -import torch -import triton -import triton.language as tl - - -@triton.jit -def _fused_scale_kernel( - weight_ptr, - q_scale_ptr, - out_ptr, - numel, - out_scale, - BLOCK: tl.constexpr, -): - pid = tl.program_id(0) - offs = pid * BLOCK + tl.arange(0, BLOCK) - mask = offs < numel - - w = tl.load(weight_ptr + offs, mask=mask) - qs = tl.load(q_scale_ptr + offs, mask=mask) - - acc = w.to(tl.float32) * out_scale * qs.to(tl.float32) - tl.store(out_ptr + offs, acc.to(out_ptr.dtype.element_ty), mask=mask) - - -def fused_scale( - weight: torch.Tensor, - out_scale: float, - q_scale: torch.Tensor, -) -> torch.Tensor: - assert weight.is_contiguous() and q_scale.is_contiguous() - B, H = weight.shape - numel = B * H - out_dtype = torch.promote_types(weight.dtype, q_scale.dtype) - out = torch.empty((B, H, 1), device=weight.device, dtype=out_dtype) - BLOCK = 1024 - grid = (triton.cdiv(numel, BLOCK),) - _fused_scale_kernel[grid]( - weight, - q_scale, - out, - numel, - out_scale, - BLOCK=BLOCK, - ) - return out diff --git a/python/sglang/kernels/ops/attention/dsv4/tilelang_kernel.py b/python/sglang/kernels/ops/attention/dsv4/tilelang_kernel.py deleted file mode 100644 index f94c97146..000000000 --- a/python/sglang/kernels/ops/attention/dsv4/tilelang_kernel.py +++ /dev/null @@ -1,123 +0,0 @@ -import functools -from typing import Any - -import tilelang -import tilelang.language as T -import torch - -from sglang.srt.utils import is_hip - -if is_hip(): - FP8 = "float8_e5m2fnuz" - FP8_ = torch.float8_e5m2 -else: - FP8 = "float8_e4m3" - FP8_ = torch.float8_e4m3fn -FP32 = "float32" -INT32 = "int32" - - -@functools.cache -def fp8_paged_mqa_logits_kernel( - head_dim: int = 128, - num_heads: int = 64, - block_size: int = 64, - clear_accum: bool = True, -) -> Any: - N = T.symbolic("batch_size") - L = T.symbolic("max_table_length") - S = T.symbolic("max_seq_len") - C = T.symbolic("num_blocks") - B = block_size - D = head_dim - H = num_heads - d_0, d_1 = T.dynamic("d_0, d_1") - - assert D % 4 == 0 - assert H % 4 == 0 - assert D == 128 - - @tilelang.jit - def fp8_paged_mqa_logits( - q: T.Tensor[(N, H, D), FP8], - kvcache: T.StridedTensor[(C, B, D), (d_0, D, 1), FP8], - kvcache_scale: T.StridedTensor[(C, B), (d_1, 1), FP32], - weight: T.Tensor[(N, H), FP32], - seq_lens: T.Tensor[(N,), INT32], - page_table: T.Tensor[(N, L), INT32], - o: T.Tensor[(N, S), FP32], - ) -> None: - _ = N, L, S, C, D, H, B, d_0, d_1 - with T.Kernel(N) as bx: - seq_len = seq_lens[bx] - q_smem = T.alloc_shared((H, D), FP8) - q_s_frag = T.alloc_fragment((H,), FP32) - T.copy(q[bx, 0, 0], q_smem) - T.copy(weight[bx, 0], q_s_frag) - - for i in T.Pipelined(T.ceildiv(seq_len, B), num_stages=2): - page = page_table[bx, i] - k_smem = T.alloc_shared((B, D), FP8) - k_s_frag = T.alloc_fragment((B,), FP32) - T.copy(kvcache[page, 0, 0], k_smem) - T.copy(kvcache_scale[page, 0], k_s_frag) - - logits = T.alloc_fragment((B, H), FP32) - if not clear_accum: - T.fill(logits, 0.0) - T.gemm( - k_smem, - q_smem, - logits, - transpose_A=False, - transpose_B=True, - clear_accum=clear_accum, - ) - - for h, j in T.Parallel(H, B): - logits[j, h] = T.max(logits[j, h], 0.0) * q_s_frag[h] - logits_sum = T.alloc_fragment((B,), FP32) - T.reduce_sum(logits, logits_sum, dim=1) - for j in T.Parallel(B): - logits_sum[j] *= k_s_frag[j] - T.copy(logits_sum, o[bx, i * B]) - - return fp8_paged_mqa_logits - - -def tilelang_fp8_paged_mqa_logits( - q_fp8: torch.Tensor, - kvcache_fp8: torch.Tensor, - weight: torch.Tensor, - seq_lens: torch.Tensor, - page_table: torch.Tensor, - deep_gemm_metadata: Any, - max_seq_len: int, - clean_logits: bool = True, -) -> torch.Tensor: - _ = deep_gemm_metadata - batch_size, _, num_heads, head_dim = q_fp8.shape - block_size = kvcache_fp8.shape[1] - assert head_dim == 128, "TODO" - assert block_size == 64, "TODO" - assert q_fp8.shape == (batch_size, 1, num_heads, head_dim) - assert kvcache_fp8.shape[1:] == (block_size, 1, head_dim + 4) - assert weight.shape == (batch_size, num_heads) - assert seq_lens.shape == (batch_size,) - assert page_table.shape[0] == batch_size - assert clean_logits == False - - logits = page_table.new_empty((batch_size, max_seq_len), dtype=torch.float32) - kernel = fp8_paged_mqa_logits_kernel( - head_dim=head_dim, - num_heads=num_heads, - block_size=block_size, - clear_accum=clean_logits, - ) - q_fp8 = q_fp8.view(batch_size, num_heads, head_dim) - kvcache_fp8 = kvcache_fp8.view(-1, block_size * (head_dim + 4)) - kvcache = kvcache_fp8[..., : block_size * head_dim].view(dtype=FP8_) - kvcache = kvcache.view(-1, block_size, head_dim) - kvcache_scale = kvcache_fp8[..., block_size * head_dim :].view(dtype=torch.float32) - kernel(q_fp8, kvcache, kvcache_scale, weight, seq_lens, page_table, logits) - return logits diff --git a/python/sglang/kernels/ops/attention/fla/chunk_scaled_dot_kkt.py b/python/sglang/kernels/ops/attention/fla/chunk_scaled_dot_kkt.py deleted file mode 100644 index f401ac6f5..000000000 --- a/python/sglang/kernels/ops/attention/fla/chunk_scaled_dot_kkt.py +++ /dev/null @@ -1,147 +0,0 @@ -# Adapted from https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/common/chunk_scaled_dot_kkt.py -# -*- coding: utf-8 -*- -# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang - -from typing import Optional - -import torch -import triton -import triton.language as tl - -from sglang.kernels.ops.attention.fla.index import prepare_chunk_indices -from sglang.kernels.ops.attention.fla.op import safe_exp - - -# @triton.autotune( -# configs=[ -# triton.Config({"BK": BK}, num_warps=num_warps, num_stages=num_stages) -# for BK in [32, 64, 128] -# for num_warps in [2, 4, 8] -# for num_stages in [2, 3, 4] -# ], -# key=["H", "K", "BT", "IS_VARLEN"], -# ) -@triton.jit(do_not_specialize=["T"]) -def chunk_scaled_dot_kkt_fwd_kernel( - k, - beta, - g_cumsum, - A, - cu_seqlens, - chunk_indices, - T, - H: tl.constexpr, - Hg: tl.constexpr, - K: tl.constexpr, - BT: tl.constexpr, - BK: tl.constexpr, - IS_VARLEN: tl.constexpr, - USE_G: tl.constexpr, -): - i_t, i_bh = tl.program_id(0), tl.program_id(1) - i_b, i_h = i_bh // H, i_bh % H - if IS_VARLEN: - i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load( - chunk_indices + i_t * 2 + 1 - ).to(tl.int32) - bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load( - cu_seqlens + i_n + 1 - ).to(tl.int32) - T = eos - bos - else: - bos, eos = i_b * T, i_b * T + T - o_t = tl.arange(0, BT) - - p_beta = tl.make_block_ptr( - beta + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,) - ) - b_beta = tl.load(p_beta, boundary_check=(0,)) - - b_A = tl.zeros([BT, BT], dtype=tl.float32) - for i_k in range(tl.cdiv(K, BK)): - p_k = tl.make_block_ptr( - k + (bos * Hg + i_h // (H // Hg)) * K, - (T, K), - (Hg * K, 1), - (i_t * BT, i_k * BK), - (BT, BK), - (1, 0), - ) - b_k = tl.load(p_k, boundary_check=(0, 1)) - b_A += tl.dot(b_k, tl.trans(b_k)) - - if USE_G: - p_g = tl.make_block_ptr( - g_cumsum + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,) - ) - b_g = tl.load(p_g, boundary_check=(0,)) - b_g_diff = b_g[:, None] - b_g[None, :] - b_A = b_A * safe_exp(b_g_diff) - - b_A *= b_beta[:, None] - b_A = tl.where(o_t[:, None] > o_t[None, :], b_A, 0) - p_A = tl.make_block_ptr( - A + (bos * H + i_h) * BT, (T, BT), (BT * H, 1), (i_t * BT, 0), (BT, BT), (1, 0) - ) - tl.store(p_A, b_A.to(p_A.dtype.element_ty), boundary_check=(0, 1)) - - -def chunk_scaled_dot_kkt_fwd( - k: torch.Tensor, - beta: torch.Tensor, - g_cumsum: Optional[torch.Tensor] = None, - cu_seqlens: Optional[torch.LongTensor] = None, - chunk_size: int = 64, - output_dtype: torch.dtype = torch.float32, -) -> torch.Tensor: - r""" - Compute beta * K * K^T. - - Args: - k (torch.Tensor): - The key tensor of shape `[B, T, H, K]`. - beta (torch.Tensor): - The beta tensor of shape `[B, T, H]`. - g_cumsum (torch.Tensor): - The cumulative sum of the gate tensor of shape `[B, T, H]`. - Default: None - cu_seqlens (torch.LongTensor): - The cumulative sequence lengths of the input tensor. - Default: None - chunk_size (int): - The chunk size. Default: 64. - output_dtype (torch.dtype): - The dtype of the output tensor. Default: `torch.float32` - - Returns: - beta * K * K^T of shape `[B, T, H, BT]` where `BT` is the chunk size. - """ - - B, T, Hg, K = k.shape - - H = beta.shape[-1] - BT = chunk_size - chunk_indices = ( - prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None - ) - NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices) - A = torch.empty(B, T, H, BT, device=k.device, dtype=output_dtype) - chunk_scaled_dot_kkt_fwd_kernel[(NT, B * H)]( - k=k, - beta=beta, - g_cumsum=g_cumsum, - A=A, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - T=T, - H=H, - Hg=Hg, - K=K, - BT=BT, - BK=64, - IS_VARLEN=cu_seqlens is not None, - USE_G=g_cumsum is not None, - num_warps=8, - num_stages=3, - ) - return A diff --git a/python/sglang/kernels/ops/attention/fla/solve_tril.py b/python/sglang/kernels/ops/attention/fla/solve_tril.py deleted file mode 100644 index 0640a4019..000000000 --- a/python/sglang/kernels/ops/attention/fla/solve_tril.py +++ /dev/null @@ -1,464 +0,0 @@ -# Adapt from https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/utils/solve_tril.py -# -*- coding: utf-8 -*- -# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang - -from typing import Optional - -import torch -import triton -import triton.language as tl - -from sglang.kernels.ops.attention.fla.index import prepare_chunk_indices -from sglang.kernels.ops.attention.fla.utils import input_guard - - -# @triton.autotune( -# configs=[ -# triton.Config({}, num_warps=num_warps, num_stages=num_stages) -# for num_warps in [1, 2, 4, 8] -# for num_stages in [2, 3, 4, 5] -# ], -# key=["BT"], -# ) -@triton.jit(do_not_specialize=["T"]) -def solve_tril_16x16_kernel( - A, - Ad, - cu_seqlens, - chunk_indices, - T, - H: tl.constexpr, - BT: tl.constexpr, - IS_VARLEN: tl.constexpr, -): - i_t, i_bh = tl.program_id(0), tl.program_id(1) - i_b, i_h = i_bh // H, i_bh % H - if IS_VARLEN: - i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load( - chunk_indices + i_t * 2 + 1 - ).to(tl.int32) - bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load( - cu_seqlens + i_n + 1 - ).to(tl.int32) - T = eos - bos - else: - bos, eos = i_b * T, i_b * T + T - - A = A + (bos * H + i_h) * BT - Ad = Ad + (bos * H + i_h) * 16 - - offset = (i_t * 16) % BT - p_A = tl.make_block_ptr( - A, (T, BT), (H * BT, 1), (i_t * 16, offset), (16, 16), (1, 0) - ) - p_Ai = tl.make_block_ptr(Ad, (T, 16), (H * 16, 1), (i_t * 16, 0), (16, 16), (1, 0)) - b_A = tl.load(p_A, boundary_check=(0, 1)).to(tl.float32) - b_A = -tl.where(tl.arange(0, 16)[:, None] > tl.arange(0, 16)[None, :], b_A, 0) - - o_i = tl.arange(0, 16) - for i in range(1, min(16, T - i_t * 16)): - b_a = -tl.load(A + (i_t * 16 + i) * H * BT + o_i + offset) - b_a = b_a + tl.sum(b_a[:, None] * b_A, 0) - mask = o_i == i - b_A = tl.where(mask[:, None], b_a, b_A) - b_A += o_i[:, None] == o_i[None, :] - tl.store( - p_Ai, - b_A.to(p_Ai.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - - -# @triton.autotune( -# configs=[ -# triton.Config({}, num_warps=num_warps, num_stages=num_stages) -# for num_warps in [1, 2, 4, 8] -# for num_stages in [2, 3, 4, 5] -# ], -# key=["H", "BT", "IS_VARLEN"], -# ) -@triton.jit(do_not_specialize=["T"]) -def merge_16x16_to_32x32_inverse_kernel( - A, - Ad, - Ai, - cu_seqlens, - chunk_indices, - T, - H: tl.constexpr, - BT: tl.constexpr, - IS_VARLEN: tl.constexpr, -): - i_t, i_bh = tl.program_id(0), tl.program_id(1) - i_b, i_h = i_bh // H, i_bh % H - if IS_VARLEN: - i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load( - chunk_indices + i_t * 2 + 1 - ).to(tl.int32) - bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load( - cu_seqlens + i_n + 1 - ).to(tl.int32) - T = eos - bos - else: - bos, eos = i_b * T, i_b * T + T - - A += (bos * H + i_h) * 32 - Ad += (bos * H + i_h) * 16 - Ai += (bos * H + i_h) * 32 - - p_A_21 = tl.make_block_ptr( - A, (T, 32), (H * 32, 1), (i_t * 32 + 16, 0), (16, 16), (1, 0) - ) - p_Ad_11 = tl.make_block_ptr( - Ad, (T, 16), (H * 16, 1), (i_t * 32, 0), (16, 16), (1, 0) - ) - p_Ad_22 = tl.make_block_ptr( - Ad, (T, 16), (H * 16, 1), (i_t * 32 + 16, 0), (16, 16), (1, 0) - ) - p_Ai_11 = tl.make_block_ptr( - Ai, (T, 32), (H * 32, 1), (i_t * 32, 0), (16, 16), (1, 0) - ) - p_Ai_22 = tl.make_block_ptr( - Ai, (T, 32), (H * 32, 1), (i_t * 32 + 16, 16), (16, 16), (1, 0) - ) - p_Ai_21 = tl.make_block_ptr( - Ai, (T, 32), (H * 32, 1), (i_t * 32 + 16, 0), (16, 16), (1, 0) - ) - - A_21 = tl.load(p_A_21, boundary_check=(0, 1)).to(tl.float32) - Ai_11 = tl.load(p_Ad_11, boundary_check=(0, 1)).to(tl.float32) - Ai_22 = tl.load(p_Ad_22, boundary_check=(0, 1)).to(tl.float32) - Ai_21 = -tl.dot( - tl.dot(Ai_22, A_21, input_precision="ieee"), Ai_11, input_precision="ieee" - ) - tl.store( - p_Ai_11, - Ai_11.to(p_Ai_11.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_22, - Ai_22.to(p_Ai_22.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_21, - Ai_21.to(p_Ai_21.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - - -# @triton.autotune( -# configs=[ -# triton.Config({}, num_warps=num_warps, num_stages=num_stages) -# for num_warps in [2, 4, 8] -# for num_stages in [2, 3, 4, 5] -# ], -# key=["H", "BT", "IS_VARLEN"], -# ) -@triton.jit(do_not_specialize=["T"]) -def merge_16x16_to_64x64_inverse_kernel( - A, - Ad, - Ai, - cu_seqlens, - chunk_indices, - T, - H: tl.constexpr, - BT: tl.constexpr, - IS_VARLEN: tl.constexpr, -): - i_t, i_bh = tl.program_id(0), tl.program_id(1) - i_b, i_h = i_bh // H, i_bh % H - if IS_VARLEN: - i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load( - chunk_indices + i_t * 2 + 1 - ).to(tl.int32) - bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load( - cu_seqlens + i_n + 1 - ).to(tl.int32) - T = eos - bos - else: - bos, eos = i_b * T, i_b * T + T - - A += (bos * H + i_h) * 64 - Ad += (bos * H + i_h) * 16 - Ai += (bos * H + i_h) * 64 - - p_A_21 = tl.make_block_ptr( - A, (T, 64), (H * 64, 1), (i_t * 64 + 16, 0), (16, 16), (1, 0) - ) - p_A_32 = tl.make_block_ptr( - A, (T, 64), (H * 64, 1), (i_t * 64 + 32, 16), (16, 16), (1, 0) - ) - p_A_31 = tl.make_block_ptr( - A, (T, 64), (H * 64, 1), (i_t * 64 + 32, 0), (16, 16), (1, 0) - ) - p_A_43 = tl.make_block_ptr( - A, (T, 64), (H * 64, 1), (i_t * 64 + 48, 32), (16, 16), (1, 0) - ) - p_A_42 = tl.make_block_ptr( - A, (T, 64), (H * 64, 1), (i_t * 64 + 48, 16), (16, 16), (1, 0) - ) - p_A_41 = tl.make_block_ptr( - A, (T, 64), (H * 64, 1), (i_t * 64 + 48, 0), (16, 16), (1, 0) - ) - p_Ad_11 = tl.make_block_ptr( - Ad, (T, 16), (H * 16, 1), (i_t * 64, 0), (16, 16), (1, 0) - ) - p_Ad_22 = tl.make_block_ptr( - Ad, (T, 16), (H * 16, 1), (i_t * 64 + 16, 0), (16, 16), (1, 0) - ) - p_Ad_33 = tl.make_block_ptr( - Ad, (T, 16), (H * 16, 1), (i_t * 64 + 32, 0), (16, 16), (1, 0) - ) - p_Ad_44 = tl.make_block_ptr( - Ad, (T, 16), (H * 16, 1), (i_t * 64 + 48, 0), (16, 16), (1, 0) - ) - - A_21 = tl.load(p_A_21, boundary_check=(0, 1)).to(tl.float32) - A_32 = tl.load(p_A_32, boundary_check=(0, 1)).to(tl.float32) - A_31 = tl.load(p_A_31, boundary_check=(0, 1)).to(tl.float32) - A_43 = tl.load(p_A_43, boundary_check=(0, 1)).to(tl.float32) - A_42 = tl.load(p_A_42, boundary_check=(0, 1)).to(tl.float32) - A_41 = tl.load(p_A_41, boundary_check=(0, 1)).to(tl.float32) - - Ai_11 = tl.load(p_Ad_11, boundary_check=(0, 1)).to(tl.float32) - Ai_22 = tl.load(p_Ad_22, boundary_check=(0, 1)).to(tl.float32) - Ai_33 = tl.load(p_Ad_33, boundary_check=(0, 1)).to(tl.float32) - Ai_44 = tl.load(p_Ad_44, boundary_check=(0, 1)).to(tl.float32) - - Ai_21 = -tl.dot( - tl.dot(Ai_22, A_21, input_precision="ieee"), Ai_11, input_precision="ieee" - ) - Ai_32 = -tl.dot( - tl.dot(Ai_33, A_32, input_precision="ieee"), Ai_22, input_precision="ieee" - ) - Ai_43 = -tl.dot( - tl.dot(Ai_44, A_43, input_precision="ieee"), Ai_33, input_precision="ieee" - ) - - Ai_31 = -tl.dot( - Ai_33, - tl.dot(A_31, Ai_11, input_precision="ieee") - + tl.dot(A_32, Ai_21, input_precision="ieee"), - input_precision="ieee", - ) - Ai_42 = -tl.dot( - Ai_44, - tl.dot(A_42, Ai_22, input_precision="ieee") - + tl.dot(A_43, Ai_32, input_precision="ieee"), - input_precision="ieee", - ) - Ai_41 = -tl.dot( - Ai_44, - tl.dot(A_41, Ai_11, input_precision="ieee") - + tl.dot(A_42, Ai_21, input_precision="ieee") - + tl.dot(A_43, Ai_31, input_precision="ieee"), - input_precision="ieee", - ) - - p_Ai_11 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64, 0), (16, 16), (1, 0) - ) - p_Ai_22 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64 + 16, 16), (16, 16), (1, 0) - ) - p_Ai_33 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64 + 32, 32), (16, 16), (1, 0) - ) - p_Ai_44 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64 + 48, 48), (16, 16), (1, 0) - ) - p_Ai_21 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64 + 16, 0), (16, 16), (1, 0) - ) - p_Ai_31 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64 + 32, 0), (16, 16), (1, 0) - ) - p_Ai_32 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64 + 32, 16), (16, 16), (1, 0) - ) - p_Ai_41 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64 + 48, 0), (16, 16), (1, 0) - ) - p_Ai_42 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64 + 48, 16), (16, 16), (1, 0) - ) - p_Ai_43 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64 + 48, 32), (16, 16), (1, 0) - ) - tl.store( - p_Ai_11, - Ai_11.to(p_Ai_11.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_22, - Ai_22.to(p_Ai_22.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_33, - Ai_33.to(p_Ai_33.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_44, - Ai_44.to(p_Ai_44.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_21, - Ai_21.to(p_Ai_21.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_31, - Ai_31.to(p_Ai_31.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_32, - Ai_32.to(p_Ai_32.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_41, - Ai_41.to(p_Ai_41.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_42, - Ai_42.to(p_Ai_42.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_43, - Ai_43.to(p_Ai_43.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - - fill_zeros = tl.zeros((16, 16), dtype=tl.float32) - p_Ai_12 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64, 16), (16, 16), (1, 0) - ) - p_Ai_13 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64, 32), (16, 16), (1, 0) - ) - p_Ai_14 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64, 48), (16, 16), (1, 0) - ) - p_Ai_23 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64 + 16, 32), (16, 16), (1, 0) - ) - p_Ai_24 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64 + 16, 48), (16, 16), (1, 0) - ) - p_Ai_34 = tl.make_block_ptr( - Ai, (T, 64), (H * 64, 1), (i_t * 64 + 32, 48), (16, 16), (1, 0) - ) - tl.store( - p_Ai_12, - fill_zeros.to(p_Ai_12.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_13, - fill_zeros.to(p_Ai_13.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_14, - fill_zeros.to(p_Ai_14.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_23, - fill_zeros.to(p_Ai_23.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_24, - fill_zeros.to(p_Ai_24.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - tl.store( - p_Ai_34, - fill_zeros.to(p_Ai_34.dtype.element_ty, fp_downcast_rounding="rtne"), - boundary_check=(0, 1), - ) - - -@input_guard -def solve_tril( - A: torch.Tensor, - cu_seqlens: Optional[torch.Tensor] = None, - output_dtype: torch.dtype = torch.float, -) -> torch.Tensor: - """ - Compute the inverse of the lower triangular matrix - A should be strictly lower triangular, i.e., A.triu() == 0. - - Args: - A (torch.Tensor): - [B, T, H, K] - cu_seqlens (torch.Tensor): - The cumulative sequence lengths of the input tensor. - Default: None. - output_dtype (torch.dtype): - The dtype of the output tensor. Default: `torch.float` - - Returns: - (I + A)^-1 with the same shape as A - """ - assert A.shape[-1] in [16, 32, 64] - - B, T, H, BT = A.shape - Ad = torch.empty( - B, T, H, 16, device=A.device, dtype=torch.float if BT != 16 else output_dtype - ) - - chunk_indices = ( - prepare_chunk_indices(cu_seqlens, 16) if cu_seqlens is not None else None - ) - NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, 16) - solve_tril_16x16_kernel[NT, B * H]( - A=A, - Ad=Ad, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - T=T, - H=H, - BT=BT, - IS_VARLEN=cu_seqlens is not None, - num_warps=1, - num_stages=4, - ) - if BT == 16: - return Ad - - Ai = torch.empty(B, T, H, BT, device=A.device, dtype=output_dtype) - merge_fn = ( - merge_16x16_to_32x32_inverse_kernel - if BT == 32 - else merge_16x16_to_64x64_inverse_kernel - ) - chunk_indices = ( - prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None - ) - NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, BT) - merge_fn[NT, B * H]( - A=A, - Ad=Ad, - Ai=Ai, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - T=T, - H=H, - BT=BT, - IS_VARLEN=cu_seqlens is not None, - num_warps=4, - num_stages=3, - ) - return Ai diff --git a/python/sglang/kernels/ops/attention/flash_attn/cute/barrier.py b/python/sglang/kernels/ops/attention/flash_attn/cute/barrier.py deleted file mode 100644 index 13cc1cb78..000000000 --- a/python/sglang/kernels/ops/attention/flash_attn/cute/barrier.py +++ /dev/null @@ -1,76 +0,0 @@ -import cutlass -import cutlass.cute as cute -from cutlass import Int32 -from cutlass._mlir.dialects import llvm -from cutlass.cutlass_dsl import T, dsl_user_op - - -@dsl_user_op -def ld_acquire(lock_ptr: cute.Pointer, *, loc=None, ip=None) -> cutlass.Int32: - lock_ptr_i64 = lock_ptr.toint(loc=loc, ip=ip).ir_value() - state = llvm.inline_asm( - T.i32(), - [lock_ptr_i64], - "ld.global.acquire.gpu.b32 $0, [$1];", - "=r,l", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - ) - return cutlass.Int32(state) - - -@dsl_user_op -def red_relaxed( - lock_ptr: cute.Pointer, val: cutlass.Constexpr[Int32], *, loc=None, ip=None -) -> None: - lock_ptr_i64 = lock_ptr.toint(loc=loc, ip=ip).ir_value() - llvm.inline_asm( - None, - [lock_ptr_i64, Int32(val).ir_value(loc=loc, ip=ip)], - "red.relaxed.gpu.global.add.s32 [$0], $1;", - "l,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - ) - - -@dsl_user_op -def red_release( - lock_ptr: cute.Pointer, val: cutlass.Constexpr[Int32], *, loc=None, ip=None -) -> None: - lock_ptr_i64 = lock_ptr.toint(loc=loc, ip=ip).ir_value() - llvm.inline_asm( - None, - [lock_ptr_i64, Int32(val).ir_value(loc=loc, ip=ip)], - "red.release.gpu.global.add.s32 [$0], $1;", - "l,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - ) - - -@cute.jit -def wait_eq( - lock_ptr: cute.Pointer, thread_idx: int | Int32, flag_offset: int, val: Int32 -) -> None: - flag_ptr = lock_ptr + flag_offset - if thread_idx == 0: - read_val = Int32(0) - while read_val != val: - read_val = ld_acquire(flag_ptr) - - -@cute.jit -def arrive_inc( - lock_ptr: cute.Pointer, - thread_idx: int | Int32, - flag_offset: int, - val: cutlass.Constexpr[Int32], -) -> None: - flag_ptr = lock_ptr + flag_offset - if thread_idx == 0: - red_release(flag_ptr, val) - # red_relaxed(flag_ptr, val) diff --git a/python/sglang/kernels/ops/attention/flash_attn/cute/compute_block_sparsity.py b/python/sglang/kernels/ops/attention/flash_attn/cute/compute_block_sparsity.py deleted file mode 100644 index 27dcac6e5..000000000 --- a/python/sglang/kernels/ops/attention/flash_attn/cute/compute_block_sparsity.py +++ /dev/null @@ -1,591 +0,0 @@ -from functools import partial -from typing import Callable, Optional, Tuple - -import cutlass -import cutlass.cute as cute -import torch -from cutlass import Boolean, Int8, Int32, const_expr - -from sglang.kernels.ops.attention.flash_attn.cute.block_sparse_utils import ( - get_curr_blocksparse_tensors, -) -from sglang.kernels.ops.attention.flash_attn.cute.block_sparsity import ( - BlockSparseTensors, - BlockSparseTensorsTorch, - to_cute_block_sparse_tensors, -) -from sglang.kernels.ops.attention.flash_attn.cute.cute_dsl_utils import ( - get_aux_tensor_metadata, - to_cute_aux_tensor, - to_cute_tensor, -) -from sglang.kernels.ops.attention.flash_attn.cute.mask import call_mask_mod -from sglang.kernels.ops.attention.flash_attn.cute.seqlen_info import SeqlenInfoQK -from sglang.kernels.ops.attention.flash_attn.cute.testing import is_fake_mode -from sglang.kernels.ops.attention.flash_attn.cute.utils import ( - AuxData, - get_batch_from_cu_tensor, - hash_callable, - scalar_to_ssa, - ssa_to_scalar, -) - - -class BlockSparsityKernel: - """Block sparsity kernel for FlexAttention. - - This kernel computes `mask_mod` for every token of each block - to determine if an n block is full, masked, or neither. - - Writes block counts and indices to a BlockSparseTensors object. - - When use_fast_sampling=True, uses 5-point sampling (4 corners + center) - which is much faster but only suitable for masks where this is sufficient. - - TODO: - - optimize mask_mod evaluation - - transposed tensors for bwd pass - """ - - def __init__( - self, - mask_mod: Callable, - tile_mn: Tuple[int, int], - compute_full_blocks: bool = True, - use_aux_tensors: bool = False, - use_fast_sampling: bool = False, - ): - self.mask_mod = mask_mod - self.tile_mn = tile_mn - self.compute_full_blocks = compute_full_blocks - self.use_aux_tensors = use_aux_tensors - self.use_fast_sampling = use_fast_sampling - - @cute.jit - def __call__( - self, - blocksparse_tensors: BlockSparseTensors, - seqlen_q: Int32, - seqlen_k: Int32, - mCuSeqlensQ: Optional[cute.Tensor] = None, - mCuSeqlensK: Optional[cute.Tensor] = None, - mSeqUsedQ: Optional[cute.Tensor] = None, - mSeqUsedK: Optional[cute.Tensor] = None, - aux_data: AuxData = AuxData(), - ): - ( - mask_cnt, - mask_idx, - full_cnt, - full_idx, - mCuTotalMBlocks, - mCuBlockIdxOffsets, - *_, - ) = blocksparse_tensors - - self.is_varlen_q = const_expr(mCuSeqlensQ is not None) - - if const_expr(self.compute_full_blocks): - assert ( - full_cnt is not None and full_idx is not None - ), "full block tensors must be provided when computing full blocks" - if const_expr(not self.is_varlen_q): - batch_size, num_heads, num_m_blocks, _ = mask_idx.shape - total_m_blocks = batch_size * num_m_blocks - else: - assert const_expr( - mCuTotalMBlocks is not None - ), "mCuTotalMBlocks must be provided when varlen q" - num_heads, total_m_blocks = mask_cnt.shape # num_m_blocks is total_m_blocks - batch_size = mCuSeqlensQ.shape[0] - 1 - - if const_expr(self.use_fast_sampling): - num_threads = 5 - self.num_warps = 1 - else: - num_threads = self.tile_mn[0] - self.num_warps = (num_threads + 32 - 1) // 32 - - if const_expr(not self.is_varlen_q): - grid = [num_m_blocks, num_heads, batch_size] - else: - grid = [total_m_blocks, num_heads, 1] - - self.kernel( - blocksparse_tensors, - seqlen_q, - seqlen_k, - batch_size, - mCuSeqlensQ, - mCuSeqlensK, - mSeqUsedQ, - mSeqUsedK, - mCuTotalMBlocks, - mCuBlockIdxOffsets, - aux_data, - ).launch(grid=grid, block=[num_threads, 1, 1]) - - @cute.kernel - def kernel( - self, - blocksparse_tensors: BlockSparseTensors, - seqlen_q: Int32, - seqlen_k: Int32, - batch_size: Int32, - mCuSeqlensQ: Optional[cute.Tensor] = None, - mCuSeqlensK: Optional[cute.Tensor] = None, - mSeqUsedQ: Optional[cute.Tensor] = None, - mSeqUsedK: Optional[cute.Tensor] = None, - mCuTotalMBlocks: Optional[cute.Tensor] = None, - mCuBlockIdxOffsets: Optional[cute.Tensor] = None, - aux_data: AuxData = AuxData(), - ): - tidx, _, _ = cute.arch.thread_idx() - warp_idx = cute.arch.warp_idx() - lane_id = cute.arch.lane_idx() - - ssa = partial(scalar_to_ssa, dtype=Int32) - - @cute.struct - class SharedStorage: - reduction_buffer_smem: cute.struct.Align[ - cute.struct.MemRange[cutlass.Int8, 2 * self.num_warps], 1024 - ] - - smem = cutlass.utils.SmemAllocator() - storage = smem.allocate(SharedStorage, 16) - - reduction_buffer = storage.reduction_buffer_smem.get_tensor( - cute.make_layout((self.num_warps, 2)) - ) - SeqlenInfoCls = partial( - SeqlenInfoQK.create, - seqlen_q_static=seqlen_q, - seqlen_k_static=seqlen_k, - mCuSeqlensQ=mCuSeqlensQ, - mCuSeqlensK=mCuSeqlensK, - mSeqUsedQ=mSeqUsedQ, - mSeqUsedK=mSeqUsedK, - mCuTotalMBlocks=mCuTotalMBlocks, - mCuBlockIdxOffsets=mCuBlockIdxOffsets, - tile_m=self.tile_mn[0], - tile_n=self.tile_mn[1], - ) - - if const_expr(not self.is_varlen_q): - m_block, head_idx, batch_idx = cute.arch.block_idx() - else: - global_m_block, head_idx, _ = cute.arch.block_idx() - batch_idx = get_batch_from_cu_tensor(global_m_block, mCuTotalMBlocks) - m_block = global_m_block - mCuTotalMBlocks[batch_idx] - - seqlen = SeqlenInfoCls(batch_idx) - seqlen_q = seqlen.seqlen_q - seqlen_k = seqlen.seqlen_k - global_m_block = seqlen.m_block_offset + m_block - - num_n_blocks = (seqlen_k + self.tile_mn[1] - 1) // self.tile_mn[1] - - _, curr_mask_idx, _, curr_full_idx = get_curr_blocksparse_tensors( - batch_idx, head_idx, m_block, blocksparse_tensors, seqlen - ) - - num_mask_blocks = Int32(0) - num_full_blocks = Int32(0) - - m_base = m_block * self.tile_mn[0] - if const_expr(self.use_fast_sampling): - # Loop-invariant per-thread q_idx for the 5 sample points - # (tidx 0, 1: top corners; 2, 3: bottom corners; 4: center). - q_idx_sample = m_base - if tidx == 2 or tidx == 3: - q_idx_sample = cutlass.min(m_base + self.tile_mn[0] - 1, seqlen_q - 1) - elif tidx == 4: - q_idx_sample = ( - m_base + cutlass.min(seqlen_q - m_base, self.tile_mn[0]) // 2 - ) - else: - q_idx_thread = m_base + tidx - thread_in_bounds = Boolean( - tidx < self.tile_mn[0] and q_idx_thread < seqlen_q - ) - - for n_block in cutlass.range(num_n_blocks): - n_base = n_block * self.tile_mn[1] - - if const_expr(self.use_fast_sampling): - # 5-point sampling (4 corners + center). Interior n_blocks - # (n_base + tile_n <= seqlen_k) skip the OOB clamp on the right / - # center samples. - is_interior = (n_base + self.tile_mn[1]) <= seqlen_k - n_right = Int32(0) - n_mid = Int32(0) - if is_interior: - n_right = n_base + self.tile_mn[1] - 1 - n_mid = n_base + self.tile_mn[1] // 2 - else: - n_right = cutlass.min(n_base + self.tile_mn[1] - 1, seqlen_k - 1) - n_mid = ( - n_base + cutlass.min(seqlen_k - n_base, self.tile_mn[1]) // 2 - ) - - kv_idx = n_base - if tidx == 1 or tidx == 3: - kv_idx = n_right - elif tidx == 4: - kv_idx = n_mid - - thread_result = Boolean(False) - thread_is_valid = Boolean(False) - if tidx < 5: - thread_is_valid = Boolean(True) - thread_result = ssa_to_scalar( - call_mask_mod( - self.mask_mod, - ssa(batch_idx), - ssa(head_idx), - ssa(q_idx_sample), - ssa(kv_idx), - seqlen, - aux_data, - ) - ) - - has_unmasked = cute.arch.vote_any_sync(thread_result & thread_is_valid) - has_masked = cute.arch.vote_any_sync( - Boolean(not thread_result) & thread_is_valid - ) - - else: - # Full path. Interior blocks (n_base + tile_n <= seqlen_k) drop the - # per-element bound check; the boundary block (at most one) keeps it. - thread_has_unmasked = Boolean(False) - thread_has_masked = Boolean(False) - kv_idx = Int32(0) - is_interior = (n_base + self.tile_mn[1]) <= seqlen_k - - if is_interior: - if thread_in_bounds: - for c in cutlass.range(self.tile_mn[1], unroll_full=True): - mask_val = ssa_to_scalar( - call_mask_mod( - self.mask_mod, - ssa(batch_idx), - ssa(head_idx), - ssa(q_idx_thread), - ssa(n_base + c), - seqlen, - aux_data, - ) - ) - thread_has_unmasked |= Boolean(mask_val) - thread_has_masked |= Boolean(not mask_val) - else: - if thread_in_bounds: - for c in cutlass.range(self.tile_mn[1], unroll_full=True): - kv_idx = n_base + c - if kv_idx < seqlen_k: - mask_val = ssa_to_scalar( - call_mask_mod( - self.mask_mod, - ssa(batch_idx), - ssa(head_idx), - ssa(q_idx_thread), - ssa(kv_idx), - seqlen, - aux_data, - ) - ) - thread_has_unmasked |= Boolean(mask_val) - thread_has_masked |= Boolean(not mask_val) - - warp_unmasked = cute.arch.vote_any_sync( - thread_has_unmasked & thread_in_bounds - ) - warp_masked = cute.arch.vote_any_sync( - thread_has_masked & thread_in_bounds - ) - if lane_id == 0: - reduction_buffer[warp_idx, 0] = ( - Int8(1) if warp_unmasked else Int8(0) - ) - reduction_buffer[warp_idx, 1] = Int8(1) if warp_masked else Int8(0) - cute.arch.sync_threads() - - # Cross-warp OR via warp 0; thread 0 (lane 0 of warp 0) holds the result. - has_unmasked = Boolean(False) - has_masked = Boolean(False) - if warp_idx == 0: - lane_unmasked = Boolean(False) - lane_masked = Boolean(False) - if lane_id < self.num_warps: - lane_unmasked = reduction_buffer[lane_id, 0] != Int8(0) - lane_masked = reduction_buffer[lane_id, 1] != Int8(0) - has_unmasked = cute.arch.vote_any_sync(lane_unmasked) - has_masked = cute.arch.vote_any_sync(lane_masked) - - # Only thread 0 updates the output arrays (common to both paths) - if tidx == 0: - # Block classification based on what we found: - # - If has_masked and has_unmasked: partial block (needs masking) - # - If only has_unmasked: full block (no masking needed) - # - If only has_masked: skip this block entirely - is_partial = Boolean(has_masked and has_unmasked) - is_full = Boolean(has_unmasked and (not has_masked)) - - if is_partial: - curr_mask_idx[num_mask_blocks] = n_block - num_mask_blocks += 1 - elif is_full and const_expr(self.compute_full_blocks): - curr_full_idx[num_full_blocks] = n_block - num_full_blocks += 1 - - # Only thread 0 writes back the counts - if tidx == 0: - mask_cnt, _, full_cnt, *_ = blocksparse_tensors - if const_expr(self.is_varlen_q): - mask_cnt[head_idx, global_m_block] = num_mask_blocks - if const_expr(self.compute_full_blocks): - full_cnt[head_idx, global_m_block] = num_full_blocks - else: - mask_cnt[batch_idx, head_idx, m_block] = num_mask_blocks - if const_expr(self.compute_full_blocks): - full_cnt[batch_idx, head_idx, m_block] = num_full_blocks - - -def compute_block_sparsity( - tile_m, - tile_n, - batch_size, - num_heads, - seqlen_q, - seqlen_k, - mask_mod: Callable, - aux_tensors: Optional[list], - device, - aux_scalars: Optional[tuple] = None, - cu_seqlens_q: Optional[torch.Tensor] = None, - cu_seqlens_k: Optional[torch.Tensor] = None, - seqused_q: Optional[torch.Tensor] = None, - seqused_k: Optional[torch.Tensor] = None, - cu_total_m_blocks: Optional[torch.Tensor] = None, - cu_block_idx_offsets: Optional[torch.Tensor] = None, - compute_full_blocks: bool = True, - use_fast_sampling: bool = False, -) -> BlockSparseTensorsTorch: - """ - Computes block sparsity for a given `mask_mod`. - - Args: - tile_m: The tile size for the m dimension. - tile_n: The tile size for the n dimension. - batch_size: The batch size. - num_heads: The number of heads. - seqlen_q: The sequence length for the query. - seqlen_k: The sequence length for the key. - mask_mod: The `mask_mod` callable to use. - aux_tensors: A list of auxiliary tensors. - device: The device to use. - cu_seqlens_q: Cumulative q sequence lengths for varlen - cu_seqlens_k: Cumulative k sequence lengths for varlen - seqused_q: Per-batch effective q sequence lengths - seqused_k: Per-batch effective k sequence lengths - cu_total_m_blocks: Cumulative total m blocks tensor for varlen q - cu_block_idx_offsets: Cumulative offsets into the packed mask_block_idx / - full_block_idx tensors per batch (== cumsum of M_b * N_b). - compute_full_blocks: Whether to compute full blocks. If False, only partially-masked blocks are computed. - use_fast_sampling: Whether to use 5-point sampling (4 corners + center). This is much faster, but only suitable for masks where this check is sufficient. - - Returns: - BlockSparseTensorsTorch - """ - aux_scalars = tuple(aux_scalars) if aux_scalars else None - - # Check if mask_mod is marked as suitable for 5-point sampling - use_fast_sampling = getattr(mask_mod, "use_fast_sampling", use_fast_sampling) - - num_m_blocks = (seqlen_q + tile_m - 1) // tile_m - num_n_blocks = (seqlen_k + tile_n - 1) // tile_n - - if cu_seqlens_q is not None: - assert ( - cu_total_m_blocks is not None - ), "total m blocks must be provided when varlen q" - total_m_blocks = cu_total_m_blocks[-1].item() - if cu_block_idx_offsets is None and ( - cu_seqlens_k is not None or seqused_k is not None - ): - # Derive cu_block_idx_offsets from per-batch K seqlens. - cu_block_idx_offsets_list = [0] - for batch_idx in range(batch_size): - batch_seqlen_q = ( - cu_seqlens_q[batch_idx + 1].item() - cu_seqlens_q[batch_idx].item() - ) - if cu_seqlens_k is not None: - batch_seqlen_k = ( - cu_seqlens_k[batch_idx + 1].item() - - cu_seqlens_k[batch_idx].item() - ) - else: - batch_seqlen_k = seqused_k[batch_idx].item() - num_m_blocks_batch = (batch_seqlen_q + tile_m - 1) // tile_m - num_n_blocks_batch = (batch_seqlen_k + tile_n - 1) // tile_n - cu_block_idx_offsets_list.append( - cu_block_idx_offsets_list[-1] - + num_m_blocks_batch * num_n_blocks_batch - ) - cu_block_idx_offsets = torch.tensor( - cu_block_idx_offsets_list, dtype=torch.int32, device=device - ) - if cu_block_idx_offsets is not None: - total_n_blocks = cu_block_idx_offsets[-1].item() - else: - # Uniform-K varlen-Q: every batch has the same K seqlen. - total_n_blocks = total_m_blocks * num_n_blocks - - mask_block_cnt = torch.zeros( - (num_heads, total_m_blocks), device=device, dtype=torch.int32 - ) - mask_block_idx = torch.zeros( - (num_heads, total_n_blocks), device=device, dtype=torch.int32 - ) - full_block_cnt = ( - torch.zeros((num_heads, total_m_blocks), device=device, dtype=torch.int32) - if compute_full_blocks - else None - ) - full_block_idx = ( - torch.zeros((num_heads, total_n_blocks), device=device, dtype=torch.int32) - if compute_full_blocks - else None - ) - else: - total_m_blocks = batch_size * num_m_blocks - total_n_blocks = batch_size * num_m_blocks * num_n_blocks - - mask_block_cnt = torch.zeros( - (batch_size, num_heads, num_m_blocks), device=device, dtype=torch.int32 - ) - mask_block_idx = torch.zeros( - (batch_size, num_heads, num_m_blocks, num_n_blocks), - device=device, - dtype=torch.int32, - ) - full_block_cnt = ( - torch.zeros( - (batch_size, num_heads, num_m_blocks), device=device, dtype=torch.int32 - ) - if compute_full_blocks - else None - ) - full_block_idx = ( - torch.zeros( - (batch_size, num_heads, num_m_blocks, num_n_blocks), - device=device, - dtype=torch.int32, - ) - if compute_full_blocks - else None - ) - - blocksparse_tensors_torch = BlockSparseTensorsTorch( - mask_block_cnt=mask_block_cnt, - mask_block_idx=mask_block_idx, - full_block_cnt=full_block_cnt, - full_block_idx=full_block_idx, - cu_total_m_blocks=cu_total_m_blocks, - cu_block_idx_offsets=cu_block_idx_offsets, - block_size=(tile_m, tile_n), - ) - - mask_mod_hash = hash_callable(mask_mod) - if aux_tensors is not None: - aux_tensor_metadata = get_aux_tensor_metadata(aux_tensors) - else: - aux_tensor_metadata = None - aux_scalar_metadata = ( - tuple(type(s) for s in aux_scalars) if aux_scalars is not None else None - ) - - compile_key = ( - tile_m, - tile_n, - mask_mod_hash, - aux_tensor_metadata, - aux_scalar_metadata, - compute_full_blocks, - cu_seqlens_q is None, - cu_seqlens_k is None, - seqused_q is None, - seqused_k is None, - aux_tensors is not None, - use_fast_sampling, - ) - if compile_key not in compute_block_sparsity.compile_cache: - ( - cu_seqlens_q_tensor, - cu_seqlens_k_tensor, - seqused_q_tensor, - seqused_k_tensor, - ) = [ - to_cute_tensor(t, assumed_align=4, leading_dim=0) if t is not None else None - for t in ( - cu_seqlens_q, - cu_seqlens_k, - seqused_q, - seqused_k, - ) - ] - blocksparse_tensors = to_cute_block_sparse_tensors( - blocksparse_tensors_torch, enable_tvm_ffi=True - ) - if aux_tensors is not None: - cute_aux_tensors = [to_cute_aux_tensor(buf) for buf in aux_tensors] - else: - cute_aux_tensors = None - kernel = BlockSparsityKernel( - mask_mod, - tile_mn=(tile_m, tile_n), - compute_full_blocks=compute_full_blocks, - use_aux_tensors=aux_tensors is not None, - use_fast_sampling=use_fast_sampling, - ) - - compute_block_sparsity.compile_cache[compile_key] = cute.compile( - kernel, - blocksparse_tensors, - seqlen_q, - seqlen_k, - cu_seqlens_q_tensor, - cu_seqlens_k_tensor, - seqused_q_tensor, - seqused_k_tensor, - AuxData(cute_aux_tensors, aux_scalars), - options="--enable-tvm-ffi", - ) - - if not is_fake_mode(): - compute_block_sparsity.compile_cache[compile_key]( - ( - blocksparse_tensors_torch.mask_block_cnt, - blocksparse_tensors_torch.mask_block_idx, - blocksparse_tensors_torch.full_block_cnt, - blocksparse_tensors_torch.full_block_idx, - blocksparse_tensors_torch.cu_total_m_blocks, - blocksparse_tensors_torch.cu_block_idx_offsets, - blocksparse_tensors_torch.dq_write_order, - blocksparse_tensors_torch.dq_write_order_full, - ), - seqlen_q, - seqlen_k, - cu_seqlens_q, - cu_seqlens_k, - seqused_q, - seqused_k, - AuxData(aux_tensors, aux_scalars), - ) - - return blocksparse_tensors_torch - - -compute_block_sparsity.compile_cache = {} diff --git a/python/sglang/kernels/ops/attention/fused_qk_norm.py b/python/sglang/kernels/ops/attention/fused_qk_norm.py deleted file mode 100644 index ce4bc0e42..000000000 --- a/python/sglang/kernels/ops/attention/fused_qk_norm.py +++ /dev/null @@ -1,157 +0,0 @@ -"""Fused Q/K RMSNorm in a single Triton kernel launch. - -Ported from ATOM (atom/model_ops/layernorm.py). Fuses per-head Q RMSNorm -(optionally weightless) and KV RMSNorm into one kernel, halving the number -of norm kernel launches per attention layer. -""" - -from typing import Optional, Tuple - -import torch -import triton -import triton.language as tl - - -@triton.jit -def _fused_qk_norm_kernel( - q_ptr, - k_ptr, - q_out_ptr, - k_out_ptr, - q_weight_ptr, - k_weight_ptr, - eps, - num_tokens, - head_dim, - q_in_stride0, - k_in_stride0, - q_out_stride0, - k_out_stride0, - num_q_heads, - num_k_heads, - Q_HAS_WEIGHT: tl.constexpr, - RBLOCK: tl.constexpr, - XBLOCK: tl.constexpr, -): - num_q_rows = num_tokens * num_q_heads - total_rows = num_tokens * (num_q_heads + num_k_heads) - - xoffset = tl.program_id(0) * XBLOCK - xindex = xoffset + tl.arange(0, XBLOCK)[:, None] - xmask = xindex < total_rows - cols = tl.arange(0, RBLOCK)[None, :] - col_mask = cols < head_dim - - is_q = xindex < num_q_rows - row_in_section = tl.where(is_q, xindex, xindex - num_q_rows) - cur_num_heads = tl.where(is_q, num_q_heads, num_k_heads) - - tokens = row_in_section // cur_num_heads - heads = row_in_section % cur_num_heads - - in_stride = tl.where(is_q, q_in_stride0, k_in_stride0) - in_bases = tokens * in_stride + heads * head_dim - - out_stride0 = tl.where(is_q, q_out_stride0, k_out_stride0) - out_bases = tokens * out_stride0 + heads * head_dim - - mask = xmask & col_mask - - if Q_HAS_WEIGHT: - qw = tl.load( - q_weight_ptr + cols, mask=col_mask, other=0.0, eviction_policy="evict_last" - ).to(tl.float32) - else: - qw = tl.full((RBLOCK,), 1.0, tl.float32) - kw = tl.load( - k_weight_ptr + cols, mask=col_mask, other=0.0, eviction_policy="evict_last" - ).to(tl.float32) - w = tl.where(is_q, qw, kw) - - x = tl.load( - q_ptr + in_bases + cols, - mask=mask & is_q, - other=0.0, - eviction_policy="evict_first", - ).to(tl.float32) - x = x + tl.load( - k_ptr + in_bases + cols, - mask=mask & ~is_q, - other=0.0, - eviction_policy="evict_first", - ).to(tl.float32) - - var = tl.sum(x * x, 1)[:, None] - rstd = tl.rsqrt(var / head_dim + eps) - - out = (x * rstd * w).to(q_out_ptr.dtype.element_ty) - tl.store( - q_out_ptr + out_bases + cols, - out, - mask=mask & is_q, - eviction_policy="evict_first", - ) - tl.store( - k_out_ptr + out_bases + cols, - out, - mask=mask & ~is_q, - eviction_policy="evict_first", - ) - - -def fused_qk_norm( - q: torch.Tensor, - k: torch.Tensor, - q_weight: Optional[torch.Tensor], - k_weight: torch.Tensor, - eps: float, -) -> Tuple[torch.Tensor, torch.Tensor]: - """Fused Q/K RMSNorm in a single Triton kernel launch. - - Args: - q: [num_tokens, num_heads, head_dim] - k: [num_tokens, num_kv_heads, head_dim] - q_weight: [head_dim] norm weight, or None for weightless Q norm - k_weight: [head_dim] norm weight (always required) - eps: epsilon for numerical stability - - Returns: - (q_normed, k_normed) same shapes as inputs - """ - head_dim = k_weight.shape[0] - if q_weight is not None: - assert q_weight.shape[0] == head_dim - num_tokens = q.shape[0] - num_q_heads = q.shape[1] - num_k_heads = k.shape[1] - total_rows = num_tokens * (num_q_heads + num_k_heads) - RBLOCK = triton.next_power_of_2(head_dim) - - q_out = torch.empty_like(q) - k_out = torch.empty_like(k) - - XBLOCK = 2 if total_rows > 8192 else 1 - NUM_WARPS = 1 - q_weight_arg = q_weight if q_weight is not None else k_weight - _fused_qk_norm_kernel[((total_rows + XBLOCK - 1) // XBLOCK,)]( - q, - k, - q_out, - k_out, - q_weight_arg, - k_weight, - eps, - num_tokens, - head_dim, - q.stride(0), - k.stride(0), - q_out.stride(0), - k_out.stride(0), - num_q_heads, - num_k_heads, - Q_HAS_WEIGHT=q_weight is not None, - RBLOCK=RBLOCK, - XBLOCK=XBLOCK, - num_warps=NUM_WARPS, - ) - return q_out, k_out diff --git a/python/sglang/kernels/ops/attention/utils.py b/python/sglang/kernels/ops/attention/utils.py index ce11e8ed1..55fe89086 100644 --- a/python/sglang/kernels/ops/attention/utils.py +++ b/python/sglang/kernels/ops/attention/utils.py @@ -196,192 +196,6 @@ def concat_mla_absorb_q_general(q_nope, q_rope): return torch.cat([q_nope, q_rope], dim=-1) -@triton.jit -def reshape_and_cache_flash( - key_ptr, - value_ptr, - key_cache_ptr, - value_cache_ptr, - slot_mapping_ptr, - swa_slot_mapping_ptr, - k_scale_ptr, - v_scale_ptr, - block_stride, - key_stride, - value_stride, - num_heads, - head_size, - block_size, - HEAD_BLOCK: tl.constexpr, - BLOCK_D: tl.constexpr, - HAS_SWA: tl.constexpr, - USE_SCALE: tl.constexpr, -): - """ - Triton kernel for reshaping per-token K/V tensors into paged KV cache layout. - - Source layout: - key/value: [num_tokens, num_heads, head_size] - - Target cache layout: - cache: [num_blocks, block_size, num_heads, head_size] - - Each Triton program instance handles: - - one token (program_id(0)) - - one block of heads (program_id(1)) - - Features: - - optional SWA slot remapping - - optional FP8 scale dequantization before cache write - - Args: - key_ptr: Pointer to source key tensor. - value_ptr: Pointer to source value tensor. - key_cache_ptr: Pointer to destination key cache tensor. - value_cache_ptr: Pointer to destination value cache tensor. - slot_mapping_ptr: Maps token -> cache slot. - swa_slot_mapping_ptr: Optional second-stage slot remap for SWA mode. - k_scale_ptr: Optional key scaling factor pointer. - v_scale_ptr: Optional value scaling factor pointer. - block_stride: Stride between cache blocks. - key_stride: Stride between source key tokens. - value_stride: Stride between source value tokens. - num_heads: Number of attention heads. - head_size: Hidden dimension per head. - block_size: Number of slots per cache block. - HEAD_BLOCK: Number of heads processed per program. - BLOCK_D: Vectorized dimension size (power-of-2 padded). - HAS_SWA: Enable SWA remapping. - USE_SCALE: Enable scale division before storing. - """ - - # ---------------------------------- - # program ids - # pid0 = token - # pid1 = head block - # ---------------------------------- - token_idx = tl.program_id(0) - head_block_idx = tl.program_id(1) - - # ---------------------------------- - # slot mapping - # ---------------------------------- - slot_idx = tl.load(slot_mapping_ptr + token_idx) - - if HAS_SWA: - slot_idx = tl.load(swa_slot_mapping_ptr + slot_idx) - - if slot_idx < 0: - return - - block_idx = slot_idx // block_size - block_offset = slot_idx % block_size - - # ---------------------------------- - # head range - # ---------------------------------- - head_idx = head_block_idx * HEAD_BLOCK + tl.arange(0, HEAD_BLOCK) - - head_mask = head_idx < num_heads - - dim_idx = tl.arange(0, BLOCK_D) - - # shape = [HEAD_BLOCK, BLOCK_D] - offs = head_idx[:, None] * head_size + dim_idx[None, :] - - mask = head_mask[:, None] & (dim_idx[None, :] < head_size) - - # ---------------------------------- - # source load - # ---------------------------------- - src_key = token_idx * key_stride + offs - src_value = token_idx * value_stride + offs - - k = tl.load(key_ptr + src_key, mask=mask) - v = tl.load(value_ptr + src_value, mask=mask) - - # ---------------------------------- - # optional scale - # ---------------------------------- - if USE_SCALE: - k_scale = tl.load(k_scale_ptr) - v_scale = tl.load(v_scale_ptr) - - k = k / k_scale - v = v / v_scale - - # ---------------------------------- - # target layout - # [block_idx, block_offset, head, dim] - # ---------------------------------- - tgt = block_idx * block_stride + block_offset * num_heads * head_size + offs - - tl.store(key_cache_ptr + tgt, k, mask=mask) - tl.store(value_cache_ptr + tgt, v, mask=mask) - - -def launch_reshape_and_cache_flash( - key, - value, - key_cache, - value_cache, - slot_mapping, - swa_slot_mapping=None, - k_scale=None, - v_scale=None, -): - """ - Launch wrapper for reshape_and_cache_flash Triton kernel. - - This wrapper prepares launch configuration and dispatches the Triton kernel - that writes token-major K/V tensors into paged KV cache layout. - - Args: - key: Source key tensor [num_tokens, num_heads, head_size] - value: Source value tensor [num_tokens, num_heads, head_size] - key_cache: Destination key cache [num_blocks, block_size, num_heads, head_size] - value_cache: Destination value cache [num_blocks, block_size, num_heads, head_size] - slot_mapping: Token-to-cache slot mapping - swa_slot_mapping: Optional SWA remapping table - k_scale: Optional key scaling factor - v_scale: Optional value scaling factor - """ - - num_tokens = key.shape[0] - num_heads = key.shape[1] - head_size = key.shape[2] - - HEAD_BLOCK = 4 - - BLOCK_D = triton.next_power_of_2(head_size) - - grid = ( - num_tokens, - triton.cdiv(num_heads, HEAD_BLOCK), - ) - - reshape_and_cache_flash[grid]( - key, - value, - key_cache, - value_cache, - slot_mapping, - swa_slot_mapping, - k_scale if k_scale is not None else key, - v_scale if v_scale is not None else key, - key_cache.stride(0), - key.stride(0), - value.stride(0), - num_heads, - head_size, - key_cache.shape[1], - HEAD_BLOCK=HEAD_BLOCK, - BLOCK_D=BLOCK_D, - HAS_SWA=(swa_slot_mapping is not None), - USE_SCALE=(k_scale is not None), - ) - - @triton.jit def reshape_and_cache_shuffle_5d( key_ptr, @@ -646,739 +460,6 @@ def launch_gather_shuffle_5d_to_linear( return key_out, value_out -@triton.jit -def _get_gptj_rotated_x( - x, - x_rotated_mask, - BLOCK_D: tl.constexpr, - BLOCK_D_HALF: tl.constexpr, -): - # GPT-J rotary layout: - # Pair adjacent dimensions and apply: - # [x0, x1, x2, x3] -> [-x1, x0, -x3, x2] - - # Apply sign inversion on odd positions. - x_rotated = tl.where(x_rotated_mask, x, -x) - # Reshape into (D/2, 2) pairs. - x_rotated = tl.reshape(x_rotated, (BLOCK_D_HALF, 2)) - # Swap each pair. - x_rotated = tl.flip(x_rotated, 1) - # Flatten back to original shape. - x_rotated = tl.reshape(x_rotated, (BLOCK_D,)) - return x_rotated - - -@triton.jit -def _get_neox_rotated_x( - x, - x_rotated_mask, - BLOCK_D: tl.constexpr, - BLOCK_D_HALF: tl.constexpr, -): - # GPT-NeoX rotary layout: - # Split head dimension into two halves: - # [x0, x1, x2, x3] -> [-x2, -x3, x0, x1] - - # Keep first half positive, second half negative. - x_rotated = tl.where(x_rotated_mask, x, -x) - # Reshape into (2, D/2). - x_rotated = tl.reshape(x_rotated, (2, BLOCK_D_HALF)) - # Reverse each half. - x_rotated = tl.flip(x_rotated, 1) - # Flatten and reverse full vector. - x_rotated = tl.reshape(x_rotated, (BLOCK_D,)) - x_rotated = tl.flip(x_rotated, 0) - return x_rotated - - -@triton.jit -def _unit_rope( - x_ptrs, - cos, - sin, - d_pe_offs, - IS_NEOX: tl.constexpr, - BLOCK_D_pe: tl.constexpr, - BLOCK_D_HALF_pe: tl.constexpr, -): - # Load one full attention head vector. - x_pe = tl.load(x_ptrs) - - # Stage 1: Build rotated vector according to rotary layout. - if IS_NEOX: - x_rotated_mask = d_pe_offs < BLOCK_D_HALF_pe - x_pe_rotated = _get_neox_rotated_x( - x_pe, x_rotated_mask, BLOCK_D_pe, BLOCK_D_HALF_pe - ) - else: - x_rotated_mask = d_pe_offs % 2 == 0 - x_pe_rotated = _get_gptj_rotated_x( - x_pe, x_rotated_mask, BLOCK_D_pe, BLOCK_D_HALF_pe - ) - - # Stage 2: Apply RoPE transform: - # x' = x*cos + rotate(x)*sin - x_pe = x_pe * cos + x_pe_rotated * sin - - return x_pe - - -@triton.jit -def _load_cos_sin( - cos_sin_ptr, - pos, - d_cos_offs, - stride_t, - stride_d, - freq_dim, -): - base = pos * stride_t - cos = tl.load(cos_sin_ptr + base + d_cos_offs * stride_d) - sin = tl.load(cos_sin_ptr + base + (d_cos_offs + freq_dim) * stride_d) - return cos, sin - - -@triton.jit -def _fused_qk_rope_reshape_and_cache_kernel( - q_ptr, - k_ptr, - v_ptr, - pos_ptr, - cos_sin_ptr, - offs_ptr, - key_cache_ptr, - value_cache_ptr, - slot_mapping_ptr, - swa_slot_mapping_ptr, - q_out_ptr, - k_out_ptr, - zeros_out_ptr, - T, - T_slot, - q_stride_t, - q_stride_h, - q_stride_d, - k_stride_t, - k_stride_h, - k_stride_d, - v_stride_t, - v_stride_h, - v_stride_d, - cos_sin_stride_t, - cos_sin_stride_d, - q_out_stride_t, - q_out_stride_h, - q_out_stride_d, - k_out_stride_t, - k_out_stride_h, - k_out_stride_d, - key_cache_stride_t, - key_cache_stride_h, - key_cache_stride_d, - key_cache_stride_b, - key_cache_stride_x, - value_cache_stride_t, - value_cache_stride_h, - value_cache_stride_d, - value_cache_stride_b, - value_cache_stride_slot_chunk, - value_cache_stride_x, - zeros_out_stride_t, - zeros_out_stride_h, - zeros_out_stride_d, - k_scale_ptr, - v_scale_ptr, - QH_PER_KH: tl.constexpr, - QH: tl.constexpr, - KH: tl.constexpr, - REUSE_FREQS_FRONT_PART: tl.constexpr, - IS_NEOX: tl.constexpr, - BLOCK_D_pe: tl.constexpr, - BLOCK_D_HALF_pe: tl.constexpr, - BLOCK_SIZE: tl.constexpr, - X_SIZE: tl.constexpr, - FLASH_LAYOUT: tl.constexpr, - VALUE_SHUFFLE_LAYOUT: tl.constexpr = False, - HAVE_POS: tl.constexpr = False, - HAVE_K_SCALE: tl.constexpr = False, - HAVE_V_SCALE: tl.constexpr = False, - HAVE_ZEROS: tl.constexpr = False, - HAS_SWA: tl.constexpr = False, -): - # ============================================================ - # Stage 0: Static stride assumptions for Triton compiler - # - # These assumptions help Triton optimize pointer arithmetic and - # simplify generated address calculations. - # ============================================================ - - tl.assume(q_stride_t >= 0) - tl.assume(q_stride_h >= 0) - tl.assume(q_stride_d >= 0) - tl.assume(k_stride_t >= 0) - tl.assume(k_stride_h >= 0) - tl.assume(k_stride_d >= 0) - tl.assume(v_stride_t >= 0) - tl.assume(v_stride_h >= 0) - tl.assume(v_stride_d >= 0) - tl.assume(cos_sin_stride_t >= 0) - tl.assume(cos_sin_stride_d >= 0) - tl.assume(q_out_stride_t >= 0) - tl.assume(q_out_stride_h >= 0) - tl.assume(q_out_stride_d >= 0) - tl.assume(k_out_stride_t >= 0) - tl.assume(k_out_stride_h >= 0) - tl.assume(k_out_stride_d >= 0) - tl.assume(key_cache_stride_t >= 0) - tl.assume(key_cache_stride_h >= 0) - tl.assume(key_cache_stride_d >= 0) - tl.assume(key_cache_stride_b >= 0) - tl.assume(key_cache_stride_x >= 0) - tl.assume(value_cache_stride_t >= 0) - tl.assume(value_cache_stride_h >= 0) - tl.assume(value_cache_stride_d >= 0) - tl.assume(value_cache_stride_b >= 0) - tl.assume(value_cache_stride_slot_chunk >= 0) - tl.assume(value_cache_stride_x >= 0) - tl.assume(zeros_out_stride_t >= 0) - tl.assume(zeros_out_stride_h >= 0) - tl.assume(zeros_out_stride_d >= 0) - - # ============================================================ - # Stage 1: Program instance mapping - # - # Each program handles: - # - one (token, q_head) for Q path - # - selected KV ownership for cache write path - # - # pid layout: - # [0, T*QH) -> decode Q path - # [T*QH, extra KV) -> KV-only path - # ============================================================ - - pid = tl.program_id(0) - tl.assume(pid >= 0) - - d_pe_offs = tl.arange(0, BLOCK_D_pe).to(tl.int64) - - # ============================================================ - # Stage 2: Main decode path (Q always active) - # ============================================================ - - if pid < T * QH: - pid_t = pid // QH - pid_hq = pid % QH - - # -------------------------------------------------------- - # Stage 2.1: Compute rotary frequency offsets - # - # RoPE frequencies may be stored as: - # D/2 frequencies (shared front-half) - # D frequencies (full explicit) - # -------------------------------------------------------- - - if REUSE_FREQS_FRONT_PART: - if IS_NEOX: - d_cos_offs = d_pe_offs - d_cos_offs = tl.where( - (d_cos_offs >= BLOCK_D_HALF_pe) & (d_cos_offs < BLOCK_D_pe), - d_cos_offs - BLOCK_D_HALF_pe, - d_cos_offs, - ).to(d_cos_offs.dtype) - # d_cos_mask = d_cos_offs < BLOCK_D_pe - else: - d_cos_offs = d_pe_offs // 2 - # d_cos_mask = d_cos_offs < BLOCK_D_HALF_pe - else: - d_cos_offs = d_pe_offs - # d_cos_mask = d_cos_offs < BLOCK_D_pe - - # -------------------------------------------------------- - # Stage 2.2: Load token position and optional offset - # - # offs_ptr is used by chunked prefill / sliding-window decode. - # -------------------------------------------------------- - pos = tl.load(pos_ptr + pid_t) - if HAVE_POS: - offset = tl.load(offs_ptr + pid_t) - pos = pos + offset - - # -------------------------------------------------------- - # Stage 2.3: Load cosine / sine table - # -------------------------------------------------------- - # cos_offs = pos * cos_stride_t + d_cos_offs * cos_stride_d - # cos = tl.load(cos_ptr + cos_offs) - # sin = tl.load(sin_ptr + cos_offs) - - freq_dim = BLOCK_D_HALF_pe if REUSE_FREQS_FRONT_PART else BLOCK_D_pe - - cos, sin = _load_cos_sin( - cos_sin_ptr, - pos, - d_cos_offs, - cos_sin_stride_t, - cos_sin_stride_d, - freq_dim, - ) - - # -------------------------------------------------------- - # Stage 2.4: Apply RoPE to Q - # -------------------------------------------------------- - q_ptrs = ( - q_ptr + pid_t * q_stride_t + pid_hq * q_stride_h + d_pe_offs * q_stride_d - ) - q_pe = _unit_rope( - q_ptrs, - cos, - sin, - d_pe_offs, - IS_NEOX, - BLOCK_D_pe, - BLOCK_D_HALF_pe, - ) - - # Store rotated Q output. - q_out_ptrs = ( - q_out_ptr - + pid_t * q_out_stride_t - + pid_hq * q_out_stride_h - + d_pe_offs * q_out_stride_d - ) - tl.store(q_out_ptrs, q_pe.to(q_out_ptr.dtype.element_ty)) - - if HAVE_ZEROS: - z = tl.zeros((BLOCK_D_pe,), dtype=zeros_out_ptr.dtype.element_ty) - zeros_out_ptrs = ( - zeros_out_ptr - + pid_t * zeros_out_stride_t - + pid_hq * zeros_out_stride_h - + d_pe_offs * zeros_out_stride_d - ) - tl.store(zeros_out_ptrs, z) - - # ======================================================== - # Stage 3: KV ownership path - # - # Only one Q group leader writes KV: - # pid_hq % QH_PER_KH == 0 - # - # This prevents duplicated KV cache writes. - # ======================================================== - - if pid_hq % QH_PER_KH == 0: - # ---------------------------------------------------- - # Stage 3.1: Resolve cache slot - # ---------------------------------------------------- - pid_slot = tl.load(slot_mapping_ptr + pid_t).to(tl.int64) - if HAS_SWA: - pid_slot = tl.load(swa_slot_mapping_ptr + pid_slot) - - # ------------------------------------------------ - # Stage 3.2: Apply RoPE to K - # ------------------------------------------------ - if pid_slot >= 0: - pid_t_slot = pid_slot // BLOCK_SIZE - pid_b = pid_slot % BLOCK_SIZE - pid_hk = pid_hq // QH_PER_KH - if HAVE_K_SCALE: - k_scale = tl.load(k_scale_ptr) - else: - k_scale = 1 - k_ptrs = ( - k_ptr - + pid_t * k_stride_t - + pid_hk * k_stride_h - + d_pe_offs * k_stride_d - ) - k_pe = _unit_rope( - k_ptrs, - cos, - sin, - d_pe_offs, - IS_NEOX, - BLOCK_D_pe, - BLOCK_D_HALF_pe, - ) - - k_out_ptrs = ( - k_out_ptr - + pid_t * k_out_stride_t - + pid_hk * k_out_stride_h - + d_pe_offs * k_out_stride_d - ) - tl.store(k_out_ptrs, k_pe.to(k_out_ptr.dtype.element_ty)) - - # ------------------------------------------------ - # Stage 3.3: Optional fp8 scaling before cache - # ------------------------------------------------ - - k_scale_rcprl = 1 / k_scale - k_pe = k_pe * k_scale_rcprl - - # ------------------------------------------------ - # Stage 3.4: Write K cache - # - # Two layouts supported: - # FLASH_LAYOUT - # paged KV layout - # ------------------------------------------------ - - if FLASH_LAYOUT: - k_out_ptrs = ( - key_cache_ptr - + pid_t_slot * key_cache_stride_t - + pid_b * key_cache_stride_b - + pid_hk * key_cache_stride_h - + d_pe_offs * key_cache_stride_d - ) - else: - k_pe = tl.reshape(k_pe, (BLOCK_D_pe // X_SIZE, X_SIZE)) - dx_offs = tl.arange(0, BLOCK_D_pe // X_SIZE).to(tl.int64) - x_offs = tl.arange(0, X_SIZE).to(tl.int64) - k_out_ptrs = ( - key_cache_ptr - + pid_t_slot * key_cache_stride_t - + pid_hk * key_cache_stride_h - + dx_offs[:, None] * key_cache_stride_d - + pid_b * key_cache_stride_b - + x_offs[None, :] * key_cache_stride_x - ) - - tl.store(k_out_ptrs, k_pe.to(key_cache_ptr.dtype.element_ty)) - - # ------------------------------------------------ - # Stage 3.5: Write V cache - # - # Supports: - # normal layout - # shuffle layout - # ------------------------------------------------ - - v_ptrs = ( - v_ptr - + pid_t * v_stride_t - + pid_hk * v_stride_h - + d_pe_offs * v_stride_d - ) - if HAVE_V_SCALE: - v_scale = tl.load(v_scale_ptr) - else: - v_scale = 1 - v_scale_rcprl = 1 / v_scale - v = tl.load(v_ptrs) * v_scale_rcprl - if VALUE_SHUFFLE_LAYOUT: - slot_chunk = pid_b // X_SIZE - x_off = pid_b % X_SIZE - v_out_ptrs = ( - value_cache_ptr - + pid_t_slot * value_cache_stride_t - + pid_hk * value_cache_stride_h - + slot_chunk * value_cache_stride_slot_chunk - + d_pe_offs.to(tl.int64) * value_cache_stride_d - + x_off * value_cache_stride_x - ) - else: - v_out_ptrs = ( - value_cache_ptr - + pid_t_slot * value_cache_stride_t - + pid_hk * value_cache_stride_h - + d_pe_offs.to(tl.int64) * value_cache_stride_d - + pid_b * value_cache_stride_b - ) - tl.store(v_out_ptrs, v.to(value_cache_ptr.dtype.element_ty)) - # ============================================================ - # Stage 4: Extra KV-only path - # - # Handles tokens that only require cache update: - # T_slot > T - # - # No Q / no RoPE on Q branch. - # ============================================================ - else: - pid = pid - T * QH + T * KH - if pid < T_slot * KH: - pid_t = pid // KH - pid_hk = pid % KH - pid_slot = tl.load(slot_mapping_ptr + pid_t).to(tl.int64) - if HAS_SWA: - pid_slot = tl.load(swa_slot_mapping_ptr + pid_slot) - - if pid_slot >= 0: - pid_t_slot = pid_slot // BLOCK_SIZE - pid_b = pid_slot % BLOCK_SIZE - if HAVE_K_SCALE: - k_scale = tl.load(k_scale_ptr) - else: - k_scale = 1 - k_ptrs = ( - k_ptr - + pid_t * k_stride_t - + pid_hk * k_stride_h - + d_pe_offs * k_stride_d - ) - - k_pe = tl.load(k_ptrs) - - k_out_ptrs = ( - k_out_ptr - + pid_t * k_out_stride_t - + pid_hk * k_out_stride_h - + d_pe_offs * k_out_stride_d - ) - tl.store(k_out_ptrs, k_pe.to(k_out_ptr.dtype.element_ty)) - - k_scale_rcprl = 1 / k_scale - k_pe = k_pe * k_scale_rcprl - - if FLASH_LAYOUT: - k_out_ptrs = ( - key_cache_ptr - + pid_t_slot * key_cache_stride_t - + d_pe_offs * key_cache_stride_d - + pid_b * key_cache_stride_b - + pid_hk * key_cache_stride_h - ) - else: - k_pe = tl.reshape(k_pe, (BLOCK_D_pe // X_SIZE, X_SIZE)) - dx_offs = tl.arange(0, BLOCK_D_pe // X_SIZE).to(tl.int64) - x_offs = tl.arange(0, X_SIZE).to(tl.int64) - k_out_ptrs = ( - key_cache_ptr - + pid_t_slot * key_cache_stride_t - + pid_hk * key_cache_stride_h - + dx_offs[:, None] * key_cache_stride_d - + pid_b * key_cache_stride_b - + x_offs[None, :] * key_cache_stride_x - ) - tl.store(k_out_ptrs, k_pe.to(key_cache_ptr.dtype.element_ty)) - - v_ptrs = ( - v_ptr - + pid_t * v_stride_t - + pid_hk * v_stride_h - + d_pe_offs * v_stride_d - ) - if HAVE_V_SCALE: - v_scale = tl.load(v_scale_ptr) - else: - v_scale = 1 - v_scale_rcprl = 1 / v_scale - v = tl.load(v_ptrs) * v_scale_rcprl - if VALUE_SHUFFLE_LAYOUT: - slot_chunk = pid_b // X_SIZE - x_off = pid_b % X_SIZE - v_out_ptrs = ( - value_cache_ptr - + pid_t_slot * value_cache_stride_t - + pid_hk * value_cache_stride_h - + slot_chunk * value_cache_stride_slot_chunk - + d_pe_offs * value_cache_stride_d - + x_off * value_cache_stride_x - ) - else: - v_out_ptrs = ( - value_cache_ptr - + pid_t_slot * value_cache_stride_t - + pid_hk * value_cache_stride_h - + d_pe_offs * value_cache_stride_d - + pid_b * value_cache_stride_b - ) - tl.store(v_out_ptrs, v.to(value_cache_ptr.dtype.element_ty)) - - -def fused_qk_rope_reshape_and_cache( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - key_cache: torch.Tensor, - value_cache: torch.Tensor, - slot_mapping: torch.Tensor, - pos: torch.Tensor, - cos_sin: torch.Tensor, - k_scale: torch.Tensor, - v_scale: torch.Tensor, - is_neox: bool, - flash_layout: bool, - apply_scale: bool = True, - offs: torch.Tensor = None, - q_out: torch.Tensor = None, - k_out: torch.Tensor = None, - output_zeros: bool = True, - zeros_out: torch.Tensor = None, - swa_slot_mapping=None, -): - """ - Perform RoPE on q and k and along the last dimension and copy k and v in to key_cache and value_cache inplace - - Key parameters: - - q: shape (T, QH, D). - - k: shape (T_slot, KH, D). - - v: shape (T_slot, KH, D). - - if flash_layout: - - key_cache: shape (T_cache, block_size, KH, D). - - value_cache: shape (T_cache, block_size, KH, D). - - else: - - key_cache: shape (T_cache, KH, D // x, block_size, x). - - value_cache: shape (T_cache, KH, D, block_size). - - slot_mapping: shape (T_slot, ). - - T is the number of decode tokens, T_cahce * block_size is the max number of tokens of kv_cache - QH must be multiple of KH - - Returns: - - q_out: same shape as input q. - - k_out: same shape as input k. - - key_cache: same shape as input key_cache (inplace). - - value_cache: same shape as input value_cache (inplace). - - zeros_out: same shape as input q. - """ - - t, qh, d = q.shape - tk, kh, dk = k.shape - tv, vh, dv = v.shape - if flash_layout: - t_cache, block_size, kh_cache, dk_cache = key_cache.shape - t_cache_v, block_size_v, vh_cache, dv_cache = value_cache.shape - value_shuffle_layout = False - else: - t_cache, kh_cache, dkx_cache, block_size, x_cache = key_cache.shape - if value_cache.ndim == 5: - # value_cache shuffle: (num_blocks, num_kv_heads, block_size // x, head_size, x) - t_cache_v, vh_cache, slot_chunk_v, dv_cache, x_v = value_cache.shape - value_shuffle_layout = True - block_size_v = slot_chunk_v * x_v - assert block_size_v == block_size and x_v == x_cache, ( - f"value_cache shuffle (T,KH,block_size//x,D,x) must match key: " - f"{block_size_v=} {block_size=} {x_v=} {x_cache=}" - ) - else: - t_cache_v, vh_cache, dv_cache, block_size_v = value_cache.shape - value_shuffle_layout = False - (t_slot,) = slot_mapping.shape - - assert ( - t == tk == tv and t_slot <= tk - ), f"Number of tokens should be identical for q, kand v. The number of tokens of slot_mapping should no more than that of q, k and v, {t=} {tk=} {tv=} {t_slot=}" - assert ( - block_size == block_size_v - ), f"block size should be identical for key_cache, and value_cache {block_size} {block_size_v}" - assert ( - kh == vh == kh_cache == vh_cache - ), "KV head should be identical for k, v, key_cache, and value_cache" - assert ( - t_cache == t_cache_v - ), "Number of tokens should be identical for key_cache, and value_cache" - if flash_layout: - assert ( - d == dk == dv == dk_cache == dv_cache - ), "D dimension should be identical for q, k, and v" - else: - assert ( - d == dk == dv == dkx_cache * x_cache == dv_cache - ), "D dimension should be identical for q, k, and v" - assert x_cache == triton.next_power_of_2(x_cache), "x_size should be power of 2" - - assert d == triton.next_power_of_2(d), "D dimension should be power of 2" - assert block_size == triton.next_power_of_2( - block_size - ), "block_size should be power of 2" - assert qh % kh == 0, "Q heads must be multiple of H heads" - d_freq = cos_sin.shape[-1] // 2 - assert (d_freq == d // 2) or ( - d_freq == d - ), "cos/sin last dim should be the same or half of the qk last dim" - reuse_freqs_front_part = d_freq == d // 2 - - if q_out is None: - q_out = torch.empty((t, qh, d), dtype=q.dtype, device=q.device) - - if k_out is None: - k_out = torch.empty((tk, kh, dk), dtype=k.dtype, device=q.device) - - if zeros_out is not None: - tz, qhz, dz = zeros_out.shape - assert ( - t == tz and qh == qhz and d == dz - ), f"q and zeros shape mismatch {q.shape=} {zeros_out.shape=}" - output_zeros = True - elif output_zeros: - zeros_out = torch.empty((t, qh, d), dtype=q.dtype, device=q.device) - else: - zeros_out = None - - n_pid = t * qh + (t_slot - t) * kh if t_slot >= t else t * qh - grid = (n_pid, 1, 1) - _fused_qk_rope_reshape_and_cache_kernel[grid]( - q, - k, - v, - pos, - cos_sin, - offs, - key_cache, - value_cache, - slot_mapping, - swa_slot_mapping, - q_out, - k_out, - zeros_out, - t, - t_slot, - *q.stride(), - *k.stride(), - *v.stride(), - cos_sin.stride(0), - cos_sin.stride(-1), - *q_out.stride(), - *k_out.stride(), - key_cache.stride(0) if not flash_layout else key_cache.stride(0), - key_cache.stride(1) if not flash_layout else key_cache.stride(2), - key_cache.stride(2) if not flash_layout else key_cache.stride(3), - key_cache.stride(3) if not flash_layout else key_cache.stride(1), - key_cache.stride(4) if not flash_layout else 0, - value_cache.stride(0) if not flash_layout else value_cache.stride(0), - value_cache.stride(1) if not flash_layout else value_cache.stride(2), - ( - value_cache.stride(3) - if (not flash_layout and value_shuffle_layout) - else (value_cache.stride(2) if not flash_layout else value_cache.stride(3)) - ), - ( - 0 - if (not flash_layout and value_shuffle_layout) - else (value_cache.stride(3) if not flash_layout else value_cache.stride(1)) - ), - value_cache.stride(2) if (not flash_layout and value_shuffle_layout) else 0, - value_cache.stride(4) if (not flash_layout and value_shuffle_layout) else 0, - zeros_out.stride(0) if zeros_out is not None else 0, - zeros_out.stride(1) if zeros_out is not None else 0, - zeros_out.stride(2) if zeros_out is not None else 0, - k_scale_ptr=k_scale, - v_scale_ptr=v_scale, - QH_PER_KH=qh // kh, - QH=qh, - KH=kh, - REUSE_FREQS_FRONT_PART=reuse_freqs_front_part, - IS_NEOX=is_neox, - BLOCK_D_pe=d, - BLOCK_D_HALF_pe=d // 2, - BLOCK_SIZE=block_size, - X_SIZE=x_cache if not flash_layout else 0, - FLASH_LAYOUT=flash_layout, - VALUE_SHUFFLE_LAYOUT=value_shuffle_layout, - HAVE_POS=(offs is not None), - HAVE_K_SCALE=(k_scale is not None and apply_scale), - HAVE_V_SCALE=(v_scale is not None and apply_scale), - HAVE_ZEROS=output_zeros, - HAS_SWA=(swa_slot_mapping is not None), - num_warps=1, - ) - - if zeros_out is not None: - return q_out.view(-1, qh * d), k_out, key_cache, value_cache, zeros_out - return q_out.view(-1, qh * d), k_out, key_cache, value_cache - - def assert_buffer_fits(used: int, capacity: int, what: str, **context) -> None: """Safety guard: a preallocated cuda-graph buffer must hold the runtime write. diff --git a/python/sglang/kernels/ops/diffusion/__init__.py b/python/sglang/kernels/ops/diffusion/__init__.py index 4baeb0ce5..55fbecef5 100644 --- a/python/sglang/kernels/ops/diffusion/__init__.py +++ b/python/sglang/kernels/ops/diffusion/__init__.py @@ -114,6 +114,6 @@ register_kernel( KernelSpec( op="diffusion.sparse_linear_attn_fwd", backend=KernelBackend.TRITON, - target="sglang.kernels.ops.diffusion.sparse_linear_attn_kernels:get_block_map", + target="sglang.kernels.ops.diffusion.sparse_linear_attn_kernels:_attn_fwd", ) ) diff --git a/python/sglang/kernels/ops/elementwise/__init__.py b/python/sglang/kernels/ops/elementwise/__init__.py index d895402cc..da0d838fc 100644 --- a/python/sglang/kernels/ops/elementwise/__init__.py +++ b/python/sglang/kernels/ops/elementwise/__init__.py @@ -2,10 +2,10 @@ Home for cross-cutting pointwise kernels that do not belong to a single functional group: the fused-pointwise Triton collection (``elementwise``: -softcap, sigmoid-mul, gated-activation and fused-rmsnorm variants shared -across models) and the ``add_constant`` JIT reference kernel used by the -developer guide. Individual functions register (or are imported) under the -functional op id they logically belong to. +sigmoid-mul, gated-activation and fused-rmsnorm variants shared across models) +and the ``add_constant`` JIT reference kernel used by the developer guide. +Individual functions register (or are imported) under the functional op id +they logically belong to. """ __all__ = [] diff --git a/python/sglang/kernels/ops/elementwise/elementwise.py b/python/sglang/kernels/ops/elementwise/elementwise.py index c7a65cd6f..f0eccba29 100644 --- a/python/sglang/kernels/ops/elementwise/elementwise.py +++ b/python/sglang/kernels/ops/elementwise/elementwise.py @@ -1,38 +1,13 @@ -from typing import Optional, Tuple - import torch import triton import triton.language as tl from sglang.kernels.jit.utils import is_arch_support_pdl -from sglang.kernels.ops.activation.softcap import softcap_out as fused_softcap from sglang.srt.utils import is_hip -from sglang.srt.utils.custom_op import register_custom_op _is_hip = is_hip() -# cast to float + softcap -class Softcap: - def __init__(self, softcap_const: float): - self.softcap_const = softcap_const - - def __call__(self, *args, **kwargs): - return self.forward(*args, **kwargs) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - if x.is_cuda: - return self.forward_cuda(x) - else: - return self.forward_native(x) - - def forward_native(self, x: torch.Tensor) -> torch.Tensor: - return torch.tanh(x.float() / self.softcap_const) * self.softcap_const - - def forward_cuda(self, x: torch.Tensor, autotune=False) -> torch.Tensor: - return fused_softcap(x, self.softcap_const, autotune=autotune) - - rmsnorm_autotune = triton.autotune( configs=[ triton.Config(kwargs={"BLOCK_SIZE": 1024}, num_warps=4, num_stages=1), @@ -213,136 +188,6 @@ def fused_rmsnorm(x, weight, eps, autotune=False, inplace=False): return output -class FusedDualResidualRMSNorm: - """ - Fused implementation of - y = RMSNorm2(RMSNorm1(x) + residual)) - """ - - def __init__(self, rmsnorm1, rmsnorm2) -> None: # the one after rmsnorm1 - self.rmsnorm1 = rmsnorm1 - self.rmsnorm2 = rmsnorm2 - self.variance_epsilon = self.rmsnorm1.variance_epsilon - assert self.rmsnorm1.variance_epsilon == self.rmsnorm2.variance_epsilon - assert self.rmsnorm1.weight.shape == self.rmsnorm2.weight.shape - - def __call__(self, *args, **kwargs): - return self.forward(*args, **kwargs) - - def forward( - self, x: torch.Tensor, residual: torch.Tensor - ) -> Tuple[torch.Tensor, torch.Tensor]: - if x.is_cuda: - return self.forward_cuda(x, residual) - else: - return self.forward_flashinfer(x, residual) - - def forward_cuda( - self, x: torch.Tensor, residual: torch.Tensor, autotune=False - ) -> Tuple[torch.Tensor, torch.Tensor]: - return fused_dual_residual_rmsnorm( - x, - residual, - self.rmsnorm1.weight, - self.rmsnorm2.weight, - self.variance_epsilon, - autotune=autotune, - ) - - def forward_flashinfer( - self, - x: torch.Tensor, - residual: torch.Tensor, - ) -> Tuple[torch.Tensor, torch.Tensor]: - normed1 = self.rmsnorm1(x) - residual = normed1 + residual - return self.rmsnorm2(residual), residual - - def forward_native( - self, - x: torch.Tensor, - residual: torch.Tensor, - ) -> Tuple[torch.Tensor, torch.Tensor]: - normed1 = self.rmsnorm1.forward_native(x) - residual = normed1 + residual - return self.rmsnorm2.forward_native(residual), residual - - -@triton.jit -def experts_combine_kernel( - out_hidden_states, - moe_hidden_states, - mlp_hidden_states, - combine_k: tl.constexpr, - hidden_dim: tl.constexpr, - BLOCK_SIZE: tl.constexpr, -): - pid = tl.program_id(0) - start_index_mlp = pid * hidden_dim - start_index_rmoe = pid * hidden_dim * combine_k - offsets = tl.arange(0, BLOCK_SIZE) - mask = offsets < hidden_dim - combine_k_offsets = tl.arange(0, combine_k) - - moe_x = tl.load( - moe_hidden_states - + start_index_rmoe - + combine_k_offsets[:, None] * hidden_dim - + offsets[None, :], - mask=mask[None, :], - other=0.0, - ) - moe_x = tl.sum(moe_x, axis=0) - mlp_x = tl.load(mlp_hidden_states + start_index_mlp + offsets, mask=mask, other=0.0) - combined_x = (moe_x + mlp_x) / 1.4142135623730951 - - tl.store(out_hidden_states + start_index_mlp + offsets, combined_x, mask=mask) - - -@register_custom_op(out_shape="mlp_hidden_states") -def experts_combine_triton( - moe_hidden_states: torch.Tensor, - mlp_hidden_states: torch.Tensor, - output_buffer: Optional[torch.Tensor] = None, -) -> torch.Tensor: - assert moe_hidden_states.is_contiguous() - assert mlp_hidden_states.is_contiguous() - - if len(moe_hidden_states.shape) == 2: - combine_k = 1 # pre-combined - else: - combine_k = moe_hidden_states.shape[1] - - if output_buffer is None: - out_hidden_states = torch.empty_like(mlp_hidden_states) - else: - flat_output_buffer = output_buffer.view(mlp_hidden_states.dtype).reshape(-1) - assert flat_output_buffer.numel() >= mlp_hidden_states.numel() - out_hidden_states = flat_output_buffer[: mlp_hidden_states.numel()].reshape( - mlp_hidden_states.shape - ) - - bs, hidden_dim = mlp_hidden_states.shape - - config = { - "BLOCK_SIZE": triton.next_power_of_2(hidden_dim), - "num_warps": max( - min(triton.next_power_of_2(triton.cdiv(hidden_dim, 1024)), 8), 4 - ), - } - - experts_combine_kernel[(bs,)]( - out_hidden_states, - moe_hidden_states, - mlp_hidden_states, - combine_k, - hidden_dim, - **config, - ) - - return out_hidden_states - - # gelu on first half of vector @triton.jit def gelu_and_mul_kernel( diff --git a/python/sglang/kernels/ops/kvcache/__init__.py b/python/sglang/kernels/ops/kvcache/__init__.py index b6501baaa..c04e8212c 100644 --- a/python/sglang/kernels/ops/kvcache/__init__.py +++ b/python/sglang/kernels/ops/kvcache/__init__.py @@ -67,7 +67,6 @@ _TRITON_KERNELS = [ ("kv_indices", "get_num_kv_index_blocks_flashmla"), ("kv_indices", "get_num_page_per_block_flashmla"), ("rope_cache", "fused_qk_rope_reshape_and_cache"), - ("trtllm_fp8_kv_kernel", "fused_fp8_set_kv_buffer"), ("trtllm_mha_page_table", "build_trtllm_mha_page_table"), ("trtllm_mha_graph_metadata", "update_trtllm_mha_graph_metadata"), ("aiter_unified_attention", "scatter_ragged_to_page_table_kernel"), diff --git a/python/sglang/kernels/ops/mamba/triton_ops/layernorm_gated.py b/python/sglang/kernels/ops/mamba/triton_ops/layernorm_gated.py deleted file mode 100644 index 524e8d399..000000000 --- a/python/sglang/kernels/ops/mamba/triton_ops/layernorm_gated.py +++ /dev/null @@ -1,172 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -# Copyright (c) 2024, Tri Dao. -# Adapted from https://github.com/state-spaces/mamba/blob/60dadf2e0ee730ac337035d5533de10bc26e4847/mamba_ssm/ops/triton/layernorm_gated.py - -import torch -import triton -import triton.language as tl - - -@triton.heuristics({"HAS_BIAS": lambda args: args["B"] is not None}) -@triton.heuristics({"HAS_Z": lambda args: args["Z"] is not None}) -@triton.jit -def _layer_norm_fwd_1pass_kernel( - X, # pointer to the input - Y, # pointer to the output - W, # pointer to the weights - B, # pointer to the biases - Z, # pointer to the other branch - Mean, # pointer to the mean - Rstd, # pointer to the 1/std - stride_x_row: tl.int64, - stride_y_row: tl.int64, - stride_z_row: tl.int64, - M: tl.int64, # number of rows in X - N: tl.int64, # number of columns in X - eps, # epsilon to avoid division by zero - BLOCK_N: tl.constexpr, - HAS_BIAS: tl.constexpr, - HAS_Z: tl.constexpr, - NORM_BEFORE_GATE: tl.constexpr, - IS_RMS_NORM: tl.constexpr, -): - # Map the program id to the row of X and Y it should compute. - row = tl.program_id(0) - group = tl.program_id(1) - X += row * stride_x_row + group * N - Y += row * stride_y_row + group * N - if HAS_Z: - Z += row * stride_z_row + group * N - if not IS_RMS_NORM: - Mean += group * M - Rstd += group * M - W += group * N - if HAS_BIAS: - B += group * N - # Compute mean and variance - cols = tl.arange(0, BLOCK_N) - x = tl.load(X + cols, mask=cols < N, other=0.0).to(tl.float32) - if HAS_Z and not NORM_BEFORE_GATE: - z = tl.load(Z + cols, mask=cols < N).to(tl.float32) - x *= z * tl.sigmoid(z) - if not IS_RMS_NORM: - mean = tl.sum(x, axis=0) / N - tl.store(Mean + row, mean) - xbar = tl.where(cols < N, x - mean, 0.0) - var = tl.sum(xbar * xbar, axis=0) / N - else: - xbar = tl.where(cols < N, x, 0.0) - var = tl.sum(xbar * xbar, axis=0) / N - rstd = 1 / tl.sqrt(var + eps) - tl.store(Rstd + row, rstd) - # Normalize and apply linear transformation - mask = cols < N - w = tl.load(W + cols, mask=mask).to(tl.float32) - if HAS_BIAS: - b = tl.load(B + cols, mask=mask).to(tl.float32) - x_hat = (x - mean) * rstd if not IS_RMS_NORM else x * rstd - y = x_hat * w + b if HAS_BIAS else x_hat * w - if HAS_Z and NORM_BEFORE_GATE: - z = tl.load(Z + cols, mask=mask).to(tl.float32) - y *= z * tl.sigmoid(z) - # Write output - tl.store(Y + cols, y, mask=mask) - - -def _layer_norm_fwd( - x, - weight, - bias, - eps, - z=None, - out=None, - group_size=None, - norm_before_gate=True, - is_rms_norm=False, -): - M, N = x.shape - if group_size is None: - group_size = N - assert N % group_size == 0 - ngroups = N // group_size - assert x.stride(-1) == 1 - if z is not None: - assert z.stride(-1) == 1 - assert z.shape == (M, N) - assert weight.shape == (N,) - assert weight.stride(-1) == 1 - if bias is not None: - assert bias.stride(-1) == 1 - assert bias.shape == (N,) - # allocate output - if out is not None: - assert out.shape == x.shape - else: - out = torch.empty_like(x) - assert out.stride(-1) == 1 - mean = ( - torch.empty((ngroups * M,), dtype=torch.float32, device=x.device) - if not is_rms_norm - else None - ) - rstd = torch.empty((ngroups * M,), dtype=torch.float32, device=x.device) - # Less than 64KB per feature: enqueue fused kernel - MAX_FUSED_SIZE = 65536 // x.element_size() - BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(group_size)) - if group_size > BLOCK_N: - raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.") - # heuristics for number of warps - num_warps = min(max(BLOCK_N // 256, 1), 8) - grid = (M, ngroups) - with torch.get_device_module(x.device).device(x.device.index): - _layer_norm_fwd_1pass_kernel[grid]( - x, - out, - weight, - bias, - z, - mean, - rstd, - x.stride(0), - out.stride(0), - z.stride(0) if z is not None else 0, - M, - group_size, - eps, - BLOCK_N=BLOCK_N, - NORM_BEFORE_GATE=norm_before_gate, - IS_RMS_NORM=is_rms_norm, - num_warps=num_warps, - ) - return out, mean, rstd - - -def rms_norm_gated( - x, weight, bias, z=None, eps=1e-6, group_size=None, norm_before_gate=True -): - x_shape_og = x.shape - # reshape input data into 2D tensor - x = x.reshape(-1, x.shape[-1]) - if x.stride(-1) != 1: - x = x.contiguous() - if z is not None: - assert z.shape == x_shape_og - z = z.reshape(-1, z.shape[-1]) - if z.stride(-1) != 1: - z = z.contiguous() - weight = weight.contiguous() - if bias is not None: - bias = bias.contiguous() - y, _, _ = _layer_norm_fwd( - x, - weight, - bias, - eps, - z=z, - group_size=group_size, - norm_before_gate=norm_before_gate, - is_rms_norm=True, - ) - - return y.reshape(x_shape_og) diff --git a/python/sglang/kernels/ops/moe/kpool_topk_transform.py b/python/sglang/kernels/ops/moe/kpool_topk_transform.py deleted file mode 100644 index 330293a68..000000000 --- a/python/sglang/kernels/ops/moe/kpool_topk_transform.py +++ /dev/null @@ -1,74 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING, Optional - -import torch - -from sglang.kernels.jit.utils import cache_once, load_jit - -if TYPE_CHECKING: - from tvm_ffi.module import Module - -# Pool-level top-k values that have dedicated, validated kernel instantiations. -SUPPORTED_GROUP_TOPK = (128, 160, 192, 224, 256, 512) - - -@cache_once -def _jit_kpool_topk_transform_module(group_topk: int) -> Module: - """Compile and cache the kpool top-k transform module for a given group_topk.""" - assert group_topk in SUPPORTED_GROUP_TOPK, ( - "fast_kpool_topk_transform supports pool-level topk " - f"{SUPPORTED_GROUP_TOPK}, got {group_topk}" - ) - return load_jit( - f"kpool_topk_transform_{group_topk}", - cuda_files=["dsa/kpool_topk_transform.cuh"], - cuda_wrappers=[("kpool_topk_transform", "KpoolTopKTransformKernel::transform")], - extra_cuda_cflags=[f"-DSGL_GROUP_TOPK={group_topk}"], - ) - - -def fast_kpool_topk_transform_fused( - score: torch.Tensor, - lengths: torch.Tensor, - pool_size: int, - topk: int, - page_table: Optional[torch.Tensor] = None, - topk_indices_offset: Optional[torch.Tensor] = None, - row_starts: Optional[torch.Tensor] = None, - seq_lens: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """ - Pool-level radix top-k for the NSA kpool indexer. - - Selects pool groups from ``score`` at pool granularity, expands each selected - group to ``pool_size`` token indices, and optionally transforms those token - indices through a page table or a ragged offset. - """ - assert topk % pool_size == 0 - group_topk = topk // pool_size - assert group_topk in SUPPORTED_GROUP_TOPK, ( - "fast_kpool_topk_transform supports pool-level topk " - f"{SUPPORTED_GROUP_TOPK}, got {group_topk}" - ) - assert score.dim() == 2 - assert page_table is None or topk_indices_offset is None - if seq_lens is not None: - assert seq_lens.dim() == 1 - assert seq_lens.shape[0] == score.shape[0] - - out_cols = topk + (pool_size - 1 if seq_lens is not None else 0) - dst_token_indices = score.new_empty((score.shape[0], out_cols), dtype=torch.int32) - - module = _jit_kpool_topk_transform_module(group_topk) - module.kpool_topk_transform( - score, - lengths, - dst_token_indices, - pool_size, - page_table, - topk_indices_offset, - row_starts, - seq_lens, - ) - return dst_token_indices diff --git a/python/sglang/kernels/ops/moe/router.py b/python/sglang/kernels/ops/moe/router.py index 5c0b86e58..bce31503b 100644 --- a/python/sglang/kernels/ops/moe/router.py +++ b/python/sglang/kernels/ops/moe/router.py @@ -1,10 +1,9 @@ -from typing import Optional, Tuple +from typing import Optional import torch import triton import triton.language as tl -from sglang.srt.layers.moe.topk import fused_topk from sglang.srt.utils import is_hip _is_hip = is_hip() @@ -177,7 +176,6 @@ def fused_moe_router_tensorcore_kernel( stride_bn: tl.constexpr, dp_attn_workaround_flag: tl.constexpr, ): - # 1. get block id pid = tl.program_id(axis=0) @@ -387,42 +385,3 @@ def fused_moe_router_shim( moe_softcapping=moe_softcapping, correction_bias=correction_bias, ) - - -class FusedMoeRouter: - def __init__(self, router_linear, topk, moe_softcapping) -> None: - self.router_linear = router_linear - self.topk = topk - self.moe_softcapping = moe_softcapping - - def __call__(self, *args, **kwargs): - return self.forward(*args, **kwargs) - - def forward( - self, x: torch.Tensor, residual: torch.Tensor - ) -> Tuple[torch.Tensor, torch.Tensor]: - if x.is_cuda: - return self.forward_cuda(x, residual) - else: - return self.forward_vllm(x, residual) - - def forward_cuda( - self, x: torch.Tensor, autotune=False - ) -> Tuple[torch.Tensor, torch.Tensor]: - return fused_moe_router_shim( - moe_softcapping=self.moe_softcapping, - hidden_states=x, - gating_output=self.router_linear.weight, - topk=self.topk, - renormalize=False, - ) - - def forward_torch( - self, - x: torch.Tensor, - ) -> Tuple[torch.Tensor, torch.Tensor]: - g = x.float() @ self.router_linear.weight.T.float() - - g = torch.tanh(g.float() / self.moe_softcapping) * self.moe_softcapping - - return fused_topk(x, g, self.topk, False) diff --git a/python/sglang/kernels/ops/quantization/mxfp8.py b/python/sglang/kernels/ops/quantization/mxfp8.py deleted file mode 100644 index 8f65ba7db..000000000 --- a/python/sglang/kernels/ops/quantization/mxfp8.py +++ /dev/null @@ -1,136 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -import torch - -from sglang.kernels.jit.utils import ( - cache_once, - load_jit, - make_cpp_args, - override_jit_cuda_arch, -) - -if TYPE_CHECKING: - from tvm_ffi.module import Module - - -def _mxfp8_cuda_flags() -> list[str]: - return [ - "-DNDEBUG", - "-DCUTLASS_ENABLE_TENSOR_CORE_MMA=1", - "-DCUTLASS_VERSIONS_GENERATED", - "-DCUTLASS_DEBUG_TRACE_LEVEL=0", - "--expt-extended-lambda", - ] - - -def _mxfp8_arch_env(): - if not torch.cuda.is_available(): - raise RuntimeError("MXFP8 JIT kernels require CUDA.") - major, minor = torch.cuda.get_device_capability() - if major < 10: - raise RuntimeError( - f"MXFP8 JIT kernels require compute capability >= 10.0, got {major}.{minor}." - ) - # MXFP8 kernels use architecture-family-specific instructions and must be - # compiled for `sm_*a` targets (e.g. sm_100a), not plain sm_100. - # JIT compilation targets only the current device, unlike AOT fat-binaries; - # adding extra architectures here would clash with the single SGL_CUDA_ARCH - # value injected by load_jit(). - return override_jit_cuda_arch(major, minor, suffix="a") - - -@cache_once -def _jit_es_sm100_mxfp8_blockscaled_group_quant(dtype: torch.dtype) -> Module: - args = make_cpp_args(dtype) - with _mxfp8_arch_env(): - return load_jit( - "es_sm100_mxfp8_blockscaled_group_quant", - *args, - cuda_files=[ - "moe/expert_specialization/es_sm100_mxfp8_blockscaled_group_quant.cuh" - ], - cuda_wrappers=[ - ( - "es_sm100_mxfp8_blockscaled_group_quant", - f"EsSm100MXFP8BlockscaledGroupQuant<{args}>::run", - ) - ], - extra_dependencies=["cutlass"], - extra_cuda_cflags=_mxfp8_cuda_flags(), - ) - - -@cache_once -def _jit_es_sm100_mxfp8_blockscaled_moe_group_gemm(dtype: torch.dtype) -> Module: - args = make_cpp_args(dtype) - with _mxfp8_arch_env(): - return load_jit( - "es_sm100_mxfp8_blockscaled_moe_group_gemm", - *args, - cuda_files=[ - "moe/expert_specialization/es_sm100_mxfp8_blockscaled_moe_group_gemm.cuh" - ], - cuda_wrappers=[ - ( - "es_sm100_mxfp8_blockscaled_moe_group_gemm", - f"EsSm100MXFP8BlockscaledMoeGroupGemm<{args}>::run", - ) - ], - extra_dependencies=["cutlass"], - extra_cuda_cflags=_mxfp8_cuda_flags(), - ) - - -def es_sm100_mxfp8_blockscaled_grouped_quant( - input: torch.Tensor, - tokens_per_expert: torch.Tensor, - expert_offsets: torch.Tensor, - blockscale_offsets: torch.Tensor, - quant_output: torch.Tensor, - scale_factor: torch.Tensor, -) -> None: - module = _jit_es_sm100_mxfp8_blockscaled_group_quant(input.dtype) - module.es_sm100_mxfp8_blockscaled_group_quant( - input, - tokens_per_expert, - expert_offsets, - blockscale_offsets, - quant_output, - scale_factor, - ) - - -def es_sm100_mxfp8_blockscaled_moe_grouped_gemm( - a: torch.Tensor, - b: torch.Tensor, - sfa: torch.Tensor, - sfb: torch.Tensor, - expert_offsets: torch.Tensor, - blockscale_offsets: torch.Tensor, - tokens_per_expert: torch.Tensor, - workspace: torch.Tensor, - dtype: torch.dtype, -) -> torch.Tensor: - num_experts, m, tokens = a.shape[0], a.shape[1], b.shape[0] - d = torch.empty((tokens, m), device=a.device, dtype=dtype) - d_ptrs = torch.empty((num_experts,), device=a.device, dtype=torch.int64) - b_ptrs = torch.empty((num_experts,), device=a.device, dtype=torch.int64) - sfb_ptrs = torch.empty((num_experts,), device=a.device, dtype=torch.int64) - module = _jit_es_sm100_mxfp8_blockscaled_moe_group_gemm(dtype) - module.es_sm100_mxfp8_blockscaled_moe_group_gemm( - a, - b, - sfa, - sfb, - expert_offsets, - blockscale_offsets, - tokens_per_expert, - b_ptrs, - sfb_ptrs, - d, - d_ptrs, - workspace, - ) - return d diff --git a/python/sglang/kernels/ops/speculative/resolve_future_token_ids.py b/python/sglang/kernels/ops/speculative/resolve_future_token_ids.py deleted file mode 100644 index 43e786b10..000000000 --- a/python/sglang/kernels/ops/speculative/resolve_future_token_ids.py +++ /dev/null @@ -1,41 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -import torch - -from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args - -if TYPE_CHECKING: - from tvm_ffi.module import Module - - -@cache_once -def _jit_resolve_future_token_ids_module(dtype: torch.dtype) -> Module: - """Compile and cache the JIT module for a given dtype.""" - args = make_cpp_args(dtype) - return load_jit( - "resolve_future_token_ids", - *args, - cuda_files=["elementwise/resolve_future_token_ids.cuh"], - cuda_wrappers=[ - ( - "resolve_future_token_ids", - f"ResolveFutureTokenIds<{args}>::run", - ) - ], - ) - - -def resolve_future_token_ids_cuda( - input_ids: torch.Tensor, future_token_ids_map: torch.Tensor -) -> None: - """Resolve future token IDs in-place on CUDA. - - For each negative value in input_ids, replaces it with - future_token_ids_map[-value]. Non-negative values are unchanged. - - Supported dtypes: torch.int32, torch.int64. - """ - module = _jit_resolve_future_token_ids_module(input_ids.dtype) - module.resolve_future_token_ids(input_ids, future_token_ids_map) diff --git a/test/registered/kernels/benchmark/quantization/bench_mxfp8_moe.py b/test/registered/kernels/benchmark/quantization/bench_mxfp8_moe.py deleted file mode 100644 index 6b75424a2..000000000 --- a/test/registered/kernels/benchmark/quantization/bench_mxfp8_moe.py +++ /dev/null @@ -1,292 +0,0 @@ -from __future__ import annotations - -import sys -from typing import Any - -import torch -import triton - -from sglang.kernels.jit.benchmark.utils import get_benchmark_range, run_benchmark -from sglang.kernels.ops.quantization.mxfp8 import ( - es_sm100_mxfp8_blockscaled_grouped_quant, - es_sm100_mxfp8_blockscaled_moe_grouped_gemm, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci( - est_time=5, stage="base-b-kernel-benchmark", runner_config="1-gpu-large" -) - - -def is_sm100_supported(device=None) -> bool: - if not torch.cuda.is_available(): - return False - return (torch.cuda.get_device_capability(device)[0] == 10) and ( - torch.version.cuda >= "12.8" - ) - - -_SM100_SUPPORTED = is_sm100_supported() - - -def _probe_sgl_kernel_group_mm() -> tuple[bool, str]: - if not _SM100_SUPPORTED: - return False, "MXFP8 MoE benchmark requires sm100+ with CUDA 12.8+." - try: - import sgl_kernel # noqa: F401 - except Exception as e: - return False, f"import sgl_kernel failed: {e}" - if not hasattr(sgl_kernel, "es_sm100_mxfp8_blockscaled_grouped_mm"): - return False, "sgl_kernel.es_sm100_mxfp8_blockscaled_grouped_mm is missing." - try: - pass - - # We assume if it's imported, it works - except Exception as e: - return False, f"calling sgl-kernel grouped_mm op failed: {e}" - return True, "" - - -_SGL_KERNEL_AVAILABLE, _SGL_KERNEL_REASON = _probe_sgl_kernel_group_mm() - - -def align(val: int, alignment: int = 128) -> int: - return int((val + alignment - 1) // alignment * alignment) - - -def _prepare_case( - total_tokens: int, n_g: int, k_g: int, num_experts: int, dtype: torch.dtype -) -> dict[str, Any]: - device = torch.device("cuda") - base = total_tokens // num_experts - rem = total_tokens % num_experts - m_per_expert = [base + (1 if i < rem else 0) for i in range(num_experts)] - - expert_offset = 0 - expert_offsets = [] - aux_expert_offset = 0 - aux_expert_offsets = [] - a_blockscale_offset = 0 - a_blockscale_offsets = [] - b_blockscale_offset = 0 - b_blockscale_offsets = [] - tokens_per_expert_list = [] - expert_ranges = [] - problem_sizes = [] - - a_list = [] - b_list = [] - for g in range(num_experts): - m_g = m_per_expert[g] - tokens_per_expert_list.append(m_g) - expert_ranges.append((expert_offset, expert_offset + m_g)) - expert_offsets.append(expert_offset) - expert_offset += m_g - - aux_expert_offsets.append(aux_expert_offset) - aux_expert_offset += n_g - - a_blockscale_offsets.append(a_blockscale_offset) - a_blockscale_offset += align(m_g, 128) - - b_blockscale_offsets.append(b_blockscale_offset) - b_blockscale_offset += n_g # n_g already align to 128 in practice - - problem_sizes.append([m_g, n_g, k_g]) - - a = torch.randn((m_g, k_g), device=device, dtype=dtype) * 0.1 - b = torch.randn((n_g, k_g), device=device, dtype=dtype) * 0.1 - a_list.append(a) - b_list.append(b) - - a = torch.concat(a_list, dim=0) - b = torch.concat(b_list, dim=0) - - _expert_offsets = torch.tensor(expert_offsets).to(device=device, dtype=torch.int32) - _aux_expert_offsets = torch.tensor(aux_expert_offsets).to( - device=device, dtype=torch.int32 - ) - _a_blockscale_offsets = torch.tensor(a_blockscale_offsets).to( - device=device, dtype=torch.int32 - ) - _b_blockscale_offsets = torch.tensor(b_blockscale_offsets).to( - device=device, dtype=torch.int32 - ) - _tokens_per_expert = torch.tensor(tokens_per_expert_list).to( - device=device, dtype=torch.int32 - ) - _problem_sizes = torch.tensor(problem_sizes).to(device=device, dtype=torch.int32) - - a_quant = torch.zeros_like(a, dtype=torch.float8_e4m3fn, device=device) - a_scale_factor = torch.zeros( - (a_blockscale_offset, k_g // 32), dtype=torch.uint8, device=device - ) - - b_quant = torch.zeros_like(b, dtype=torch.float8_e4m3fn, device=device) - b_scale_factor = torch.zeros( - (num_experts * n_g, k_g // 32), dtype=torch.uint8, device=device - ) - - # Use a global workspace to avoid allocating 1GB every time - workspace = torch.empty((1024, 1024, 1024), dtype=torch.uint8, device=device) - - es_sm100_mxfp8_blockscaled_grouped_quant( - a, - _tokens_per_expert, - _expert_offsets, - _a_blockscale_offsets, - a_quant, - a_scale_factor, - ) - - es_sm100_mxfp8_blockscaled_grouped_quant( - b, - torch.ones_like(_tokens_per_expert) * n_g, - _aux_expert_offsets, - _b_blockscale_offsets, - b_quant, - b_scale_factor, - ) - - b_quant = b_quant.view(num_experts, n_g, k_g) - b_scale_factor = b_scale_factor.view(num_experts, n_g, k_g // 32) - - sgl_b_quant = b_quant.transpose(1, 2) - sgl_b_scale_factor = b_scale_factor.transpose(1, 2) - - return { - "a": a, - "b": b.view(num_experts, n_g, k_g), - "b_quant": b_quant, - "a_quant": a_quant, - "b_scale_factor": b_scale_factor, - "a_scale_factor": a_scale_factor, - "expert_offsets": _expert_offsets, - "a_blockscale_offsets": _a_blockscale_offsets, - "tokens_per_expert": _tokens_per_expert, - "problem_sizes": _problem_sizes, - "sgl_b_quant": sgl_b_quant, - "sgl_b_scale_factor": sgl_b_scale_factor, - "workspace": workspace, - "expert_ranges": expert_ranges, - "dtype": dtype, - } - - -def _sgl_kernel_group_mm(case: dict[str, Any]) -> torch.Tensor: - from sgl_kernel import es_sm100_mxfp8_blockscaled_grouped_mm - - a_quant = case["a_quant"] - sgl_b_quant = case["sgl_b_quant"] - a_scale_factor = case["a_scale_factor"] - sgl_b_scale_factor = case["sgl_b_scale_factor"] - problem_sizes = case["problem_sizes"] - expert_offsets = case["expert_offsets"] - a_blockscale_offsets = case["a_blockscale_offsets"] - dtype = case["dtype"] - - total_tokens = a_quant.shape[0] - n_g = sgl_b_quant.shape[2] - - # sgl-kernel takes output pre-allocated - d = torch.empty((total_tokens, n_g), device=a_quant.device, dtype=dtype) - es_sm100_mxfp8_blockscaled_grouped_mm( - d, - a_quant, - sgl_b_quant, - a_scale_factor, - sgl_b_scale_factor, - problem_sizes, - expert_offsets, - a_blockscale_offsets, - ) - return d - - -shape_range = get_benchmark_range( - full_range=[ - # (total_tokens, n_g, k_g, num_experts) - (1024, 4096, 4096, 64), - (2048, 4096, 4096, 64), - (4096, 4096, 4096, 64), - ] - + [ - (total_tokens, n_g, k_g, num_experts) - for total_tokens in [32 * (2**i) for i in range(9)] # 32 to 8192 - for n_g, k_g, num_experts in [ - # DeepSeek-V3/R1, gateup, TP = 1, EP = 8 - (4096, 7168, 32), - # DeepSeek-V3/R1, down, TP = 1, EP = 8 - (7168, 2048, 32), - ] - ], - ci_range=[(1024, 2048, 2048, 8)], -) - -line_vals = ["jit"] -line_names = ["JIT MXFP8 MoE GroupMM"] -styles = [("green", "-")] - -if _SGL_KERNEL_AVAILABLE: - line_vals.append("sgl_kernel") - line_names.append("sgl-kernel MXFP8 MoE GroupMM") - styles.append(("orange", "-")) - - -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["total_tokens", "n_g", "k_g", "num_experts"], - x_vals=shape_range, - x_log=False, - line_arg="provider", - line_vals=line_vals, - line_names=line_names, - styles=styles, - ylabel="us", - plot_name="mxfp8-moe-groupmm-performance", - args={}, - ) -) -def benchmark(total_tokens, n_g, k_g, num_experts, provider): - case = _prepare_case(total_tokens, n_g, k_g, num_experts, torch.bfloat16) - - if provider == "jit": - fn = lambda: es_sm100_mxfp8_blockscaled_moe_grouped_gemm( - case["b_quant"], - case["a_quant"], - case["b_scale_factor"], - case["a_scale_factor"], - case["expert_offsets"], - case["a_blockscale_offsets"], - case["tokens_per_expert"], - case["workspace"], - case["dtype"], - ) - elif provider == "sgl_kernel": - fn = lambda: _sgl_kernel_group_mm(case) - else: - raise ValueError(f"Unknown provider: {provider}") - - # Warm up - fn() - - # Profile - if provider == "jit": - torch.cuda.nvtx.range_push("jit") - fn() - torch.cuda.nvtx.range_pop() - elif provider == "sgl_kernel": - torch.cuda.nvtx.range_push("sgl_kernel") - fn() - torch.cuda.nvtx.range_pop() - - return run_benchmark(fn) - - -if __name__ == "__main__": - if not _SM100_SUPPORTED: - print("[skip] MXFP8 MoE GroupMM benchmark requires sm100+ with CUDA 12.8+.") - sys.exit(0) - if not _SGL_KERNEL_AVAILABLE: - print(f"[info] sgl-kernel baseline unavailable: {_SGL_KERNEL_REASON}") - benchmark.run(print_data=True) diff --git a/test/registered/kernels/benchmark/speculative/bench_resolve_future_token_ids.py b/test/registered/kernels/benchmark/speculative/bench_resolve_future_token_ids.py deleted file mode 100644 index 08ba11530..000000000 --- a/test/registered/kernels/benchmark/speculative/bench_resolve_future_token_ids.py +++ /dev/null @@ -1,77 +0,0 @@ -import itertools - -import torch -import triton -import triton.testing - -from sglang.kernels.jit.benchmark.utils import ( - DEFAULT_DEVICE, - get_benchmark_range, - run_benchmark, -) -from sglang.kernels.ops.speculative.resolve_future_token_ids import ( - resolve_future_token_ids_cuda, -) -from sglang.srt.utils import get_compiler_backend -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci - -register_cuda_ci( - est_time=10, stage="base-b-kernel-benchmark", runner_config="1-gpu-large" -) -register_amd_ci(est_time=10, stage="jit-kernel-benchmark", runner_config="amd") - -SIZE_LIST = get_benchmark_range( - full_range=[2**n for n in range(4, 16)], # 16 … 32K elements - ci_range=[256, 4096], -) - -configs = list(itertools.product(SIZE_LIST)) - - -def _torch_resolve(input_ids, future_map): - input_ids[:] = torch.where( - input_ids < 0, - future_map[torch.clamp(-input_ids, min=0)], - input_ids, - ) - - -_compiled_resolve = torch.compile( - _torch_resolve, dynamic=True, backend=get_compiler_backend() -) - - -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["size"], - x_vals=configs, - line_arg="provider", - line_vals=["jit", "torch_compile", "torch"], - line_names=["SGL JIT Kernel", "torch.compile", "PyTorch"], - styles=[("blue", "-"), ("green", "-."), ("red", "--")], - ylabel="us", - plot_name="resolve-future-token-ids-performance", - args={}, - ) -) -def benchmark(size: int, provider: str): - map_size = 8192 - future_map = torch.randint( - 0, 50000, (map_size,), dtype=torch.int64, device=DEFAULT_DEVICE - ) - input_ids = torch.randint( - -map_size + 1, 50000, (size,), dtype=torch.int64, device=DEFAULT_DEVICE - ) - - if provider == "jit": - fn = lambda: resolve_future_token_ids_cuda(input_ids.clone(), future_map) - elif provider == "torch_compile": - fn = lambda: _compiled_resolve(input_ids.clone(), future_map) - else: - fn = lambda: _torch_resolve(input_ids.clone(), future_map) - - return run_benchmark(fn) - - -if __name__ == "__main__": - benchmark.run(print_data=True) diff --git a/test/registered/kernels/ops/layernorm/test_kernels_namespace.py b/test/registered/kernels/ops/layernorm/test_kernels_namespace.py index 5fb7a1f09..ce7e2d541 100644 --- a/test/registered/kernels/ops/layernorm/test_kernels_namespace.py +++ b/test/registered/kernels/ops/layernorm/test_kernels_namespace.py @@ -1,6 +1,7 @@ """GPU-free import / registry / selector tests for ``sglang.kernels`` (RFC #29630).""" import importlib +import importlib.util import subprocess import sys @@ -84,6 +85,20 @@ def test_specs_well_formed(): assert sep == ":" and mod and attr, spec.target +def test_internal_registry_target_modules_exist(): + for spec in K.registry.all_specs(): + module, _, _ = spec.target.partition(":") + if module.startswith("sglang.kernels."): + assert importlib.util.find_spec(module) is not None, spec.target + + +def test_sparse_linear_attention_registry_targets_forward_kernel(): + spec = K.registry.get_backend( + "diffusion.sparse_linear_attn_fwd", KernelBackend.TRITON + ) + assert spec.target.endswith(":_attn_fwd") + + def test_single_backend_resolves_without_backend(): assert K.select_kernel("gemm.fp8_scaled_mm").backend is KernelBackend.AOT diff --git a/test/registered/kernels/ops/moe/test_kpool_topk_transform.py b/test/registered/kernels/ops/moe/test_kpool_topk_transform.py deleted file mode 100644 index 4949c6dd8..000000000 --- a/test/registered/kernels/ops/moe/test_kpool_topk_transform.py +++ /dev/null @@ -1,136 +0,0 @@ -"""Test for the kpool top-k transform JIT kernel. - -Ported from the former AOT sgl-kernel test (sgl-kernel/tests/test_topk.py). -The kernel selects pool groups at pool granularity, expands each selected group -to ``pool_size`` token indices, and optionally transforms those token indices -through a page table or a ragged offset. -""" - -from __future__ import annotations - -import sys -from typing import Optional - -import pytest -import torch - -from sglang.kernels.ops.moe.kpool_topk_transform import fast_kpool_topk_transform_fused -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="1-gpu-large") - - -def _ref_torch_kpool_transform_impl( - score: torch.Tensor, - lengths: torch.Tensor, - pool_size: int, - topk: int, - page_table: Optional[torch.Tensor] = None, - topk_indices_offset: Optional[torch.Tensor] = None, - seq_lens: Optional[torch.Tensor] = None, -) -> torch.Tensor: - rows = score.shape[0] - group_topk = topk // pool_size - offsets = torch.arange(pool_size, dtype=torch.int32, device=score.device) - out_cols = topk + (pool_size - 1 if seq_lens is not None else 0) - out = torch.full((rows, out_cols), -1, dtype=torch.int32, device=score.device) - for i in range(rows): - length = int(lengths[i].item()) - valid_count = min(length, group_topk) - write_pos = 0 - if valid_count == 0: - token_ids = torch.empty((0,), dtype=torch.int32, device=score.device) - elif length <= group_topk: - selected = torch.arange(length, dtype=torch.int32, device=score.device) - token_ids = (selected.unsqueeze(1) * pool_size + offsets).reshape(-1) - else: - selected = torch.topk( - score[i, :length], group_topk, dim=-1, sorted=False - ).indices.to(torch.int32) - token_ids = (selected.unsqueeze(1) * pool_size + offsets).reshape(-1) - if token_ids.numel() > 0: - if page_table is not None: - token_ids = page_table[i, token_ids.long()].to(torch.int32) - elif topk_indices_offset is not None: - token_ids = token_ids + topk_indices_offset[i].to(torch.int32) - write_pos = valid_count * pool_size - out[i, :write_pos] = token_ids[:write_pos] - if seq_lens is not None: - tail_count = int(seq_lens[i].item()) % pool_size - if tail_count > 0: - raw_tail = length * pool_size + torch.arange( - tail_count, dtype=torch.int32, device=score.device - ) - if page_table is not None: - tail = page_table[i, raw_tail.long()].to(torch.int32) - elif topk_indices_offset is not None: - tail = raw_tail + topk_indices_offset[i].to(torch.int32) - else: - tail = raw_tail - out[i, write_pos : write_pos + tail_count] = tail - return out - - -@pytest.mark.parametrize( - "pool_size,group_topk", - [(16, 128), (16, 160), (16, 192), (16, 224), (8, 256), (4, 512)], -) -@pytest.mark.parametrize("mode", ["raw", "paged", "ragged"]) -@pytest.mark.parametrize("append_tail", [False, True]) -@torch.inference_mode() -def test_kpool_topk_transform_kernel( - pool_size: int, group_topk: int, mode: str, append_tail: bool -) -> None: - torch.manual_seed(42) - bs = 17 - topk = pool_size * group_topk - num_groups = 4096 - score = torch.randn(bs, num_groups, dtype=torch.float32, device="cuda") - lengths = torch.randint( - group_topk + 1, num_groups + 1, (bs,), dtype=torch.int32, device="cuda" - ) - - page_table = None - topk_indices_offset = None - seq_lens = None - tail_counts = torch.randint(0, pool_size, (bs,), dtype=torch.int32, device="cuda") - if append_tail: - seq_lens = lengths * pool_size + tail_counts - if mode == "paged": - page_table = torch.arange( - bs * (num_groups * pool_size + pool_size), - dtype=torch.int32, - device="cuda", - ).view(bs, num_groups * pool_size + pool_size) - elif mode == "ragged": - topk_indices_offset = torch.randint( - 0, 2048, (bs,), dtype=torch.int32, device="cuda" - ) - - out_ref = _ref_torch_kpool_transform_impl( - score, - lengths, - pool_size, - topk, - page_table=page_table, - topk_indices_offset=topk_indices_offset, - seq_lens=seq_lens, - ) - out_our = fast_kpool_topk_transform_fused( - score, - lengths, - pool_size, - topk, - page_table=page_table, - topk_indices_offset=topk_indices_offset, - seq_lens=seq_lens, - ) - torch.cuda.synchronize() - - out_ref = torch.sort(out_ref, dim=-1).values - out_our = torch.sort(out_our, dim=-1).values - torch.testing.assert_close(out_our, out_ref, atol=0, rtol=0) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/quantization/test_mxfp8_moe.py b/test/registered/kernels/ops/quantization/test_mxfp8_moe.py deleted file mode 100644 index ce2e17ad7..000000000 --- a/test/registered/kernels/ops/quantization/test_mxfp8_moe.py +++ /dev/null @@ -1,152 +0,0 @@ -import random -import sys - -import pytest -import torch - -from sglang.kernels.ops.quantization.mxfp8 import ( - es_sm100_mxfp8_blockscaled_grouped_quant, - es_sm100_mxfp8_blockscaled_moe_grouped_gemm, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=5, stage="base-b-kernel-unit", runner_config="1-gpu-large") - - -def align(val: int, alignment: int = 128) -> int: - return int((val + alignment - 1) // alignment * alignment) - - -# Copy from: https://github.com/deepseek-ai/DeepGEMM/blob/main/deep_gemm/utils.py -def calc_diff(x, y): - x, y = x.double(), y.double() - denominator = (x * x + y * y).sum() - sim = 2 * (x * y).sum() / denominator - return 1 - sim - - -def is_sm100_supported(device=None) -> bool: - return (torch.cuda.get_device_capability(device)[0] == 10) and ( - torch.version.cuda >= "12.8" - ) - - -@pytest.mark.skipif( - not is_sm100_supported(), - reason="test_mxfp8_moe at jit kernen is only supported on sm100", -) -@pytest.mark.parametrize("num_experts", [8, 16, 32, 64]) -@pytest.mark.parametrize("out_dtype", [torch.half, torch.bfloat16]) -def test_es_sm100_mxfp8_blockscaled_grouped_mm(num_experts, out_dtype): - device = "cuda" - alignment = 128 - n_g = random.randint(1, 64) * alignment - k_g = random.randint(1, 64) * alignment - - expert_offset = 0 - expert_offsets = [] - aux_expert_offset = 0 - aux_expert_offsets = [] - a_blockscale_offset = 0 - a_blockscale_offsets = [] - b_blockscale_offset = 0 - b_blockscale_offsets = [] - a_list = [] - b_list = [] - ref_d_list = [] - tokens_per_expert = [] - - for g in range(num_experts): - m_g = random.randint(1, 512) - tokens_per_expert.append(m_g) - expert_offsets.append(expert_offset) - expert_offset += m_g - aux_expert_offsets.append(aux_expert_offset) - aux_expert_offset += n_g - a_blockscale_offsets.append(a_blockscale_offset) - a_blockscale_offset += align(m_g, 128) - b_blockscale_offsets.append(b_blockscale_offset) - b_blockscale_offset += n_g # n_g already align to 128 - - a = torch.normal( - 0.0, std=1.0, size=(m_g, k_g), device=device, dtype=out_dtype - ) # (M, K):(K, 1) - b = torch.normal( - 0.0, std=1.0, size=(n_g, k_g), device=device, dtype=out_dtype - ) # (N, K):(K, 1) - - a_list.append(a) - b_list.append(b) - ref_d = a @ b.T - ref_d_list.append(ref_d) - a = torch.concat(a_list, dim=0) - b = torch.concat(b_list, dim=0) - - _expert_offsets = torch.tensor(expert_offsets).to(device=device, dtype=torch.int32) - _aux_expert_offsets = torch.tensor(aux_expert_offsets).to( - device=device, dtype=torch.int32 - ) - _a_blockscale_offsets = torch.tensor(a_blockscale_offsets).to( - device=device, dtype=torch.int32 - ) - _b_blockscale_offsets = torch.tensor(b_blockscale_offsets).to( - device=device, dtype=torch.int32 - ) - - a_quant = torch.zeros_like(a, dtype=torch.float8_e4m3fn, device=device) - a_scale_factor = torch.zeros( - (a_blockscale_offset, k_g // 32), dtype=torch.uint8, device=device - ) - - b_quant = torch.zeros_like(b, dtype=torch.float8_e4m3fn, device=device) - b_scale_factor = torch.zeros( - (num_experts * n_g, k_g // 32), dtype=torch.uint8, device=device - ) - tokens_per_expert = torch.tensor(tokens_per_expert).to( - device=device, dtype=torch.int32 - ) - workspace = torch.empty((1024, 1024, 1024), dtype=torch.uint8, device=device) - - es_sm100_mxfp8_blockscaled_grouped_quant( - a, - tokens_per_expert, - _expert_offsets, - _a_blockscale_offsets, - a_quant, - a_scale_factor, - ) - es_sm100_mxfp8_blockscaled_grouped_quant( - b, - torch.ones_like(tokens_per_expert) * n_g, - _aux_expert_offsets, - _b_blockscale_offsets, - b_quant, - b_scale_factor, - ) - - b_quant = b_quant.view(num_experts, n_g, k_g) - b_scale_factor = b_scale_factor.view(num_experts, n_g, k_g // 32) - d = es_sm100_mxfp8_blockscaled_moe_grouped_gemm( - b_quant, - a_quant, - b_scale_factor, - a_scale_factor, - _expert_offsets, - _a_blockscale_offsets, - tokens_per_expert, - workspace, - a.dtype, - ) - - for g in range(num_experts): - baseline = ref_d_list[g] - actual = d[expert_offsets[g] : (expert_offsets[g] + tokens_per_expert[g])] - diff = calc_diff(actual, baseline) - assert diff < 0.001 - print( - f"m_g={baseline.shape[0]} n_g={n_g} k_g={k_g} num_experts={num_experts}, out_dtype={out_dtype}, diff={diff:.5f}: OK" - ) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/speculative/test_resolve_future_token_ids.py b/test/registered/kernels/ops/speculative/test_resolve_future_token_ids.py deleted file mode 100644 index 6df314484..000000000 --- a/test/registered/kernels/ops/speculative/test_resolve_future_token_ids.py +++ /dev/null @@ -1,71 +0,0 @@ -import sys - -import pytest -import torch - -from sglang.kernels.ops.speculative.resolve_future_token_ids import ( - resolve_future_token_ids_cuda, -) -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci - -register_cuda_ci(est_time=9, stage="base-b-kernel-unit", runner_config="1-gpu-large") -register_amd_ci(est_time=9, stage="jit-kernel-unit", runner_config="amd") - - -def _reference_resolve(input_ids, future_map): - """Reference implementation using plain torch.""" - result = input_ids.clone() - result[:] = torch.where( - result < 0, - future_map[torch.clamp(-result, min=0)], - result, - ) - return result - - -@pytest.mark.parametrize("size", [1, 2, 127, 128, 255, 256, 1024, 4097]) -@pytest.mark.parametrize("dtype", [torch.int32, torch.int64]) -class TestResolveFutureTokenIds: - def test_all_negative(self, size: int, dtype: torch.dtype) -> None: - map_size = 8192 - future_map = torch.randint(0, 50000, (map_size,), dtype=dtype, device="cuda") - # Negative indices in range [-map_size+1, -1] - input_ids = -torch.randint(1, map_size, (size,), dtype=dtype, device="cuda") - - expected = _reference_resolve(input_ids, future_map) - resolve_future_token_ids_cuda(input_ids, future_map) - assert torch.equal(input_ids, expected) - - def test_all_non_negative(self, size: int, dtype: torch.dtype) -> None: - map_size = 16 - future_map = torch.randint(0, 50000, (map_size,), dtype=dtype, device="cuda") - input_ids = torch.randint(0, 50000, (size,), dtype=dtype, device="cuda") - - expected = input_ids.clone() - resolve_future_token_ids_cuda(input_ids, future_map) - assert torch.equal(input_ids, expected) - - def test_mixed(self, size: int, dtype: torch.dtype) -> None: - map_size = 8192 - future_map = torch.randint(0, 50000, (map_size,), dtype=dtype, device="cuda") - # Mix of negative and non-negative - input_ids = torch.randint( - -map_size + 1, 50000, (size,), dtype=dtype, device="cuda" - ) - - expected = _reference_resolve(input_ids, future_map) - resolve_future_token_ids_cuda(input_ids, future_map) - assert torch.equal(input_ids, expected) - - def test_zeros(self, size: int, dtype: torch.dtype) -> None: - map_size = 16 - future_map = torch.randint(0, 50000, (map_size,), dtype=dtype, device="cuda") - input_ids = torch.zeros(size, dtype=dtype, device="cuda") - - expected = input_ids.clone() - resolve_future_token_ids_cuda(input_ids, future_map) - assert torch.equal(input_ids, expected) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"]))