diff --git a/python/sglang/kernels/jit/csrc/attention/kda_packed_decode.cuh b/python/sglang/kernels/jit/csrc/attention/kda_packed_decode.cuh index 4c52324de..e288352d2 100644 --- a/python/sglang/kernels/jit/csrc/attention/kda_packed_decode.cuh +++ b/python/sglang/kernels/jit/csrc/attention/kda_packed_decode.cuh @@ -21,6 +21,7 @@ #include // For bf16_t, fp32_t, device::cast #include // For LaunchKernel +#include #include @@ -48,19 +49,6 @@ struct KdaPackedDecodeParams { int32_t use_lower_bound; }; -__device__ __forceinline__ float warp_allreduce_sum(float v) { -#if defined(__HIP_PLATFORM_AMD__) - constexpr uint64_t kFullMask = 0xffffffffffffffffull; -#else - constexpr uint32_t kFullMask = 0xffffffffu; -#endif -#pragma unroll - for (int off = 16; off > 0; off >>= 1) { - v += __shfl_xor_sync(kFullMask, v, off); - } - return v; -} - // K = V = 128 specialization: one lane owns 4 consecutive K-elements (16B). template __global__ @@ -104,8 +92,8 @@ __launch_bounds__(kWarps * 32) void kda_packed_decode_kernel(const KdaPackedDeco k_sq += k[e] * k[e]; } // tl: q / sqrt(sum(q*q) + 1e-6), then * scale - const float q_inv = 1.0f / sqrtf(warp_allreduce_sum(q_sq) + 1e-6f); - const float k_inv = 1.0f / sqrtf(warp_allreduce_sum(k_sq) + 1e-6f); + const float q_inv = 1.0f / sqrtf(warp::reduce_sum<32>(q_sq) + 1e-6f); + const float k_inv = 1.0f / sqrtf(warp::reduce_sum<32>(k_sq) + 1e-6f); #pragma unroll for (int e = 0; e < kElems; ++e) { q[e] = q[e] * q_inv * params.scale; @@ -143,7 +131,7 @@ __launch_bounds__(kWarps * 32) void kda_packed_decode_kernel(const KdaPackedDeco h[e] *= decay[e]; t += h[e] * k[e]; } - t = warp_allreduce_sum(t); + t = warp::reduce_sum<32>(t); const float v_new = (cast(v_ptr[r]) - t) * beta; float o_acc = 0.0f; #pragma unroll @@ -151,7 +139,7 @@ __launch_bounds__(kWarps * 32) void kda_packed_decode_kernel(const KdaPackedDeco h[e] += v_new * k[e]; o_acc += h[e] * q[e]; } - o_acc = warp_allreduce_sum(o_acc); + o_acc = warp::reduce_sum<32>(o_acc); *reinterpret_cast(h_base + r * K + e0) = make_float4(h[0], h[1], h[2], h[3]); if (lane == 0) { o_ptr[r] = cast(o_acc); diff --git a/python/sglang/kernels/jit/csrc/inkling/inkling_ar_fused_decode.cuh b/python/sglang/kernels/jit/csrc/inkling/inkling_ar_fused_decode.cuh index 99de47c25..edbc01c8d 100644 --- a/python/sglang/kernels/jit/csrc/inkling/inkling_ar_fused_decode.cuh +++ b/python/sglang/kernels/jit/csrc/inkling/inkling_ar_fused_decode.cuh @@ -161,7 +161,7 @@ __global__ __launch_bounds__(1024, 1) void inkling_ar_sconv_norm_kernel(const __ // ---- 1. push: wait for the producer's output (PDL; no-op without a PDL // launch or an early-triggering producer), multicast-store this rank's // partial row, and issue the residual load (it lands under the barrier). ---- - asm volatile("griddepcontrol.wait;" ::: "memory"); + device::PDLWaitPrimary(); const auto* in_row = static_cast(p.in) + t * p.in_stride_t; const auto* sh_row = p.shared == nullptr ? nullptr : static_cast(p.shared) + t * p.shared_stride_t; @@ -393,7 +393,7 @@ __launch_bounds__(1024, 1) void inkling_ar_sconv_norm_verify_kernel(const __grid // GPU, so Phase 2's cross-token (neighbor) staging reads are race-free -- the // per-block barrier only synchronized the same blockIdx across ranks and did // NOT order block t-j's push before block t's read. - asm volatile("griddepcontrol.wait;" ::: "memory"); + device::PDLWaitPrimary(); auto* mc = static_cast<__nv_bfloat16*>(p.mc_stage); const auto* in = static_cast(p.in); const auto* sh = static_cast(p.shared); diff --git a/python/sglang/kernels/jit/csrc/inkling/inkling_ar_scattered_sconv.cuh b/python/sglang/kernels/jit/csrc/inkling/inkling_ar_scattered_sconv.cuh index 76da8f90c..87da53eb7 100644 --- a/python/sglang/kernels/jit/csrc/inkling/inkling_ar_scattered_sconv.cuh +++ b/python/sglang/kernels/jit/csrc/inkling/inkling_ar_scattered_sconv.cuh @@ -1028,7 +1028,7 @@ __global__ __launch_bounds__(1024, 1) void inkling_ar_col_decode_kernel(const __ static_cast(p.residual_in) + static_cast(t) * p.H + c0[i]); } } - asm volatile("griddepcontrol.wait;" ::: "memory"); + device::PDLWaitPrimary(); // ---- 1. entry: peers' producer partials visible (block t <-> peer t) ---- inkling_ar::block_system_barrier(p.state, p.flag_ptrs, p.rank); @@ -1764,7 +1764,7 @@ __launch_bounds__(1024, 1) void inkling_ar_ssconv_norm_decode_kernel(const __gri } // ---- 1. push: this rank's partial row + its window shard ---- - asm volatile("griddepcontrol.wait;" ::: "memory"); + device::PDLWaitPrimary(); const auto* in_row = static_cast(p.in) + t * p.in_stride_t; auto* slot = static_cast<__nv_bfloat16*>(p.mc_stage) + (static_cast(p.rank) * p.T + t) * p.D; auto* wslot = static_cast<__nv_bfloat16*>(p.mc_wstage) + static_cast(t) * W1 * p.D; diff --git a/python/sglang/kernels/jit/csrc/moe/moe_align_kernel.cu b/python/sglang/kernels/jit/csrc/moe/moe_align_kernel.cu index f82ab277b..5c1726d4f 100644 --- a/python/sglang/kernels/jit/csrc/moe/moe_align_kernel.cu +++ b/python/sglang/kernels/jit/csrc/moe/moe_align_kernel.cu @@ -13,10 +13,12 @@ See the License for the specific language governing permissions and limitations under the License. ==============================================================================*/ +#include #include #include #include +#include #include @@ -33,28 +35,8 @@ namespace sglang { using Vec = int4; -inline uint32_t next_pow2(uint32_t x) noexcept { - --x; - x |= x >> 1; - x |= x >> 2; - x |= x >> 4; - x |= x >> 8; - x |= x >> 16; - return x + 1; -} - namespace moe { -__device__ __forceinline__ int warp_exclusive_scan(int v, unsigned mask = 0xffffffffu) { - int original = v; -#pragma unroll - for (int offset = 1; offset < WARP_SIZE; offset <<= 1) { - int n = __shfl_up_sync(mask, v, offset); - if ((threadIdx.x & (WARP_SIZE - 1)) >= offset) v += n; - } - return v - original; -} - template __global__ void count_and_sort_expert_tokens_kernel( const scalar_t* __restrict__ topk_ids, @@ -187,14 +169,14 @@ __global__ void moe_align_block_size_kernel( const int warp_id = tid / WARP_SIZE; const int lane_id = tid & (WARP_SIZE - 1); const int num_warps_for_scan = (scan_size + WARP_SIZE - 1) / WARP_SIZE; - const int warp_sum = warp_exclusive_scan(padded_count) + padded_count; + const int warp_sum = device::warp::inclusive_sum<32>(padded_count); if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = warp_sum; __syncthreads(); // warp0 accumulate all the block's prefix sum if (tid < WARP_SIZE) { int val = (tid < num_warps_for_scan) ? warp_sums[tid] : 0; - int incl = warp_exclusive_scan(val) + val; + int incl = device::warp::inclusive_sum<32>(val); warp_sums[tid] = incl; } __syncthreads(); @@ -213,13 +195,13 @@ __global__ void moe_align_block_size_kernel( // Perform 2 level exclusive-prefix-sum to scan_buf int v = (tid < scan_size) ? scan_buf[tid] : 0; - int pre = warp_exclusive_scan(v); + int pre = device::warp::inclusive_sum<32>(v) - v; if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = pre + v; __syncthreads(); if (warp_id == 0) { int val = (lane_id < num_warps_for_scan) ? warp_sums[lane_id] : 0; - warp_sums[lane_id] = warp_exclusive_scan(val); + warp_sums[lane_id] = device::warp::inclusive_sum<32>(val) - val; } __syncthreads(); @@ -409,7 +391,7 @@ __global__ void moe_align_block_size_kernel_v2( } // Level 1: intra-warp exclusive scan on thread_sum - int32_t warp_prefix = warp_exclusive_scan(thread_sum); + int32_t warp_prefix = device::warp::inclusive_sum<32>(thread_sum) - thread_sum; int32_t warp_total = warp_prefix + thread_sum; if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = warp_total; __syncthreads(); @@ -418,7 +400,7 @@ __global__ void moe_align_block_size_kernel_v2( const int num_warps = (blockDim.x + WARP_SIZE - 1) / WARP_SIZE; if (tid < WARP_SIZE) { int val = (tid < num_warps) ? warp_sums[tid] : 0; - warp_sums[tid] = warp_exclusive_scan(val); + warp_sums[tid] = device::warp::inclusive_sum<32>(val) - val; } __syncthreads(); @@ -510,7 +492,7 @@ struct MoeAlignBlockSizeKernel { pad_sorted_token_ids, (int32_t)max_num_tokens_padded); } else if (num_experts <= 1024) { - const size_t scan_size = next_pow2(num_experts); + const size_t scan_size = host::round_up_pow2(static_cast(num_experts)); const size_t shared_mem_size = (num_experts + (num_experts + 1) + scan_size + WARP_SIZE) * sizeof(int32_t); auto align_kernel = moe::moe_align_block_size_kernel; diff --git a/python/sglang/kernels/jit/csrc/trtllm_lora_temp/moe_lora_merged_align_kernel.cu b/python/sglang/kernels/jit/csrc/trtllm_lora_temp/moe_lora_merged_align_kernel.cu index cf8d15975..a158e3bcb 100644 --- a/python/sglang/kernels/jit/csrc/trtllm_lora_temp/moe_lora_merged_align_kernel.cu +++ b/python/sglang/kernels/jit/csrc/trtllm_lora_temp/moe_lora_merged_align_kernel.cu @@ -24,10 +24,12 @@ limitations under the License. // using fused scatter for eligible shapes and two kernels otherwise. Larger // domains keep the old path through the Python dispatcher. +#include #include #include #include +#include #include @@ -44,28 +46,8 @@ namespace sglang { using Vec = int4; -inline uint32_t next_pow2(uint32_t x) noexcept { - --x; - x |= x >> 1; - x |= x >> 2; - x |= x >> 4; - x |= x >> 8; - x |= x >> 16; - return x + 1; -} - namespace moe_lora_merged { -__device__ __forceinline__ int warp_exclusive_scan(int v, unsigned mask = 0xffffffffu) { - int original = v; -#pragma unroll - for (int offset = 1; offset < WARP_SIZE; offset <<= 1) { - int n = __shfl_up_sync(mask, v, offset); - if ((threadIdx.x & (WARP_SIZE - 1)) >= offset) v += n; - } - return v - original; -} - // Inline mirror of _fused_virtual_topk_ids_kernel (virtual_experts.py). Returns // the merged virtual expert id for flat slot `i` (range [-1, virtual_num_experts); // -1 is the dropped/masked sentinel). The caller adds +1 to get the histogram @@ -232,14 +214,14 @@ __global__ void moe_align_block_size_kernel( const int warp_id = tid / WARP_SIZE; const int lane_id = tid & (WARP_SIZE - 1); const int num_warps_for_scan = (scan_size + WARP_SIZE - 1) / WARP_SIZE; - const int warp_sum = warp_exclusive_scan(padded_count) + padded_count; + const int warp_sum = device::warp::inclusive_sum<32>(padded_count); if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = warp_sum; __syncthreads(); // warp0 accumulate all the block's prefix sum if (tid < WARP_SIZE) { int val = (tid < num_warps_for_scan) ? warp_sums[tid] : 0; - int incl = warp_exclusive_scan(val) + val; + int incl = device::warp::inclusive_sum<32>(val); warp_sums[tid] = incl; } __syncthreads(); @@ -258,13 +240,13 @@ __global__ void moe_align_block_size_kernel( // Perform 2 level exclusive-prefix-sum to scan_buf int v = (tid < scan_size) ? scan_buf[tid] : 0; - int pre = warp_exclusive_scan(v); + int pre = device::warp::inclusive_sum<32>(v) - v; if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = pre + v; __syncthreads(); if (warp_id == 0) { int val = (lane_id < num_warps_for_scan) ? warp_sums[lane_id] : 0; - warp_sums[lane_id] = warp_exclusive_scan(val); + warp_sums[lane_id] = device::warp::inclusive_sum<32>(val) - val; } __syncthreads(); @@ -384,12 +366,12 @@ __global__ void fused_align_scatter_kernel( padded_count = (count + block_size - 1) / block_size * block_size; scan_buf[tid] = padded_count; } - const int warp_sum = warp_exclusive_scan(padded_count) + padded_count; + const int warp_sum = device::warp::inclusive_sum<32>(padded_count); if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = warp_sum; __syncthreads(); if (tid < WARP_SIZE) { int val = (tid < num_warps_for_scan) ? warp_sums[tid] : 0; - int incl = warp_exclusive_scan(val) + val; + int incl = device::warp::inclusive_sum<32>(val); warp_sums[tid] = incl; } __syncthreads(); @@ -402,12 +384,12 @@ __global__ void fused_align_scatter_kernel( if (tid >= num_experts && tid < scan_size) scan_buf[tid] = 0; __syncthreads(); int v = (tid < scan_size) ? scan_buf[tid] : 0; - int pre = warp_exclusive_scan(v); + int pre = device::warp::inclusive_sum<32>(v) - v; if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = pre + v; __syncthreads(); if (warp_id == 0) { int val = (lane_id < num_warps_for_scan) ? warp_sums[lane_id] : 0; - warp_sums[lane_id] = warp_exclusive_scan(val); + warp_sums[lane_id] = device::warp::inclusive_sum<32>(val) - val; } __syncthreads(); int off = warp_sums[warp_id]; @@ -500,7 +482,7 @@ struct MoeLoraMergedAlignKernel { int32_t* cumsum_buffer_ptr = static_cast(cumsum_buffer.data_ptr()); size_t numel = topk_ids.numel(); - const size_t scan_size = next_pow2(num_experts); + const size_t scan_size = host::round_up_pow2(static_cast(num_experts)); if (fuse_scatter) { // One block does fill + histogram + scan + expert_ids + scatter. Extra diff --git a/python/sglang/kernels/jit/utils/__init__.py b/python/sglang/kernels/jit/utils/__init__.py index 41142291c..41c25d6f9 100644 --- a/python/sglang/kernels/jit/utils/__init__.py +++ b/python/sglang/kernels/jit/utils/__init__.py @@ -1,6 +1,7 @@ """Public interface of sglang.kernels.jit.utils.""" from sglang.kernels.jit.utils.arch import ( + get_activation_cuda_cflags, get_jit_cuda_arch, is_arch_support_pdl, override_jit_cuda_arch, @@ -27,6 +28,7 @@ __all__ = [ "make_cpp_args", "load_jit", "override_jit_cuda_arch", + "get_activation_cuda_cflags", "get_jit_cuda_arch", "is_arch_support_pdl", "KERNEL_PATH", diff --git a/python/sglang/kernels/jit/utils/arch.py b/python/sglang/kernels/jit/utils/arch.py index ade8d4927..4183ab9f3 100644 --- a/python/sglang/kernels/jit/utils/arch.py +++ b/python/sglang/kernels/jit/utils/arch.py @@ -179,3 +179,11 @@ def is_arch_support_pdl() -> bool: if is_hip_runtime() or is_musa_runtime(): return False return get_jit_cuda_arch().major >= 9 + + +def get_activation_cuda_cflags() -> list[str]: + """Match the AOT activation fast-math policy without changing other kernels.""" + # Blackwell needs precise expf; HIP clang rejects --use_fast_math. + if is_hip_runtime() or get_jit_cuda_arch().major >= 10: + return [] + return ["--use_fast_math"] diff --git a/python/sglang/kernels/ops/activation/activation.py b/python/sglang/kernels/ops/activation/activation.py index 41e63514c..9108fadd2 100644 --- a/python/sglang/kernels/ops/activation/activation.py +++ b/python/sglang/kernels/ops/activation/activation.py @@ -6,9 +6,8 @@ import torch from sglang.kernels.jit.utils import ( cache_once, - get_jit_cuda_arch, + get_activation_cuda_cflags, is_arch_support_pdl, - is_hip_runtime, load_jit, make_cpp_args, ) @@ -18,19 +17,9 @@ if TYPE_CHECKING: from tvm_ffi.module import Module -def _fast_math_flags() -> list[str]: - # Mirrors sgl-kernel's CMake policy: fast-math on SM90, precise on - # SM100+ (Blackwell needs bit-exact expf), off on HIP (clang rejects). - if is_hip_runtime(): - return [] - if get_jit_cuda_arch().major >= 10: - return [] - return ["--use_fast_math"] - - @cache_once def activation_module(dtype: torch.dtype, *, fast_math: bool = True) -> Module: - fast_math_flags = _fast_math_flags() + fast_math_flags = get_activation_cuda_cflags() if not fast_math and not fast_math_flags: return activation_module(dtype) args = make_cpp_args(dtype, is_arch_support_pdl()) diff --git a/python/sglang/kernels/ops/attention/dsa/dequant_k_cache.py b/python/sglang/kernels/ops/attention/dsa/dequant_k_cache.py index 55562283f..7ce302350 100644 --- a/python/sglang/kernels/ops/attention/dsa/dequant_k_cache.py +++ b/python/sglang/kernels/ops/attention/dsa/dequant_k_cache.py @@ -9,49 +9,6 @@ def dequantize_k_cache(quant_k_cache): return _dequantize_k_cache_fast_wrapped(quant_k_cache) -def _dequantize_k_cache_ref( - quant_k_cache: torch.Tensor, # (num_blocks, block_size, 1, bytes_per_token) - dv: int = 512, - tile_size: int = 128, - d: int = 576, -) -> torch.Tensor: - """ - De-quantize the k-cache - """ - assert dv % tile_size == 0 - original_ndim = quant_k_cache.ndim - if original_ndim == 3: - # set block_size = 1 - quant_k_cache = quant_k_cache.unsqueeze(1) - num_tiles = dv // tile_size - num_blocks, block_size, h_k, _ = quant_k_cache.shape - assert h_k == 1 - result = torch.empty( - (num_blocks, block_size, d), dtype=torch.bfloat16, device=quant_k_cache.device - ) - - quant_k_cache = quant_k_cache.view(num_blocks, block_size, -1) - - input_nope = quant_k_cache[..., :dv] - input_scale = quant_k_cache[..., dv : dv + num_tiles * 4].view(torch.float32) - input_rope = quant_k_cache[..., dv + num_tiles * 4 :].view(torch.bfloat16) - result[..., dv:] = input_rope - - for tile_idx in range(0, num_tiles): - cur_nope = input_nope[ - ..., tile_idx * tile_size : (tile_idx + 1) * tile_size - ].to(torch.float32) - cur_scales = input_scale[..., tile_idx].unsqueeze(-1) - result[..., tile_idx * tile_size : (tile_idx + 1) * tile_size] = ( - cur_nope * cur_scales - ) - - if original_ndim == 3: - return result.view(num_blocks, 1, -1) - else: - return result.view(num_blocks, block_size, 1, -1) - - def _dequantize_k_cache_fast_wrapped( quant_k_cache: torch.Tensor, dv: int = 512, diff --git a/python/sglang/kernels/ops/attention/dsa/triton_sparse_mla_decode.py b/python/sglang/kernels/ops/attention/dsa/triton_sparse_mla_decode.py index 569e2c922..95cc5d784 100644 --- a/python/sglang/kernels/ops/attention/dsa/triton_sparse_mla_decode.py +++ b/python/sglang/kernels/ops/attention/dsa/triton_sparse_mla_decode.py @@ -11,17 +11,19 @@ Two variants: 2. Split-K: adaptive split-K with fused fast path (adapted from DSv4) """ -import functools - import torch import triton import triton.language as tl from sglang.kernels.ops.attention.dsa.triton_sparse_mla import ( _PREFERRED_BLOCK_K, + _cu_count, + _kv_splits_heuristic, + _next_pow2, _no_async_copy, _row_strides, _sparse_mla_block_k, + _sparse_mla_reduce_kernel, _validate_input_dtypes, ) from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz @@ -70,44 +72,6 @@ def _get_splitk_bufs( LOG2E = 1.4426950408889634 -@functools.lru_cache(maxsize=1) -def _cu_count() -> int: - return torch.cuda.get_device_properties( - torch.cuda.current_device() - ).multi_processor_count - - -def _prev_pow2(n: int) -> int: - if n < 1: - return 1 - return 1 << (n.bit_length() - 1) - - -def _next_pow2(n: int) -> int: - if n < 1: - return 1 - return 1 << (n - 1).bit_length() - - -def _kv_splits_heuristic( - T: int, - H: int, - block_h: int, - num_cu: int | None = None, - target_wg_per_cu: float = 2.0, - max_kv_splits: int = 64, -) -> int: - if num_cu is None: - num_cu = _cu_count() - target_wg = max(1, int(target_wg_per_cu * num_cu)) - head_blocks = max(1, (H + block_h - 1) // block_h) - base_ctas = max(1, T * head_blocks) - if base_ctas >= target_wg: - return 1 - splits_to_fill = max(1, target_wg // base_ctas) - return _prev_pow2(min(splits_to_fill, max_kv_splits)) - - @triton.jit def _sparse_mla_decode_fused_kernel( q_nope_ptr, # [N, H, D_V] @@ -512,64 +476,6 @@ def _sparse_mla_decode_split_kernel( ) -@triton.jit -def _sparse_mla_decode_reduce_kernel( - lse_partial_ptr, # [N, KV_SPLITS, H_padded] fp32 - acc_partial_ptr, # [N, KV_SPLITS, H_padded, D_V] bf16 - out_ptr, # [N, H, D_V] - H: tl.constexpr, - D_V: tl.constexpr, - KV_SPLITS: tl.constexpr, - ACTIVE_SPLITS: tl.constexpr, - ACTIVE_SPLITS_POW2: tl.constexpr, - D_CHUNK: tl.constexpr, - BLOCK_K: tl.constexpr, -): - t = tl.program_id(0) - h = tl.program_id(1) - dc = tl.program_id(2) - - d_offs = dc * D_CHUNK + tl.arange(0, D_CHUNK) - # tl.arange needs a power-of-two extent, but ACTIVE_SPLITS is only a power - # of two when topk // BLOCK_K is. Iterate over the padded range and mask the - # tail: -3.4e38 drives exp2() to 0 without the NaN an -inf would produce. - k_offs = tl.arange(0, ACTIVE_SPLITS_POW2) - k_mask = k_offs < ACTIVE_SPLITS - d_mask = d_offs < D_V - - H_padded = tl.cdiv(H, 16) * 16 - - lse_base = t * KV_SPLITS * H_padded - lse_p = tl.load( - lse_partial_ptr + lse_base + k_offs * H_padded + h, - mask=k_mask, - other=-3.4e38, - ) - - ap_base = t * KV_SPLITS * H_padded * D_V - a_p = tl.load( - acc_partial_ptr - + ap_base - + k_offs[:, None] * H_padded * D_V - + h * D_V - + d_offs[None, :], - mask=k_mask[:, None] & d_mask[None, :], - other=0.0, - ).to(tl.float32) - - lse_max = tl.max(lse_p, axis=0) - weights = tl.exp2(lse_p - lse_max) - w_sum = tl.sum(weights, axis=0) - scale = tl.exp2(lse_p - lse_max - tl.log2(tl.maximum(w_sum, 1.0e-30))) - out = tl.sum(a_p * scale[:, None], axis=0) - - tl.store( - out_ptr + t * H * D_V + h * D_V + d_offs, - out.to(tl.bfloat16), - mask=d_mask, - ) - - def triton_sparse_mla_decode_splitk( q_nope: torch.Tensor, q_rope: torch.Tensor, @@ -710,7 +616,7 @@ def triton_sparse_mla_decode_splitk( D_CHUNK = 64 grid_reduce = (bs, H, (d_v + D_CHUNK - 1) // D_CHUNK) - _sparse_mla_decode_reduce_kernel[grid_reduce]( + _sparse_mla_reduce_kernel[grid_reduce]( lse_partial, acc_partial, out, diff --git a/python/sglang/kernels/ops/attention/linear/gdn_blackwell/kernel_h.py b/python/sglang/kernels/ops/attention/linear/gdn_blackwell/kernel_h.py index 9f1dc2c02..ddf95dcd9 100644 --- a/python/sglang/kernels/ops/attention/linear/gdn_blackwell/kernel_h.py +++ b/python/sglang/kernels/ops/attention/linear/gdn_blackwell/kernel_h.py @@ -17,6 +17,10 @@ from sglang.kernels.ops.attention.cute_utils import ( fence_before_tma_store, simple_tma_copy, ) +from sglang.kernels.ops.attention.linear.tma import ( + make_chunk_tma_args, + make_recurrent_state_tma_args, +) class Sm100ChunkHKernel: @@ -53,46 +57,6 @@ class Sm100ChunkHKernel: self.num_stages = num_stages self.num_warps = 10 - @cute.jit - def _make_bf16_tma_args( - self, - tensor: cute.Tensor, - dim: cutlass.Constexpr[int], - op: cpasync.TmaCopyOp, - stages: cutlass.Constexpr[int], - ): - swizzle_128B = cute.make_swizzle(3, 4, 3) - slayout = cute.make_layout( - (self.BT, 1, (64, dim // 64), stages), - stride=(64, 0, (1, self.BT * 64), self.BT * dim), - ) - slayout = cute.make_composed_layout(swizzle_128B, 0, slayout) - atom, tma_tensor = cpasync.make_tiled_tma_atom( - op, - cute.logical_divide(tensor, (None, None, 64)), - slayout, - cta_tiler=(self.BT, 1, dim), - ) - return atom, tma_tensor, slayout - - @cute.jit - def _make_h_tma_args(self, tensor: cute.Tensor, op: cpasync.TmaCopyOp): - # number of elements to fill 128B - num_elems = 128 // (tensor.element_type.width // 8) - swizzle_128B = cute.make_swizzle(3, 4, 3) - slayout = cute.make_layout( - (1, 1, self.V_dim, (num_elems, self.K_dim // num_elems)), - stride=(0, 0, num_elems, (1, self.V_dim * num_elems)), - ) - slayout = cute.make_composed_layout(swizzle_128B, 0, slayout) - atom, tma_tensor = cpasync.make_tiled_tma_atom( - op, - cute.logical_divide(tensor, (None, None, None, num_elems)), - slayout, - cta_tiler=(1, 1, self.V_dim, self.K_dim), - ) - return atom, tma_tensor, slayout - @cute.jit def __call__( self, @@ -112,13 +76,13 @@ class Sm100ChunkHKernel: tma_g2s = cpasync.CopyBulkTensorTileG2SOp() tma_s2g = cpasync.CopyBulkTensorTileS2GOp() - K_args = self._make_bf16_tma_args(K, self.K_dim, tma_g2s, self.num_stages) - V_args = self._make_bf16_tma_args(V, self.V_dim, tma_g2s, self.num_stages) - W_args = self._make_bf16_tma_args(W, self.K_dim, tma_g2s, self.num_stages) - V_new_args = self._make_bf16_tma_args(V_new, self.V_dim, tma_s2g, 1) - H0_args = self._make_h_tma_args(h0, tma_g2s) - HT_args = self._make_h_tma_args(ht, tma_s2g) - H_args = self._make_h_tma_args(h, tma_s2g) + K_args = make_chunk_tma_args(K, self.K_dim, tma_g2s, self.num_stages, self.BT) + V_args = make_chunk_tma_args(V, self.V_dim, tma_g2s, self.num_stages, self.BT) + W_args = make_chunk_tma_args(W, self.K_dim, tma_g2s, self.num_stages, self.BT) + V_new_args = make_chunk_tma_args(V_new, self.V_dim, tma_s2g, 1, self.BT) + H0_args = make_recurrent_state_tma_args(h0, tma_g2s, self.K_dim, self.V_dim) + HT_args = make_recurrent_state_tma_args(ht, tma_s2g, self.K_dim, self.V_dim) + H_args = make_recurrent_state_tma_args(h, tma_s2g, self.K_dim, self.V_dim) # h0/ht may be the full state pool ([num_slots, ...]) rather than a # per-sequence gather, so the sequence count comes from cu_seqlens and diff --git a/python/sglang/kernels/ops/attention/linear/gdn_blackwell/kernel_kkt_inv_uw.py b/python/sglang/kernels/ops/attention/linear/gdn_blackwell/kernel_kkt_inv_uw.py index c600018e4..48570a1b1 100644 --- a/python/sglang/kernels/ops/attention/linear/gdn_blackwell/kernel_kkt_inv_uw.py +++ b/python/sglang/kernels/ops/attention/linear/gdn_blackwell/kernel_kkt_inv_uw.py @@ -18,6 +18,9 @@ from sglang.kernels.ops.attention.cute_utils import ( mma_bf16, simple_tma_copy, ) +from sglang.kernels.ops.attention.linear.tma import ( + make_chunk_tma_args, +) class Sm100ChunkUWKernel: @@ -50,33 +53,6 @@ class Sm100ChunkUWKernel: self.BT = 64 self.num_warps = 2 + 4 + 4 - @cute.jit - def _make_tma_args( - self, - tensor: cute.Tensor, - dim: cutlass.Constexpr[int], - num_stages: int, - op: cpasync.TmaCopyOp, - ): - # logical layout: [BT, dim] - # permute for TMA: [dim/64, BT, 64] with swizzling - swizzle_128B = cute.make_swizzle(3, 4, 3) - slayout = cute.make_layout( - (self.BT, 1, (64, dim // 64), num_stages), - stride=(64, 0, (1, self.BT * 64), self.BT * dim), - ) - slayout = cute.make_composed_layout(swizzle_128B, 0, slayout) - - # we need to convert gmem layout to (T, H, (64, D/64)) for make_tiled_tma_atom() - # to emit a single 4D TMA. otherwise, it will emit (D/64)x 3D TMA. - atom, tma_tensor = cpasync.make_tiled_tma_atom( - op, - cute.logical_divide(tensor, (None, None, 64)), - slayout, - cta_tiler=(self.BT, 1, dim), - ) - return atom, tma_tensor, slayout - @cute.jit def __call__( self, @@ -96,10 +72,10 @@ class Sm100ChunkUWKernel: tma_g2s = cpasync.CopyBulkTensorTileG2SOp() tma_s2g = cpasync.CopyBulkTensorTileS2GOp() - K_args = self._make_tma_args(K, self.K_dim, self.num_stages, tma_g2s) - V_args = self._make_tma_args(V, self.V_dim, self.num_stages, tma_g2s) - U_args = self._make_tma_args(U, self.V_dim, 1, tma_s2g) - W_args = self._make_tma_args(W, self.K_dim, 1, tma_s2g) + K_args = make_chunk_tma_args(K, self.K_dim, tma_g2s, self.num_stages, self.BT) + V_args = make_chunk_tma_args(V, self.V_dim, tma_g2s, self.num_stages, self.BT) + U_args = make_chunk_tma_args(U, self.V_dim, tma_s2g, 1, self.BT) + W_args = make_chunk_tma_args(W, self.K_dim, tma_s2g, 1, self.BT) grid = (num_sms // self.Hv, self.Hv, 1) block = (self.num_warps * 32, 1, 1) diff --git a/python/sglang/kernels/ops/attention/linear/gdn_blackwell/kernel_o.py b/python/sglang/kernels/ops/attention/linear/gdn_blackwell/kernel_o.py index 203e4e2fd..6d79efefc 100644 --- a/python/sglang/kernels/ops/attention/linear/gdn_blackwell/kernel_o.py +++ b/python/sglang/kernels/ops/attention/linear/gdn_blackwell/kernel_o.py @@ -17,6 +17,10 @@ from sglang.kernels.ops.attention.cute_utils import ( fence_before_tma_store, simple_tma_copy, ) +from sglang.kernels.ops.attention.linear.tma import ( + make_chunk_tma_args, + make_output_state_tma_args, +) class Sm100ChunkOKernel: @@ -48,50 +52,6 @@ class Sm100ChunkOKernel: self.num_stages = num_stages self.num_warps = 10 - @cute.jit - def _make_bf16_tma_args( - self, - tensor: cute.Tensor, - dim: cutlass.Constexpr[int], - op: cpasync.TmaCopyOp, - stages: cutlass.Constexpr[int], - ): - swizzle_128B = cute.make_swizzle(3, 4, 3) - slayout = cute.make_layout( - (self.BT, 1, (64, dim // 64), stages), - stride=(64, 0, (1, self.BT * 64), self.BT * dim), - ) - slayout = cute.make_composed_layout(swizzle_128B, 0, slayout) - atom, tma_tensor = cpasync.make_tiled_tma_atom( - op, - cute.logical_divide(tensor, (None, None, 64)), - slayout, - cta_tiler=(self.BT, 1, dim), - ) - return atom, tma_tensor, slayout - - @cute.jit - def _make_h_tma_args( - self, - tensor: cute.Tensor, - op: cpasync.TmaCopyOp, - stages: cutlass.Constexpr[int], - ): - num_elems = 128 // (tensor.element_type.width // 8) - swizzle_128B = cute.make_swizzle(3, 4, 3) - slayout = cute.make_layout( - (1, self.V_dim, (num_elems, self.K_dim // num_elems), stages), - stride=(0, num_elems, (1, self.V_dim * num_elems), self.V_dim * self.K_dim), - ) - slayout = cute.make_composed_layout(swizzle_128B, 0, slayout) - atom, tma_tensor = cpasync.make_tiled_tma_atom( - op, - cute.logical_divide(tensor, (None, None, num_elems)), - slayout, - cta_tiler=(1, self.V_dim, self.K_dim), - ) - return atom, tma_tensor, slayout - @cute.jit def __call__( self, @@ -112,13 +72,15 @@ class Sm100ChunkOKernel: block = (self.num_warps * 32, 1, 1) tma_g2s = cpasync.CopyBulkTensorTileG2SOp() tma_s2g = cpasync.CopyBulkTensorTileS2GOp() - Q_args = self._make_bf16_tma_args(q, self.K_dim, tma_g2s, self.num_stages) - K_args = self._make_bf16_tma_args(k, self.K_dim, tma_g2s, self.num_stages) - V_args = self._make_bf16_tma_args( - v_new_chunks, self.V_dim, tma_g2s, self.num_stages + Q_args = make_chunk_tma_args(q, self.K_dim, tma_g2s, self.num_stages, self.BT) + K_args = make_chunk_tma_args(k, self.K_dim, tma_g2s, self.num_stages, self.BT) + V_args = make_chunk_tma_args( + v_new_chunks, self.V_dim, tma_g2s, self.num_stages, self.BT ) - H_args = self._make_h_tma_args(h, tma_g2s, self.num_stages) - O_args = self._make_bf16_tma_args(o, self.V_dim, tma_s2g, 1) + H_args = make_output_state_tma_args( + h, tma_g2s, self.num_stages, self.K_dim, self.V_dim + ) + O_args = make_chunk_tma_args(o, self.V_dim, tma_s2g, 1, self.BT) self.kernel( Q_args, K_args, diff --git a/python/sglang/kernels/ops/attention/linear/kda_blackwell/kernel_h.py b/python/sglang/kernels/ops/attention/linear/kda_blackwell/kernel_h.py index 0e426c0c1..37eff60d5 100644 --- a/python/sglang/kernels/ops/attention/linear/kda_blackwell/kernel_h.py +++ b/python/sglang/kernels/ops/attention/linear/kda_blackwell/kernel_h.py @@ -33,6 +33,10 @@ from sglang.kernels.ops.attention.cute_utils import ( fence_before_tma_store, simple_tma_copy, ) +from sglang.kernels.ops.attention.linear.tma import ( + make_chunk_tma_args, + make_recurrent_state_tma_args, +) class Sm100KdaChunkHKernel: @@ -60,45 +64,6 @@ class Sm100KdaChunkHKernel: self.num_stages = num_stages self.num_warps = 10 - @cute.jit - def _make_bf16_tma_args( - self, - tensor: cute.Tensor, - dim: cutlass.Constexpr[int], - op: cpasync.TmaCopyOp, - stages: cutlass.Constexpr[int], - ): - swizzle_128B = cute.make_swizzle(3, 4, 3) - slayout = cute.make_layout( - (self.BT, 1, (64, dim // 64), stages), - stride=(64, 0, (1, self.BT * 64), self.BT * dim), - ) - slayout = cute.make_composed_layout(swizzle_128B, 0, slayout) - atom, tma_tensor = cpasync.make_tiled_tma_atom( - op, - cute.logical_divide(tensor, (None, None, 64)), - slayout, - cta_tiler=(self.BT, 1, dim), - ) - return atom, tma_tensor, slayout - - @cute.jit - def _make_h_tma_args(self, tensor: cute.Tensor, op: cpasync.TmaCopyOp): - num_elems = 128 // (tensor.element_type.width // 8) - swizzle_128B = cute.make_swizzle(3, 4, 3) - slayout = cute.make_layout( - (1, 1, self.V_dim, (num_elems, self.K_dim // num_elems)), - stride=(0, 0, num_elems, (1, self.V_dim * num_elems)), - ) - slayout = cute.make_composed_layout(swizzle_128B, 0, slayout) - atom, tma_tensor = cpasync.make_tiled_tma_atom( - op, - cute.logical_divide(tensor, (None, None, None, num_elems)), - slayout, - cta_tiler=(1, 1, self.V_dim, self.K_dim), - ) - return atom, tma_tensor, slayout - @cute.jit def __call__( self, @@ -118,13 +83,13 @@ class Sm100KdaChunkHKernel: tma_g2s = cpasync.CopyBulkTensorTileG2SOp() tma_s2g = cpasync.CopyBulkTensorTileS2GOp() - K_args = self._make_bf16_tma_args(K, self.K_dim, tma_g2s, self.num_stages) - V_args = self._make_bf16_tma_args(V, self.V_dim, tma_g2s, self.num_stages) - W_args = self._make_bf16_tma_args(W, self.K_dim, tma_g2s, self.num_stages) - V_new_args = self._make_bf16_tma_args(V_new, self.V_dim, tma_s2g, 1) - H0_args = self._make_h_tma_args(h0, tma_g2s) - HT_args = self._make_h_tma_args(ht, tma_s2g) - H_args = self._make_h_tma_args(h, tma_s2g) + K_args = make_chunk_tma_args(K, self.K_dim, tma_g2s, self.num_stages, self.BT) + V_args = make_chunk_tma_args(V, self.V_dim, tma_g2s, self.num_stages, self.BT) + W_args = make_chunk_tma_args(W, self.K_dim, tma_g2s, self.num_stages, self.BT) + V_new_args = make_chunk_tma_args(V_new, self.V_dim, tma_s2g, 1, self.BT) + H0_args = make_recurrent_state_tma_args(h0, tma_g2s, self.K_dim, self.V_dim) + HT_args = make_recurrent_state_tma_args(ht, tma_s2g, self.K_dim, self.V_dim) + H_args = make_recurrent_state_tma_args(h, tma_s2g, self.K_dim, self.V_dim) # h0/ht may be the full state pool ([num_slots, ...]) rather than a # per-sequence gather, so the sequence count comes from cu_seqlens and diff --git a/python/sglang/kernels/ops/attention/linear/kda_blackwell/kernel_kkt_inv_uw.py b/python/sglang/kernels/ops/attention/linear/kda_blackwell/kernel_kkt_inv_uw.py index 68bb82062..2a7262357 100644 --- a/python/sglang/kernels/ops/attention/linear/kda_blackwell/kernel_kkt_inv_uw.py +++ b/python/sglang/kernels/ops/attention/linear/kda_blackwell/kernel_kkt_inv_uw.py @@ -33,6 +33,9 @@ from sglang.kernels.ops.attention.cute_utils import ( mma_bf16, simple_tma_copy, ) +from sglang.kernels.ops.attention.linear.tma import ( + make_chunk_tma_args, +) class Sm100KdaChunkUWKernel: @@ -57,28 +60,6 @@ class Sm100KdaChunkUWKernel: self.BT = 64 self.num_warps = 2 + 4 + 4 - @cute.jit - def _make_tma_args( - self, - tensor: cute.Tensor, - dim: cutlass.Constexpr[int], - num_stages: int, - op: cpasync.TmaCopyOp, - ): - swizzle_128B = cute.make_swizzle(3, 4, 3) - slayout = cute.make_layout( - (self.BT, 1, (64, dim // 64), num_stages), - stride=(64, 0, (1, self.BT * 64), self.BT * dim), - ) - slayout = cute.make_composed_layout(swizzle_128B, 0, slayout) - atom, tma_tensor = cpasync.make_tiled_tma_atom( - op, - cute.logical_divide(tensor, (None, None, 64)), - slayout, - cta_tiler=(self.BT, 1, dim), - ) - return atom, tma_tensor, slayout - @cute.jit def __call__( self, @@ -98,12 +79,12 @@ class Sm100KdaChunkUWKernel: tma_g2s = cpasync.CopyBulkTensorTileG2SOp() tma_s2g = cpasync.CopyBulkTensorTileS2GOp() - KL_args = self._make_tma_args(KL, self.K_dim, self.num_stages, tma_g2s) - KR_args = self._make_tma_args(KR, self.K_dim, self.num_stages, tma_g2s) - KG_args = self._make_tma_args(KG, self.K_dim, self.num_stages, tma_g2s) - V_args = self._make_tma_args(V, self.V_dim, self.num_stages, tma_g2s) - U_args = self._make_tma_args(U, self.V_dim, 1, tma_s2g) - W_args = self._make_tma_args(W, self.K_dim, 1, tma_s2g) + KL_args = make_chunk_tma_args(KL, self.K_dim, tma_g2s, self.num_stages, self.BT) + KR_args = make_chunk_tma_args(KR, self.K_dim, tma_g2s, self.num_stages, self.BT) + KG_args = make_chunk_tma_args(KG, self.K_dim, tma_g2s, self.num_stages, self.BT) + V_args = make_chunk_tma_args(V, self.V_dim, tma_g2s, self.num_stages, self.BT) + U_args = make_chunk_tma_args(U, self.V_dim, tma_s2g, 1, self.BT) + W_args = make_chunk_tma_args(W, self.K_dim, tma_s2g, 1, self.BT) grid = (num_sms // self.Hv, self.Hv, 1) block = (self.num_warps * 32, 1, 1) diff --git a/python/sglang/kernels/ops/attention/linear/kda_blackwell/kernel_o.py b/python/sglang/kernels/ops/attention/linear/kda_blackwell/kernel_o.py index f78c95690..3daf4328c 100644 --- a/python/sglang/kernels/ops/attention/linear/kda_blackwell/kernel_o.py +++ b/python/sglang/kernels/ops/attention/linear/kda_blackwell/kernel_o.py @@ -32,6 +32,10 @@ from sglang.kernels.ops.attention.cute_utils import ( fence_before_tma_store, simple_tma_copy, ) +from sglang.kernels.ops.attention.linear.tma import ( + make_chunk_tma_args, + make_output_state_tma_args, +) class Sm100KdaChunkOKernel: @@ -58,50 +62,6 @@ class Sm100KdaChunkOKernel: self.num_stages = num_stages self.num_warps = 10 - @cute.jit - def _make_bf16_tma_args( - self, - tensor: cute.Tensor, - dim: cutlass.Constexpr[int], - op: cpasync.TmaCopyOp, - stages: cutlass.Constexpr[int], - ): - swizzle_128B = cute.make_swizzle(3, 4, 3) - slayout = cute.make_layout( - (self.BT, 1, (64, dim // 64), stages), - stride=(64, 0, (1, self.BT * 64), self.BT * dim), - ) - slayout = cute.make_composed_layout(swizzle_128B, 0, slayout) - atom, tma_tensor = cpasync.make_tiled_tma_atom( - op, - cute.logical_divide(tensor, (None, None, 64)), - slayout, - cta_tiler=(self.BT, 1, dim), - ) - return atom, tma_tensor, slayout - - @cute.jit - def _make_h_tma_args( - self, - tensor: cute.Tensor, - op: cpasync.TmaCopyOp, - stages: cutlass.Constexpr[int], - ): - num_elems = 128 // (tensor.element_type.width // 8) - swizzle_128B = cute.make_swizzle(3, 4, 3) - slayout = cute.make_layout( - (1, self.V_dim, (num_elems, self.K_dim // num_elems), stages), - stride=(0, num_elems, (1, self.V_dim * num_elems), self.V_dim * self.K_dim), - ) - slayout = cute.make_composed_layout(swizzle_128B, 0, slayout) - atom, tma_tensor = cpasync.make_tiled_tma_atom( - op, - cute.logical_divide(tensor, (None, None, num_elems)), - slayout, - cta_tiler=(1, self.V_dim, self.K_dim), - ) - return atom, tma_tensor, slayout - @cute.jit def __call__( self, @@ -121,14 +81,16 @@ class Sm100KdaChunkOKernel: block = (self.num_warps * 32, 1, 1) tma_g2s = cpasync.CopyBulkTensorTileG2SOp() tma_s2g = cpasync.CopyBulkTensorTileS2GOp() - Q_args = self._make_bf16_tma_args(qg2, self.K_dim, tma_g2s, self.num_stages) - Q2_args = self._make_bf16_tma_args(qg, self.K_dim, tma_g2s, self.num_stages) - K_args = self._make_bf16_tma_args(kg, self.K_dim, tma_g2s, self.num_stages) - V_args = self._make_bf16_tma_args( - v_new_chunks, self.V_dim, tma_g2s, self.num_stages + Q_args = make_chunk_tma_args(qg2, self.K_dim, tma_g2s, self.num_stages, self.BT) + Q2_args = make_chunk_tma_args(qg, self.K_dim, tma_g2s, self.num_stages, self.BT) + K_args = make_chunk_tma_args(kg, self.K_dim, tma_g2s, self.num_stages, self.BT) + V_args = make_chunk_tma_args( + v_new_chunks, self.V_dim, tma_g2s, self.num_stages, self.BT ) - H_args = self._make_h_tma_args(h, tma_g2s, self.num_stages) - O_args = self._make_bf16_tma_args(o, self.V_dim, tma_s2g, 1) + H_args = make_output_state_tma_args( + h, tma_g2s, self.num_stages, self.K_dim, self.V_dim + ) + O_args = make_chunk_tma_args(o, self.V_dim, tma_s2g, 1, self.BT) self.kernel( Q_args, Q2_args, diff --git a/python/sglang/kernels/ops/attention/linear/kda_nvidia_prefill/chunk_fwd.py b/python/sglang/kernels/ops/attention/linear/kda_nvidia_prefill/chunk_fwd.py index 480b3b621..187d088df 100644 --- a/python/sglang/kernels/ops/attention/linear/kda_nvidia_prefill/chunk_fwd.py +++ b/python/sglang/kernels/ops/attention/linear/kda_nvidia_prefill/chunk_fwd.py @@ -341,111 +341,6 @@ def _get_padded_input_buffers( return e -# Multi-seq varlen repack cache for the Phase 2.2 path. Keyed by id(cu_seqlens). -# Stores: (orig_seq_lens, padded_seq_lens, new_cu_seqlens_tensor, -# new_chunk_indices_tensor, new_T_total, padded_input_buffers). -# All tensors are GPU-side and pre-allocated at cache build time. Per-call the -# kernel reads from / writes to these buffers; we copy caller's input slices in -# and output slices back (only the valid prefix of each seq). -_multiseq_repack_cache = {} - -_caller_layout_O_cache = {} - - -def _get_caller_layout_O_buffer(multiseq_info, dtype, V_dim, device): - """Per-shape cached output buffer at caller's contiguous layout (sum of - seq_lens, no padding gaps). Filled by per-seq copies from K4's padded O.""" - caller_T = multiseq_info["caller_T"] - H_x_V = multiseq_info.get("_H_V") # not strictly needed since we fix B=1 H known - key = (caller_T, V_dim, dtype, device.index if device.index is not None else 0) - e = _caller_layout_O_cache.get(key) - if e is None: - # B=1 enforced upstream when multiseq_info is built. - e = torch.empty( - 1, - caller_T, - multiseq_info["q_pad"].shape[2], - V_dim, - dtype=dtype, - device=device, - ) - _caller_layout_O_cache[key] = e - return e - - -def _get_multiseq_repack_info(cu_seqlens, q, k, v, g, beta, BT, device): - """Build (and cache) the padded layout for multi-seq varlen with non-aligned - seqs. Returns None if all seqs are already 64-aligned (caller can use the - existing varlen_pure path).""" - import weakref - - key = id(cu_seqlens) - cached = _multiseq_repack_cache.get(key) - if cached is not None: - wref, e = cached - if wref() is cu_seqlens: - return e - # id collision after GC: rebuild - del _multiseq_repack_cache[key] - cu_cpu = cu_seqlens.cpu().tolist() - seq_lens = [cu_cpu[i + 1] - cu_cpu[i] for i in range(len(cu_cpu) - 1)] - if all(sl % BT == 0 for sl in seq_lens): - _multiseq_repack_cache[key] = (weakref.ref(cu_seqlens), None) - return None - padded_lens = [((sl + BT - 1) // BT) * BT for sl in seq_lens] - new_cu = [0] - for pl in padded_lens: - new_cu.append(new_cu[-1] + pl) - new_T_total = new_cu[-1] - B = q.shape[0] - H = q.shape[2] - K = q.shape[3] - # Pre-allocated padded input buffers. q/k/v/beta tail = 0 (zero MMA), g - # tail = -1e3 sentinel (zero gate activation). The PER-SEQ tail regions - # are between (new_cu[i] + seq_lens[i], new_cu[i+1]) — pre-fill once. - q_pad = torch.zeros(B, new_T_total, H, K, dtype=q.dtype, device=device) - k_pad = torch.zeros_like(q_pad) - v_pad = torch.zeros(B, new_T_total, H, v.shape[3], dtype=v.dtype, device=device) - beta_pad = torch.zeros(B, new_T_total, H, dtype=beta.dtype, device=device) - g_pad = torch.zeros(B, new_T_total, H, K, dtype=g.dtype, device=device) - for i, (sl, pl) in enumerate(zip(seq_lens, padded_lens)): - if sl < pl: - tail_start = new_cu[i] + sl - tail_end = new_cu[i + 1] - g_pad[:, tail_start:tail_end] = -1000.0 - new_cu_tensor = torch.tensor(new_cu, dtype=cu_seqlens.dtype, device=device) - new_chunk_indices = prepare_chunk_indices(new_cu_tensor, BT) - # Build index map: dst_indices[i] = position in padded layout where orig - # row i lives. Used by index_copy_ to do the scatter in one op (instead of - # N_seqs × 5 separate slice copies, which cost ~5us each in Python). - T_total_orig = cu_cpu[-1] - dst_indices_list = [] - for i, sl in enumerate(seq_lens): - for j in range(sl): - dst_indices_list.append(new_cu[i] + j) - dst_indices = torch.tensor(dst_indices_list, dtype=torch.long, device=device) - # Mark this cu_seqlens as VARLEN_PURE eligible — every seq in the new - # layout is 64-aligned by construction. - _varlen_pure_cache[id(new_cu_tensor)] = True - e = { - "seq_lens": seq_lens, - "padded_lens": padded_lens, - "new_cu": new_cu, - "new_T_total": new_T_total, - "new_cu_tensor": new_cu_tensor, - "new_chunk_indices": new_chunk_indices, - "q_pad": q_pad, - "k_pad": k_pad, - "v_pad": v_pad, - "g_pad": g_pad, - "beta_pad": beta_pad, - "dst_indices": dst_indices, - "T_total_orig": T_total_orig, - } - _multiseq_repack_cache[key] = (weakref.ref(cu_seqlens), e) - return e - - def _get_buffers(dev, dtype_k, B, T, H, K_dim, V_dim, NT, N_seqs, BT): """All beta fusion lives in akk_inv kernel epilogue (post-inv column-scale).""" key = (dev.index or 0, B, T, H, K_dim, V_dim, NT, N_seqs) diff --git a/python/sglang/kernels/ops/attention/linear/tma.py b/python/sglang/kernels/ops/attention/linear/tma.py new file mode 100644 index 000000000..fb4479fce --- /dev/null +++ b/python/sglang/kernels/ops/attention/linear/tma.py @@ -0,0 +1,75 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Shared TMA layouts for the Blackwell GDN and KDA chunk kernels.""" + +import cutlass +from cutlass import cute +from cutlass.cute.nvgpu import cpasync + + +@cute.jit +def make_chunk_tma_args( + tensor: cute.Tensor, + dim: cutlass.Constexpr[int], + op: cpasync.TmaCopyOp, + stages: cutlass.Constexpr[int], + chunk_size: cutlass.Constexpr[int], +): + # Divide the contiguous dimension so the descriptor issues one 4D TMA. + slayout = cute.make_layout( + (chunk_size, 1, (64, dim // 64), stages), + stride=(64, 0, (1, chunk_size * 64), chunk_size * dim), + ) + slayout = cute.make_composed_layout(cute.make_swizzle(3, 4, 3), 0, slayout) + atom, tma_tensor = cpasync.make_tiled_tma_atom( + op, + cute.logical_divide(tensor, (None, None, 64)), + slayout, + cta_tiler=(chunk_size, 1, dim), + ) + return atom, tma_tensor, slayout + + +@cute.jit +def make_recurrent_state_tma_args( + tensor: cute.Tensor, + op: cpasync.TmaCopyOp, + key_dim: cutlass.Constexpr[int], + value_dim: cutlass.Constexpr[int], +): + num_elems = 128 // (tensor.element_type.width // 8) + slayout = cute.make_layout( + (1, 1, value_dim, (num_elems, key_dim // num_elems)), + stride=(0, 0, num_elems, (1, value_dim * num_elems)), + ) + slayout = cute.make_composed_layout(cute.make_swizzle(3, 4, 3), 0, slayout) + atom, tma_tensor = cpasync.make_tiled_tma_atom( + op, + cute.logical_divide(tensor, (None, None, None, num_elems)), + slayout, + cta_tiler=(1, 1, value_dim, key_dim), + ) + return atom, tma_tensor, slayout + + +@cute.jit +def make_output_state_tma_args( + tensor: cute.Tensor, + op: cpasync.TmaCopyOp, + stages: cutlass.Constexpr[int], + key_dim: cutlass.Constexpr[int], + value_dim: cutlass.Constexpr[int], +): + num_elems = 128 // (tensor.element_type.width // 8) + slayout = cute.make_layout( + (1, value_dim, (num_elems, key_dim // num_elems), stages), + stride=(0, num_elems, (1, value_dim * num_elems), value_dim * key_dim), + ) + slayout = cute.make_composed_layout(cute.make_swizzle(3, 4, 3), 0, slayout) + atom, tma_tensor = cpasync.make_tiled_tma_atom( + op, + cute.logical_divide(tensor, (None, None, num_elems)), + slayout, + cta_tiler=(1, value_dim, key_dim), + ) + return atom, tma_tensor, slayout diff --git a/python/sglang/kernels/ops/gemm/trtllm_lora_temp/kernel_utils.py b/python/sglang/kernels/ops/gemm/trtllm_lora_temp/kernel_utils.py index a92fcd822..2ef951338 100644 --- a/python/sglang/kernels/ops/gemm/trtllm_lora_temp/kernel_utils.py +++ b/python/sglang/kernels/ops/gemm/trtllm_lora_temp/kernel_utils.py @@ -1,7 +1,7 @@ -import triton -import triton.language as tl - from sglang.kernels.jit.utils import is_arch_support_pdl +from sglang.kernels.ops.gemm.kernel_utils import ( + _resolve_token_positions as _resolve_token_positions, +) def get_pdl_launch_metadata() -> tuple[bool, dict]: @@ -12,20 +12,3 @@ def get_pdl_launch_metadata() -> tuple[bool, dict]: """ enable_pdl = is_arch_support_pdl() return enable_pdl, ({"launch_pdl": True} if enable_pdl else {}) - - -@triton.jit -def _resolve_token_positions( - sorted_token_ids, seg_start, s_offset, seg_len, SORTED_BY_ADAPTER: tl.constexpr -): - """Map logical segment offsets to physical token positions. - - When SORTED_BY_ADAPTER is True, segments are grouped by adapter and - sorted_token_ids provides the indirection to the original token rows. - When False, tokens are already contiguous starting at seg_start. - """ - if SORTED_BY_ADAPTER: - return tl.load( - sorted_token_ids + seg_start + s_offset, mask=s_offset < seg_len - ).to(tl.int64) - return (seg_start + s_offset).to(tl.int64) diff --git a/python/sglang/kernels/ops/kimi_k3/activation.py b/python/sglang/kernels/ops/kimi_k3/activation.py index 8a942f04b..8d77c0b3e 100644 --- a/python/sglang/kernels/ops/kimi_k3/activation.py +++ b/python/sglang/kernels/ops/kimi_k3/activation.py @@ -6,9 +6,8 @@ import torch from sglang.kernels.jit.utils import ( cache_once, - get_jit_cuda_arch, + get_activation_cuda_cflags, is_arch_support_pdl, - is_hip_runtime, load_jit, make_cpp_args, ) @@ -21,16 +20,6 @@ def _make_name(*args): return "kimi_k3_" + "_".join(str(a) for a in args) -def _fast_math_flags() -> list[str]: - # Mirrors sgl-kernel's CMake policy: fast-math on SM90, precise on - # SM100+ (Blackwell needs bit-exact expf), off on HIP (clang rejects). - if is_hip_runtime(): - return [] - if get_jit_cuda_arch().major >= 10: - return [] - return ["--use_fast_math"] - - @cache_once def _jit_situ_and_mul_module(in_dtype: torch.dtype, out_dtype: torch.dtype) -> Module: """Compile and cache the JIT SiTU-and-mul module for an (in, out) dtype pair.""" @@ -40,7 +29,7 @@ def _jit_situ_and_mul_module(in_dtype: torch.dtype, out_dtype: torch.dtype) -> M *args, cuda_files=["kimi_k3/situ_and_mul.cuh"], cuda_wrappers=[("run", f"SituAndMulKernel<{args}>::run")], - extra_cuda_cflags=_fast_math_flags(), + extra_cuda_cflags=get_activation_cuda_cflags(), ) diff --git a/python/sglang/kernels/ops/moe/trtllm_lora_temp/virtual_experts.py b/python/sglang/kernels/ops/moe/trtllm_lora_temp/virtual_experts.py index 4ec309dd5..e3a733206 100644 --- a/python/sglang/kernels/ops/moe/trtllm_lora_temp/virtual_experts.py +++ b/python/sglang/kernels/ops/moe/trtllm_lora_temp/virtual_experts.py @@ -12,9 +12,7 @@ import triton.language as tl from sglang.kernels.ops.gemm.trtllm_lora_temp.kernel_utils import ( get_pdl_launch_metadata, ) -from sglang.kernels.ops.moe.moe_align import ( - moe_align_block_size as jit_moe_align_block_size, -) +from sglang.kernels.ops.moe.virtual_experts import _align_block_size_large from sglang.srt.lora.trtllm_lora_temp.environ import lora_envs @@ -404,195 +402,6 @@ from sglang.srt.lora.trtllm_lora_temp.specialized_expand import ( # noqa: E402, ) -def _align_block_size_jit( - topk_ids: torch.Tensor, - block_size: int, - num_experts: int, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """CUDA JIT alignment for up to 8191 experts. - - Expert IDs are shifted by one so ``-1`` maps to a sentinel bucket. The - fused allocation stays int4-aligned for the kernel's vectorized clear. - """ - assert num_experts <= 8191, ( - f"_align_block_size_jit supports at most 8191 experts " - f"(num_moe_experts * max_loras), got {num_experts}" - ) - - device = topk_ids.device - flat_topk_ids = topk_ids.reshape(-1) - if flat_topk_ids.dtype == torch.int64: - flat_topk_ids = flat_topk_ids.to(torch.int32) - num_total_tokens = flat_topk_ids.numel() - - if num_total_tokens == 0: - empty = torch.empty(0, dtype=torch.int32, device=device) - return empty, empty, torch.zeros(1, dtype=torch.int32, device=device) - - # JIT kernel uses +1 offset convention: -1 -> bucket 0 (sentinel), - # expert i -> bucket i+1. So pass num_experts + 1 as the bucket count. - jit_num_experts = num_experts + 1 - - if num_total_tokens < jit_num_experts: - max_num_tokens_padded = num_total_tokens * block_size - else: - max_num_tokens_padded = num_total_tokens + jit_num_experts * (block_size - 1) - - # Align every sub-buffer offset to a multiple of 4 (VEC_SIZE). The CUDA - # kernel fills sorted_token_ids with vectorized int4 writes whose last - # store can spill up to 3 int32s past the logical end. With a fused - # allocation the spill would corrupt the adjacent sub-buffer. - _A4 = lambda n: (n + 3) & ~3 # noqa: E731 - max_num_tokens_padded = _A4(max_num_tokens_padded) - max_num_m_blocks = (max_num_tokens_padded + block_size - 1) // block_size - max_num_m_blocks_padded = _A4(max_num_m_blocks) - num_post_pad_size = _A4(1) # 1 element, padded to 4 - cumsum_size = _A4(jit_num_experts + 1) - - # Single allocation sliced into 4 views (zero-copy) to avoid - # per-call Python overhead of 4 separate torch.empty calls. - total_buf = ( - max_num_tokens_padded - + max_num_m_blocks_padded - + num_post_pad_size - + cumsum_size - ) - buf = torch.empty(total_buf, dtype=torch.int32, device=device) - off = 0 - sorted_token_ids = buf[off : off + max_num_tokens_padded] - off += max_num_tokens_padded - expert_ids = buf[off : off + max_num_m_blocks] - off += max_num_m_blocks_padded - num_tokens_post_padded = buf[off : off + 1] - off += num_post_pad_size - cumsum_buffer = buf[off : off + jit_num_experts + 1] - - jit_moe_align_block_size( - flat_topk_ids, - jit_num_experts, - block_size, - sorted_token_ids, - expert_ids, - num_tokens_post_padded, - cumsum_buffer, - True, # pad_sorted_token_ids - ) - - return sorted_token_ids, expert_ids, num_tokens_post_padded - - -@torch.compile(dynamic=True) -def _align_block_size_torch( - topk_ids: torch.Tensor, - block_size: int, - num_experts: int, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Pure-PyTorch align_block_size for num_experts > 1024, compiled via torch.compile. - - Fallback for platforms where the CUDA JIT kernel is unavailable (e.g. AMD/ROCm). - - Out-of-range topk_ids (negative sentinels left by EP dispatch, or virtual- - expert IDs >= num_experts produced when those sentinels are combined with - a per-adapter offset) are routed into a dedicated sentinel bucket. Without - this, indexing ``padded_offsets[sorted_expert_ids]`` would wrap (-1) or - OOB-read, and the bad expert ids would propagate into the downstream LoRA - GEMM as real expert slots. - """ - device = topk_ids.device - flat_topk_ids = topk_ids.reshape(-1).to(torch.int64) - num_total_tokens = flat_topk_ids.numel() - - sentinel = num_experts - valid_mask = (flat_topk_ids >= 0) & (flat_topk_ids < num_experts) - safe_topk_ids = torch.where( - valid_mask, - flat_topk_ids, - torch.full_like(flat_topk_ids, sentinel), - ) - - bucket_count = num_experts + 1 - max_total_padded_tokens = ( - (num_total_tokens + bucket_count * (block_size - 1) + block_size - 1) - // block_size - ) * block_size - max_num_blocks = max_total_padded_tokens // block_size - - sorted_token_ids = torch.full( - (max_total_padded_tokens,), - num_total_tokens, - dtype=torch.int32, - device=device, - ) - expert_ids = torch.full( - (max_num_blocks,), - -1, - dtype=torch.int32, - device=device, - ) - - if num_total_tokens == 0: - num_tokens_post_padded = torch.zeros((1,), dtype=torch.int32, device=device) - return sorted_token_ids, expert_ids, num_tokens_post_padded - - sorted_order = torch.argsort(safe_topk_ids) - sorted_expert_ids = safe_topk_ids[sorted_order] - expert_range = torch.arange(bucket_count, device=device, dtype=torch.int64) - counts_offsets = torch.searchsorted(sorted_expert_ids, expert_range, right=False) - counts_end = torch.searchsorted(sorted_expert_ids, expert_range, right=True) - counts = counts_end - counts_offsets - padded_counts = ((counts + block_size - 1) // block_size) * block_size - total_padded_tokens = padded_counts.sum().to(torch.int32).reshape(1) - padded_offsets = torch.cumsum(padded_counts, dim=0) - padded_counts - - token_ranks = ( - torch.arange(num_total_tokens, device=device, dtype=torch.int64) - - counts_offsets[sorted_expert_ids] - ) - output_positions = padded_offsets[sorted_expert_ids] + token_ranks - sorted_token_ids.scatter_( - 0, - output_positions.to(torch.int64), - sorted_order.to(torch.int32), - ) - - block_counts = padded_counts // block_size - real_block_counts = block_counts.clone() - real_block_counts[sentinel] = 0 - actual_num_blocks = real_block_counts.sum() - - if max_num_blocks <= 0: - return sorted_token_ids, expert_ids, total_padded_tokens - - block_offsets = torch.cumsum(real_block_counts, dim=0) - all_block_positions = torch.arange(max_num_blocks, device=device, dtype=torch.int64) - assigned_experts = torch.searchsorted( - block_offsets, all_block_positions, right=True - ).to(torch.int32) - expert_ids.copy_( - torch.where( - all_block_positions < actual_num_blocks, - assigned_experts, - torch.full_like(assigned_experts, -1), - ) - ) - - return sorted_token_ids, expert_ids, total_padded_tokens - - -def _align_block_size_large( - topk_ids: torch.Tensor, - block_size: int, - num_experts: int, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Dispatch to the CUDA JIT kernel when available, otherwise fall back to - the pure-PyTorch torch.compile path (needed on AMD/ROCm or when the JIT - module fails to load).""" - try: - return _align_block_size_jit(topk_ids, block_size, num_experts) - except Exception: - return _align_block_size_torch(topk_ids, block_size, num_experts) - - def _merged_experts_fused_moe_lora_add_fake( output: torch.Tensor, hidden_states: torch.Tensor, diff --git a/python/sglang/srt/hardware_backend/npu/moe/topk.py b/python/sglang/srt/hardware_backend/npu/moe/topk.py index 75085959b..50b90aeb8 100644 --- a/python/sglang/srt/hardware_backend/npu/moe/topk.py +++ b/python/sglang/srt/hardware_backend/npu/moe/topk.py @@ -16,20 +16,6 @@ if TYPE_CHECKING: from sglang.srt.layers.moe.topk import TopKConfig, TopKOutput -def _apply_routed_scaling_after_renorm( - topk_weights: torch.Tensor, - topk_config: "TopKConfig", -) -> torch.Tensor: - """Mirror GPU post-renorm scaling when apply_routed_scaling_factor_on_output is set.""" - if ( - topk_config.renormalize - and topk_config.apply_routed_scaling_factor_on_output - and topk_config.routed_scaling_factor is not None - ): - return topk_weights * topk_config.routed_scaling_factor - return topk_weights - - def fused_topk_npu( hidden_states: torch.Tensor, router_logits: torch.Tensor, @@ -38,7 +24,6 @@ def fused_topk_npu( expert_location_dispatch_info: Optional["ExpertLocationDispatchInfo"] = None, layer_id: Optional[int] = None, ) -> "TopKOutput": - use_grouped_topk = topk_config.use_grouped_topk renormalize = topk_config.renormalize correction_bias = topk_config.correction_bias diff --git a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py index 34c2d6cc7..3afb09f3f 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py @@ -50,83 +50,6 @@ def _extract_positions_from_plan( return positions -def _compress_forward_c128_fallback( - kv_score_buffer: torch.Tensor, - kv_score_input: torch.Tensor, - ape: torch.Tensor, - plan: Union[CompressorDecodePlan, CompressorPrefillPlan], - head_dim: int, -) -> torch.Tensor: - """PyTorch fallback for C128 compress_forward on HIP (wave64). - - Fully vectorized, compatible with CUDA graph capture. - kv_score_buffer: [num_pages, 128, head_dim * 2] - ape: [128, head_dim] - - IMPORTANT: This also performs the write to state buffer (like the JIT kernel). - The JIT kernel does: (1) write kv_score_input to buffer, (2) compress from buffer. - """ - num_total_slots = kv_score_buffer.shape[0] * kv_score_buffer.shape[1] - num_pages = kv_score_buffer.shape[0] - last_dim = kv_score_buffer.shape[-1] - - # Step 1: WRITE kv_score_input to state buffer - if num_total_slots > 0: - buf_flat = kv_score_buffer.view(-1, last_dim) - if plan.is_decode: - # Decode: plan_d has write_loc per batch item - plan_raw = plan[1].view(torch.int32) # [bs, 4] - write_locs = plan_raw[:, 1].long() - # Only write valid locations (>= 0 and < buffer size) - valid_write = (write_locs >= 0) & (write_locs < num_total_slots) - if valid_write.any(): - buf_flat[write_locs[valid_write]] = kv_score_input[valid_write] - else: - # Prefill: plan_w has {ragged_id, write_loc} per write entry - plan_w = plan[2] # [num_w, 8] uint8 = WritePlan - if plan_w.shape[0] > 0: - plan_w_raw = plan_w.view(torch.int32) # [num_w, 2] - ragged_ids = plan_w_raw[:, 0].long() & 0xFFFF - write_locs = plan_w_raw[:, 1].long() - valid_write = (write_locs >= 0) & (write_locs < num_total_slots) - ragged_ids_safe = ragged_ids.clamp( - min=0, max=kv_score_input.shape[0] - 1 - ) - if valid_write.any(): - buf_flat[write_locs[valid_write]] = kv_score_input[ - ragged_ids_safe[valid_write] - ] - - # Step 2: COMPRESS (read from buffer page and do softmax-pool) - plan_c = plan[1] # plan_d for decode, plan_c for prefill - num_tokens = plan_c.shape[0] - if num_pages == 0 or num_tokens == 0: - return kv_score_input.new_zeros(num_tokens, head_dim) - - plan_c_raw = plan_c.view(torch.int32) # [N, 4] - read_page_0 = plan_c_raw[:, 2].long() - # Use torch.where instead of clamp to handle -1 (invalid) gracefully - valid_read = (read_page_0 >= 0) & (read_page_0 < num_pages) - read_page_0_safe = torch.where( - valid_read, read_page_0, torch.zeros_like(read_page_0) - ) - - gathered = kv_score_buffer[read_page_0_safe] # [N, 128, head_dim*2] - kv = gathered[:, :, :head_dim].float() - score = gathered[:, :, head_dim:].float() + ape.float().unsqueeze(0) - weights = score.softmax(dim=1) - out = (weights * kv).sum(dim=1) - - # For decode: zero out non-boundary tokens (seq_len % 128 != 0) - # so they don't corrupt kvcache location 0 when stored. - if plan.is_decode: - seq_lens = plan_c_raw[:, 0].to(torch.int32) - is_boundary = (seq_lens % 128 == 0).unsqueeze(-1) # [N, 1] - out = torch.where(is_boundary, out, torch.zeros_like(out)) - - return out.to(kv_score_input.dtype) - - class CompressorBackendMixin: def __init__(self): super().__init__() diff --git a/python/sglang/srt/lora/marlin_lora_temp/moe_runner.py b/python/sglang/srt/lora/marlin_lora_temp/moe_runner.py index 5e59c03f4..8b00244d7 100644 --- a/python/sglang/srt/lora/marlin_lora_temp/moe_runner.py +++ b/python/sglang/srt/lora/marlin_lora_temp/moe_runner.py @@ -28,7 +28,7 @@ if _is_cuda: moe_sum_reduce_triton, ) from sglang.kernels.ops.moe.moe_wna16_marlin import moe_wna16_marlin_gemm - from sglang.kernels.ops.moe.trtllm_lora_temp.virtual_experts import ( + from sglang.kernels.ops.moe.virtual_experts import ( _align_block_size_jit as moe_align_block_size, ) from sglang.srt.layers.moe.fused_moe_triton.fused_marlin_moe import ( diff --git a/python/sglang/srt/multimodal/processors/mimo_v2.py b/python/sglang/srt/multimodal/processors/mimo_v2.py index c833d7937..e914f00b1 100644 --- a/python/sglang/srt/multimodal/processors/mimo_v2.py +++ b/python/sglang/srt/multimodal/processors/mimo_v2.py @@ -3,11 +3,9 @@ import asyncio import base64 import copy -import json import math import os import re -import subprocess from concurrent.futures import ThreadPoolExecutor, as_completed from dataclasses import dataclass, field from io import BytesIO @@ -231,39 +229,6 @@ def _decode_frames_and_timestamps(vdw, ele): return video_tensor, timestamps -def _ffprobe_has_audio(src, stdin=None, label=None) -> bool: - # Header-only audio-stream probe for HTTP URLs; avoids full download. - try: - r = subprocess.run( - [ - "ffprobe", - "-v", - "quiet", - "-print_format", - "json", - "-show_streams", - "-select_streams", - "a", - src, - ], - input=stdin, - capture_output=True, - timeout=30, - ) - if r.returncode != 0: - stderr = r.stderr.decode("utf-8", errors="replace") - raise RuntimeError(f"ffprobe failed for {label}: {stderr}") - return bool(json.loads(r.stdout).get("streams")) - except subprocess.TimeoutExpired: - logger.error("ffprobe timed out for %s", label) - raise - except FileNotFoundError as e: - raise RuntimeError("ffprobe not found; install ffmpeg") from e - except json.JSONDecodeError: - logger.error("ffprobe returned invalid JSON for %s", label) - raise - - class MiMoProcessor: def __init__( self, @@ -690,7 +655,6 @@ class MiMoProcessor: def process_video( self, video_input: VideoInput | VideoAudioInput, temporal_padding_factor=None ): - def smart_resize_video( num_total_frames, min_pixels, max_pixels, total_max_pixels, **kwargs ): diff --git a/test/registered/unit/distributed/test_mmencoder_forwards_attn_cp_size.py b/test/registered/unit/distributed/test_mmencoder_forwards_attn_cp_size.py deleted file mode 100644 index c94f74b68..000000000 --- a/test/registered/unit/distributed/test_mmencoder_forwards_attn_cp_size.py +++ /dev/null @@ -1,67 +0,0 @@ -"""MMEncoder must forward attn_cp_size to initialize_model_parallel, not just -tp_size. - -Real 2-GPU hardware confirmed a live mismatch this test guards against -statically: with `tp_size=2, attn_cp_size=2` published, calling -`initialize_model_parallel(tensor_model_parallel_size=2)` alone builds the -live attention-TP group at width 2, while `get_parallel().attn_tp_size` -(derived from the published config) answers 1 -- `VisionAttention` -(`layers/attention/vision.py`) reads that derived value as its own -weight-sharding width, so the mismatch is a real, silent wrong-sharding bug, -not just a reporting discrepancy. Forwarding -`attention_context_model_parallel_size=get_parallel().attn_cp_size` too -makes the two agree (confirmed on the same hardware). - -A full `MMEncoder` instantiation needs real weights and a live process -group, so this checks the one line that matters statically: the call passes -`attention_context_model_parallel_size` as well as -`tensor_model_parallel_size`. Not a substitute for testing `MMEncoder` -end-to-end under `--attn-cp-size > 1` on real hardware, but cheap enough to -run everywhere and catches the specific regression class (a future edit -that reverts to the tp-only call). -""" - -import ast -import os - -import sglang.srt.disaggregation.encoder.server as encoder_server_module -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=5, suite="base-a-test-cpu") - - -class TestMMEncoderForwardsAttnCpSize(CustomTestCase): - def test_initialize_model_parallel_call_forwards_attn_cp_size(self): - path = encoder_server_module.__file__ - tree = ast.parse(open(path).read()) - calls = [ - node - for node in ast.walk(tree) - if isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id == "initialize_model_parallel" - ] - self.assertEqual( - len(calls), - 1, - f"expected exactly one initialize_model_parallel(...) call in " - f"{os.path.basename(path)}, found {len(calls)} -- update this " - "test if that's now intentional", - ) - kwarg_names = {kw.arg for kw in calls[0].keywords} - self.assertIn( - "attention_context_model_parallel_size", - kwarg_names, - "MMEncoder's initialize_model_parallel(...) call must forward " - "attention_context_model_parallel_size (not just " - "tensor_model_parallel_size), or get_parallel().attn_tp_size " - "silently disagrees with the group actually built whenever " - "--attn-cp-size > 1 -- confirmed on real 2-GPU hardware", - ) - - -if __name__ == "__main__": - import unittest - - unittest.main() diff --git a/test/registered/unit/managers/test_scheduler_decision_batch_params.py b/test/registered/unit/managers/test_scheduler_decision_batch_params.py index 63b51aed2..89218ce33 100644 --- a/test/registered/unit/managers/test_scheduler_decision_batch_params.py +++ b/test/registered/unit/managers/test_scheduler_decision_batch_params.py @@ -1,4 +1,3 @@ -import inspect import unittest from types import SimpleNamespace from unittest.mock import patch @@ -8,46 +7,10 @@ from sglang.test.test_utils import maybe_stub_sgl_kernel maybe_stub_sgl_kernel() -from sglang.srt.disaggregation.decode import SchedulerDisaggregationDecodeMixin -from sglang.srt.disaggregation.prefill import SchedulerDisaggregationPrefillMixin from sglang.srt.managers.scheduler import Scheduler register_cpu_ci(est_time=12, suite="base-a-test-cpu") -FORBIDDEN_TOKENS = ("self.running_batch", "self.last_batch", "self.cur_batch") - -DECISION_METHODS = ( - Scheduler.get_next_batch_to_run, - Scheduler.get_new_batch_prefill, - Scheduler._get_new_batch_prefill_raw, - Scheduler.is_disable_overlap_for_batch, - SchedulerDisaggregationPrefillMixin.get_next_disagg_prefill_batch_to_run, - SchedulerDisaggregationPrefillMixin.process_prefill_chunk, - SchedulerDisaggregationDecodeMixin.get_new_prebuilt_batch, - SchedulerDisaggregationDecodeMixin.get_next_disagg_decode_batch_to_run, -) - - -class TestDecisionMethodsHaveNoHiddenBatchChannel(unittest.TestCase): - def test_decision_methods_take_batches_as_params_not_self(self): - """The batch decision tree must receive running/last batch as params, never via self.*.""" - for method in DECISION_METHODS: - source = inspect.getsource(inspect.unwrap(method)) - self.assertIn( - f"def {method.__name__}", - source, - msg=f"failed to read the real source of {method.__qualname__}", - ) - for token in FORBIDDEN_TOKENS: - self.assertNotIn( - token, - source, - msg=( - f"{method.__qualname__} references {token}; pass the batch " - "explicitly and return it via NextBatchPlan instead." - ), - ) - class TestMtpPhaseBoundaryOverlap(unittest.TestCase): @staticmethod diff --git a/test/registered/unit/tools/test_amd_ci_install_dependency.py b/test/registered/unit/tools/test_amd_ci_install_dependency.py deleted file mode 100644 index 45b22fb70..000000000 --- a/test/registered/unit/tools/test_amd_ci_install_dependency.py +++ /dev/null @@ -1,45 +0,0 @@ -"""Guard on the apt calls in scripts/ci/amd/amd_ci_install_dependency.sh. - -Those calls run under `set -euo pipefail`, and `apt-get update` exits 100 when -any single index is unreachable -- even though it keeps every index it did -fetch. An unguarded call therefore fails the whole "Install dependencies" step -on every AMD runner at once, which is what took out ~25 of 27 jobs in -pr-test-amd run 32399046576 when AMD's internal rocm-osdb artifactory started -404ing on an index this repo never installs from. - -The packages involved are optional -- rocm.Dockerfile builds MORI without them --- so no apt call here may be able to abort the run. -""" - -import re -import unittest -from pathlib import Path - -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=10, suite="base-a-test-cpu") - -INSTALL_SCRIPT = ( - Path(__file__).resolve().parents[4] / "scripts/ci/amd/amd_ci_install_dependency.sh" -) - - -class TestAmdCiInstallDependencyApt(CustomTestCase): - def test_apt_calls_cannot_abort_the_dependency_install(self): - unguarded = [ - line.strip() - for line in INSTALL_SCRIPT.read_text().splitlines() - if re.match(r"\s*(sudo\s+)?apt-get\b", line) and "||" not in line - ] - self.assertEqual( - unguarded, - [], - "an unguarded apt-get under `set -e` fails the dependency install on " - "every AMD runner whenever one apt source is unreachable; give it an " - "`|| echo ...` fallback", - ) - - -if __name__ == "__main__": - unittest.main()