perf(mla): TMA bulk-store set_mla_kv_buffer (up to 12× over baseline) (#25311)
This commit is contained in:
@@ -0,0 +1,127 @@
|
||||
"""Benchmark the set_mla_kv_buffer dispatcher.
|
||||
|
||||
Compares three providers across a batch-size sweep:
|
||||
- ``wrapper``: the high-level wrapper exposed by ``set_mla_kv_buffer_triton``
|
||||
(dispatches to TMA on SM90+, Triton fallback otherwise).
|
||||
- ``jit_tma``: the JIT CUDA TMA bulk-store kernel directly.
|
||||
- ``triton``: the BLOCK-tiled Triton kernel (SM<90 fallback path).
|
||||
"""
|
||||
|
||||
import itertools
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.testing
|
||||
|
||||
from sglang.jit_kernel.benchmark.utils import (
|
||||
DEFAULT_DEVICE,
|
||||
DEFAULT_DTYPE,
|
||||
DEFAULT_QUANTILES,
|
||||
get_benchmark_range,
|
||||
)
|
||||
from sglang.jit_kernel.set_mla_kv_buffer import set_mla_kv_buffer as jit_set
|
||||
from sglang.jit_kernel.utils import is_arch_support_pdl
|
||||
from sglang.srt.mem_cache.utils import set_mla_kv_buffer_kernel as sglang_triton_kernel
|
||||
from sglang.srt.mem_cache.utils import set_mla_kv_buffer_triton as sglang_wrapper
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=9, suite="stage-b-kernel-benchmark-1-gpu-large")
|
||||
|
||||
|
||||
def _triton_baseline(kv_buffer, loc, cache_k_nope, cache_k_rope):
|
||||
nope_dim = cache_k_nope.shape[-1]
|
||||
rope_dim = cache_k_rope.shape[-1]
|
||||
total_dim = nope_dim + rope_dim
|
||||
BLOCK = 128
|
||||
n_loc = loc.numel()
|
||||
grid = (n_loc, triton.cdiv(total_dim, BLOCK))
|
||||
pdl_kwargs = {"USE_GDC": True, "launch_pdl": True} if is_arch_support_pdl() else {}
|
||||
sglang_triton_kernel[grid](
|
||||
kv_buffer,
|
||||
cache_k_nope,
|
||||
cache_k_rope,
|
||||
loc,
|
||||
kv_buffer.stride(0),
|
||||
cache_k_nope.stride(0),
|
||||
cache_k_rope.stride(0),
|
||||
nope_dim,
|
||||
rope_dim,
|
||||
BLOCK=BLOCK,
|
||||
**pdl_kwargs,
|
||||
)
|
||||
|
||||
|
||||
NUM_LAYERS = 8
|
||||
CACHE_SIZE = (2 * 1024 * 1024) // NUM_LAYERS
|
||||
|
||||
NOPE_DIM = 512
|
||||
ROPE_DIM = 64
|
||||
|
||||
BS_RANGE = get_benchmark_range(
|
||||
full_range=[1, 8, 32, 128, 512, 1024, 2048, 4096, 8192, 16384],
|
||||
ci_range=[1, 128, 2048, 4096, 8192],
|
||||
)
|
||||
|
||||
LINE_VALS = ["wrapper", "jit_tma", "triton"]
|
||||
LINE_NAMES = ["Wrapper (auto)", "JIT TMA bulk-store", "Triton (BLOCK=128 baseline)"]
|
||||
STYLES = [("blue", "-"), ("green", "--"), ("red", "-.")]
|
||||
X_NAMES = ["batch_size"]
|
||||
CONFIGS = list(itertools.product(BS_RANGE))
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=X_NAMES,
|
||||
x_vals=CONFIGS,
|
||||
line_arg="provider",
|
||||
line_vals=LINE_VALS,
|
||||
line_names=LINE_NAMES,
|
||||
styles=STYLES,
|
||||
ylabel="us",
|
||||
plot_name="set-mla-kv-buffer-performance",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(batch_size: int, provider: str) -> Tuple[float, float, float]:
|
||||
cache_k_nope = torch.randn(
|
||||
(NUM_LAYERS, batch_size, 1, NOPE_DIM),
|
||||
dtype=DEFAULT_DTYPE,
|
||||
device=DEFAULT_DEVICE,
|
||||
)
|
||||
cache_k_rope = torch.randn(
|
||||
(NUM_LAYERS, batch_size, 1, ROPE_DIM),
|
||||
dtype=DEFAULT_DTYPE,
|
||||
device=DEFAULT_DEVICE,
|
||||
)
|
||||
kv_buffer = torch.randn(
|
||||
(NUM_LAYERS, CACHE_SIZE, 1, NOPE_DIM + ROPE_DIM),
|
||||
dtype=DEFAULT_DTYPE,
|
||||
device=DEFAULT_DEVICE,
|
||||
)
|
||||
loc = torch.randperm(CACHE_SIZE, device=DEFAULT_DEVICE)[:batch_size]
|
||||
torch.cuda.synchronize()
|
||||
|
||||
FN_MAP = {
|
||||
"wrapper": sglang_wrapper,
|
||||
"jit_tma": lambda buf, loc, n, r: jit_set(buf, loc, n, r),
|
||||
"triton": _triton_baseline,
|
||||
}
|
||||
|
||||
def fn():
|
||||
impl = FN_MAP[provider]
|
||||
for i in range(NUM_LAYERS):
|
||||
impl(kv_buffer[i], loc, cache_k_nope[i], cache_k_rope[i])
|
||||
|
||||
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
|
||||
fn, quantiles=DEFAULT_QUANTILES
|
||||
)
|
||||
return (
|
||||
1000 * ms / NUM_LAYERS,
|
||||
1000 * max_ms / NUM_LAYERS,
|
||||
1000 * min_ms / NUM_LAYERS,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
benchmark.run(print_data=True)
|
||||
@@ -0,0 +1,249 @@
|
||||
// JIT TMA bulk-store kernel for MLA paged-KV scatter writes.
|
||||
//
|
||||
// Each warp:
|
||||
// 1. Cooperatively loads one item's (nope, rope) row into a per-warp slot in
|
||||
// shared memory via vectorised ld/st.
|
||||
// 2. Lane 0 issues a single ``cp.async.bulk.global.shared::cta`` (TMA bulk
|
||||
// store, non-tensor variant) to scatter the row to
|
||||
// ``kv_buffer + loc[item] * stride_buffer``.
|
||||
//
|
||||
// End-of-CTA: ``cp.async.bulk.commit_group`` + ``wait_group<0>`` ensures all
|
||||
// in-flight stores commit before the kernel exits so the writes are visible
|
||||
// to subsequent kernels and the host.
|
||||
//
|
||||
// Two correctness gotchas worth a comment (easy to lose):
|
||||
// - ``fence.proxy.async.shared::cta`` between the smem fill and the TMA
|
||||
// store. The TMA engine reads via the async proxy; without the fence it
|
||||
// observes stale smem under heavy concurrency (manifests as zero rows at
|
||||
// large bs).
|
||||
// - ``wait_group`` not ``wait_group_read`` — the latter only allows early
|
||||
// smem reuse; it does not wait for the gmem store to commit globally.
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <sgl_kernel/tensor.h>
|
||||
#include <sgl_kernel/utils.h>
|
||||
|
||||
#include <sgl_kernel/tile.cuh>
|
||||
#include <sgl_kernel/utils.cuh>
|
||||
#include <sgl_kernel/vec.cuh>
|
||||
|
||||
#include <cuda/ptx>
|
||||
#include <dlpack/dlpack.h>
|
||||
#include <tvm/ffi/container/tensor.h>
|
||||
|
||||
#include <cstdint>
|
||||
|
||||
namespace {
|
||||
|
||||
struct SetMlaKVBufferParams {
|
||||
const void* __restrict__ k_nope;
|
||||
const void* __restrict__ k_rope;
|
||||
void* __restrict__ kv_buffer;
|
||||
const void* __restrict__ loc;
|
||||
int64_t stride_nope_bytes;
|
||||
int64_t stride_rope_bytes;
|
||||
int64_t stride_buffer_bytes;
|
||||
uint32_t batch_size;
|
||||
};
|
||||
|
||||
// Warp-cooperative gmem -> smem copy. Picks the widest vec width that divides
|
||||
// both the per-thread share and the byte total. Caller guarantees src is
|
||||
// 16-byte aligned (PyTorch tensors are) and dst is the start of a per-warp
|
||||
// smem slot (also 16-byte aligned by ``alignas(16)``).
|
||||
template <int64_t kBytes>
|
||||
SGL_DEVICE void warp_g2s_copy(const void* __restrict__ src, void* __restrict__ dst) {
|
||||
using namespace device;
|
||||
constexpr int64_t kAlignment = (kBytes % (16 * kWarpThreads) == 0) ? 16
|
||||
: (kBytes % (8 * kWarpThreads) == 0) ? 8
|
||||
: (kBytes % (4 * kWarpThreads) == 0) ? 4
|
||||
: (kBytes % 4 == 0) ? 4
|
||||
: 0;
|
||||
static_assert(kAlignment > 0, "kBytes must be a multiple of 4");
|
||||
|
||||
using vec_t = AlignedStorage<uint32_t, kAlignment / 4>;
|
||||
constexpr auto kLoopBytes = sizeof(vec_t) * kWarpThreads;
|
||||
constexpr auto kLoopCount = kBytes / kLoopBytes;
|
||||
constexpr int64_t kTailVecs = (kBytes - kLoopCount * kLoopBytes) / sizeof(vec_t);
|
||||
|
||||
const auto gmem = tile::Memory<vec_t>::warp();
|
||||
|
||||
#pragma unroll
|
||||
for (int64_t i = 0; i < kLoopCount; ++i) {
|
||||
const auto v = gmem.load(src, i);
|
||||
gmem.store(dst, v, i);
|
||||
}
|
||||
if constexpr (kTailVecs > 0) {
|
||||
if (gmem.in_bound(kLoopCount * kWarpThreads + kTailVecs, kLoopCount)) {
|
||||
const auto v = gmem.load(src, kLoopCount);
|
||||
gmem.store(dst, v, kLoopCount);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <int64_t kNopeBytes, int64_t kRopeBytes, int kNumWarps, bool kUsePDL, typename TLoc>
|
||||
__global__ void set_mla_kv_buffer_kernel(const __grid_constant__ SetMlaKVBufferParams params) {
|
||||
using namespace device;
|
||||
static_assert((kNopeBytes + kRopeBytes) % 16 == 0, "TMA bulk store requires total row to be 16-byte aligned");
|
||||
|
||||
constexpr int64_t kRowBytes = kNopeBytes + kRopeBytes;
|
||||
|
||||
// One contiguous smem slot per warp; align to 16 for TMA.
|
||||
__shared__ alignas(16) uint8_t smem[kNumWarps][kRowBytes];
|
||||
|
||||
const uint32_t warp_in_cta = threadIdx.x / kWarpThreads;
|
||||
const uint32_t item_id = blockIdx.x * kNumWarps + warp_in_cta;
|
||||
if (item_id >= params.batch_size) return;
|
||||
|
||||
PDLWaitPrimary<kUsePDL>();
|
||||
|
||||
const int64_t loc = static_cast<int64_t>(static_cast<const TLoc*>(params.loc)[item_id]);
|
||||
|
||||
const auto nope_src = pointer::offset(params.k_nope, item_id * params.stride_nope_bytes);
|
||||
const auto rope_src = pointer::offset(params.k_rope, item_id * params.stride_rope_bytes);
|
||||
void* const gmem_dst = pointer::offset(params.kv_buffer, loc * params.stride_buffer_bytes);
|
||||
|
||||
// Warp-cooperative load (nope, rope) into the per-warp smem slot.
|
||||
warp_g2s_copy<kNopeBytes>(nope_src, &smem[warp_in_cta][0]);
|
||||
warp_g2s_copy<kRopeBytes>(rope_src, &smem[warp_in_cta][kNopeBytes]);
|
||||
|
||||
// Fence required: TMA reads smem via the async proxy, normal sts writes
|
||||
// through the generic proxy. Without this the TMA engine can observe stale
|
||||
// values at large bs.
|
||||
__syncwarp();
|
||||
asm volatile("fence.proxy.async.shared::cta;" ::: "memory");
|
||||
|
||||
// Lane 0 issues one bulk store from the smem slot to the scattered gmem row.
|
||||
if (threadIdx.x % kWarpThreads == 0) {
|
||||
cuda::ptx::cp_async_bulk(
|
||||
cuda::ptx::space_global,
|
||||
cuda::ptx::space_shared,
|
||||
gmem_dst,
|
||||
&smem[warp_in_cta][0],
|
||||
static_cast<uint32_t>(kRowBytes));
|
||||
}
|
||||
|
||||
// Commit and wait for the CTA's bulk-stores to be globally visible before
|
||||
// returning. ``wait_group`` (not ``_read``) is the one that waits for gmem
|
||||
// commit; ``_read`` only releases smem for reuse.
|
||||
cuda::ptx::cp_async_bulk_commit_group();
|
||||
cuda::ptx::cp_async_bulk_wait_group(cuda::ptx::n32_t<0>{});
|
||||
|
||||
PDLTriggerSecondary<kUsePDL>();
|
||||
}
|
||||
|
||||
template <int64_t kNopeBytes, int64_t kRopeBytes, bool kUsePDL>
|
||||
struct SetMlaKVBufferKernel {
|
||||
static_assert(kNopeBytes > 0 && kNopeBytes % 4 == 0, "kNopeBytes must be a positive multiple of 4");
|
||||
static_assert(kRopeBytes > 0 && kRopeBytes % 4 == 0, "kRopeBytes must be a positive multiple of 4");
|
||||
static_assert(
|
||||
(kNopeBytes + kRopeBytes) % 16 == 0, "TMA bulk store requires (kNopeBytes + kRopeBytes) to be a multiple of 16");
|
||||
|
||||
template <int kNumWarps, typename TLoc>
|
||||
static constexpr auto kernel = set_mla_kv_buffer_kernel<kNopeBytes, kRopeBytes, kNumWarps, kUsePDL, TLoc>;
|
||||
|
||||
static void
|
||||
run(tvm::ffi::TensorView kv_buffer,
|
||||
tvm::ffi::TensorView loc,
|
||||
tvm::ffi::TensorView k_nope,
|
||||
tvm::ffi::TensorView k_rope,
|
||||
int64_t num_warps_per_block) {
|
||||
using namespace host;
|
||||
|
||||
auto B = SymbolicSize{"batch_size"};
|
||||
auto D_nope = SymbolicSize{"nope_dim"};
|
||||
auto D_rope = SymbolicSize{"rope_dim"};
|
||||
auto D_buf = SymbolicSize{"buffer_last_dim"};
|
||||
auto S_nope = SymbolicSize{"nope_stride"};
|
||||
auto S_rope = SymbolicSize{"rope_stride"};
|
||||
auto S_buf = SymbolicSize{"buffer_stride"};
|
||||
auto S_loc = SymbolicSize{"loc_stride"};
|
||||
auto dtype = SymbolicDType{};
|
||||
auto loc_dtype = SymbolicDType{};
|
||||
auto device = SymbolicDevice{};
|
||||
device.set_options<kDLCUDA>();
|
||||
|
||||
TensorMatcher({B, D_nope}) //
|
||||
.with_strides({S_nope, 1})
|
||||
.with_dtype(dtype)
|
||||
.with_device(device)
|
||||
.verify(k_nope);
|
||||
TensorMatcher({B, D_rope}) //
|
||||
.with_strides({S_rope, 1})
|
||||
.with_dtype(dtype)
|
||||
.with_device(device)
|
||||
.verify(k_rope);
|
||||
TensorMatcher({-1, D_buf}) //
|
||||
.with_strides({S_buf, 1})
|
||||
.with_dtype(dtype)
|
||||
.with_device(device)
|
||||
.verify(kv_buffer);
|
||||
TensorMatcher({B}) //
|
||||
.with_strides({S_loc})
|
||||
.with_dtype<int32_t, int64_t>(loc_dtype)
|
||||
.with_device(device)
|
||||
.verify(loc);
|
||||
|
||||
const int64_t dtype_size = dtype_bytes(dtype.unwrap());
|
||||
RuntimeCheck(
|
||||
kNopeBytes == dtype_size * D_nope.unwrap(),
|
||||
"kNopeBytes mismatch: expected ",
|
||||
kNopeBytes,
|
||||
", got ",
|
||||
dtype_size * D_nope.unwrap());
|
||||
RuntimeCheck(
|
||||
kRopeBytes == dtype_size * D_rope.unwrap(),
|
||||
"kRopeBytes mismatch: expected ",
|
||||
kRopeBytes,
|
||||
", got ",
|
||||
dtype_size * D_rope.unwrap());
|
||||
RuntimeCheck(dtype_size * D_buf.unwrap() >= kNopeBytes + kRopeBytes, "kv_buffer last dim too small");
|
||||
RuntimeCheck(
|
||||
(S_buf.unwrap() * dtype_size) % 16 == 0,
|
||||
"kv_buffer row stride must be a multiple of 16 bytes for TMA bulk store; got ",
|
||||
S_buf.unwrap() * dtype_size);
|
||||
|
||||
const uint32_t batch = static_cast<uint32_t>(B.unwrap());
|
||||
if (batch == 0) return;
|
||||
|
||||
const auto params = SetMlaKVBufferParams{
|
||||
.k_nope = k_nope.data_ptr(),
|
||||
.k_rope = k_rope.data_ptr(),
|
||||
.kv_buffer = kv_buffer.data_ptr(),
|
||||
.loc = loc.data_ptr(),
|
||||
.stride_nope_bytes = S_nope.unwrap() * dtype_size,
|
||||
.stride_rope_bytes = S_rope.unwrap() * dtype_size,
|
||||
.stride_buffer_bytes = S_buf.unwrap() * dtype_size,
|
||||
.batch_size = batch,
|
||||
};
|
||||
|
||||
const auto use_int32 = loc_dtype.is_type<int32_t>();
|
||||
|
||||
auto launch = [&]<int kNW>() {
|
||||
const auto kernel_ptr = use_int32 ? kernel<kNW, int32_t> : kernel<kNW, int64_t>;
|
||||
const uint32_t num_blocks = div_ceil(batch, static_cast<uint32_t>(kNW));
|
||||
const uint32_t threads_per_block = static_cast<uint32_t>(kNW) * device::kWarpThreads;
|
||||
LaunchKernel(num_blocks, threads_per_block, device.unwrap()) //
|
||||
.enable_pdl(kUsePDL)(kernel_ptr, params);
|
||||
};
|
||||
|
||||
switch (num_warps_per_block) {
|
||||
case 1:
|
||||
launch.template operator()<1>();
|
||||
break;
|
||||
case 2:
|
||||
launch.template operator()<2>();
|
||||
break;
|
||||
case 4:
|
||||
launch.template operator()<4>();
|
||||
break;
|
||||
case 8:
|
||||
launch.template operator()<8>();
|
||||
break;
|
||||
default:
|
||||
Panic("Unsupported num_warps_per_block=", num_warps_per_block);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace
|
||||
@@ -0,0 +1,121 @@
|
||||
"""JIT TMA bulk-store path for ``set_mla_kv_buffer``.
|
||||
|
||||
Each warp scatter-writes one item's (nope, rope) row via a single
|
||||
``cp.async.bulk.global.shared::cta`` store. Requires SM90+ (Hopper or later)
|
||||
for the TMA bulk-store hardware. The host-side wrapper in
|
||||
``sglang.srt.mem_cache.utils`` falls back to a Triton kernel for older arches.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.jit_kernel.utils import (
|
||||
cache_once,
|
||||
is_arch_support_pdl,
|
||||
load_jit,
|
||||
make_cpp_args,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tvm_ffi.module import Module
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@cache_once
|
||||
def _jit_set_mla_kv_buffer_module(
|
||||
nope_bytes: int, rope_bytes: int, use_pdl: bool
|
||||
) -> Module:
|
||||
args = make_cpp_args(nope_bytes, rope_bytes, use_pdl)
|
||||
return load_jit(
|
||||
f"set_mla_kv_buffer_{nope_bytes}_{rope_bytes}",
|
||||
*args,
|
||||
cuda_files=["elementwise/set_mla_kv_buffer.cuh"],
|
||||
cuda_wrappers=[
|
||||
("set_mla_kv_buffer", f"SetMlaKVBufferKernel<{args}>::run"),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@cache_once
|
||||
def can_use_set_mla_kv_buffer(nope_bytes: int, rope_bytes: int) -> bool:
|
||||
"""Whether the TMA path can be used for these row byte widths.
|
||||
|
||||
TMA bulk store requires ``(nope_bytes + rope_bytes)`` to be a multiple of
|
||||
16; both halves individually must also be a multiple of 4 (the warp-coop
|
||||
smem load lower bound).
|
||||
"""
|
||||
if nope_bytes % 4 != 0 or rope_bytes % 4 != 0:
|
||||
logger.warning(
|
||||
"Unsupported nope_bytes=%d rope_bytes=%d for JIT set_mla_kv_buffer:"
|
||||
" both must be multiples of 4",
|
||||
nope_bytes,
|
||||
rope_bytes,
|
||||
)
|
||||
return False
|
||||
if (nope_bytes + rope_bytes) % 16 != 0:
|
||||
logger.warning(
|
||||
"Unsupported nope_bytes=%d rope_bytes=%d for JIT set_mla_kv_buffer:"
|
||||
" (nope_bytes + rope_bytes) must be a multiple of 16 for TMA bulk store",
|
||||
nope_bytes,
|
||||
rope_bytes,
|
||||
)
|
||||
return False
|
||||
try:
|
||||
_jit_set_mla_kv_buffer_module(nope_bytes, rope_bytes, is_arch_support_pdl())
|
||||
return True
|
||||
except Exception as e: # pragma: no cover - compile-time only
|
||||
logger.warning(
|
||||
"Failed to load JIT set_mla_kv_buffer kernel "
|
||||
"with nope_bytes=%d rope_bytes=%d: %s",
|
||||
nope_bytes,
|
||||
rope_bytes,
|
||||
e,
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def _pick_num_warps(n_loc: int) -> int:
|
||||
# Tuned on GB300: nw=4 wins below 1024 (more CTAs spread across SMs);
|
||||
# nw=8 wins above (each CTA amortises the bulk-group commit better).
|
||||
return 4 if n_loc <= 768 else 8
|
||||
|
||||
|
||||
def set_mla_kv_buffer(
|
||||
kv_buffer: torch.Tensor,
|
||||
loc: torch.Tensor,
|
||||
cache_k_nope: torch.Tensor,
|
||||
cache_k_rope: torch.Tensor,
|
||||
num_warps: int = 0,
|
||||
) -> None:
|
||||
"""Write packed [k_nope | k_rope] rows into ``kv_buffer`` at ``loc`` indices
|
||||
via a TMA bulk-store. SM90+ only — the caller is expected to gate.
|
||||
|
||||
Shapes (last dim is treated as the row payload; any leading singleton dims
|
||||
on the source tensors are flattened away):
|
||||
kv_buffer: [num_pages, total_dim] or [num_pages, 1, total_dim]
|
||||
cache_k_nope: [n_loc, nope_dim] or [n_loc, 1, nope_dim]
|
||||
cache_k_rope: [n_loc, rope_dim] or [n_loc, 1, rope_dim]
|
||||
loc: [n_loc]
|
||||
"""
|
||||
n_loc = loc.shape[0]
|
||||
if n_loc == 0:
|
||||
return
|
||||
|
||||
src_nope = cache_k_nope.view(n_loc, -1) if cache_k_nope.dim() != 2 else cache_k_nope
|
||||
src_rope = cache_k_rope.view(n_loc, -1) if cache_k_rope.dim() != 2 else cache_k_rope
|
||||
buf = kv_buffer.view(kv_buffer.shape[0], -1) if kv_buffer.dim() != 2 else kv_buffer
|
||||
|
||||
nope_bytes = src_nope.shape[-1] * src_nope.element_size()
|
||||
rope_bytes = src_rope.shape[-1] * src_rope.element_size()
|
||||
if num_warps <= 0:
|
||||
num_warps = _pick_num_warps(n_loc)
|
||||
|
||||
module = _jit_set_mla_kv_buffer_module(
|
||||
nope_bytes, rope_bytes, is_arch_support_pdl()
|
||||
)
|
||||
module.set_mla_kv_buffer(buf, loc, src_nope, src_rope, num_warps)
|
||||
@@ -0,0 +1,126 @@
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.jit_kernel.set_mla_kv_buffer import (
|
||||
can_use_set_mla_kv_buffer,
|
||||
set_mla_kv_buffer,
|
||||
)
|
||||
from sglang.jit_kernel.utils import get_ci_test_range
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=30, suite="stage-b-kernel-unit-1-gpu-large")
|
||||
|
||||
DEVICE = "cuda"
|
||||
CACHE_SIZE = 4096
|
||||
|
||||
# (nope_dim, rope_dim) pairs: standard MLA, MLA scale buffer, FP8 nope-extended layout.
|
||||
SHAPES = get_ci_test_range(
|
||||
[(512, 64), (512, 32), (256, 64), (128, 64), (528, 64)],
|
||||
[(512, 64), (528, 64)],
|
||||
)
|
||||
BATCH_SIZES = get_ci_test_range([1, 7, 64, 257, 1024], [1, 64, 1024])
|
||||
|
||||
|
||||
def _ref(kv_buffer, loc, cache_k_nope, cache_k_rope):
|
||||
nope_dim = cache_k_nope.shape[-1]
|
||||
n_loc = loc.shape[0]
|
||||
src_nope = cache_k_nope.reshape(n_loc, -1)
|
||||
src_rope = cache_k_rope.reshape(n_loc, -1)
|
||||
kv_view = kv_buffer.view(kv_buffer.shape[0], -1)
|
||||
kv_view[loc.long(), :nope_dim] = src_nope
|
||||
kv_view[loc.long(), nope_dim : nope_dim + src_rope.shape[-1]] = src_rope
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
@pytest.mark.parametrize("shape", SHAPES)
|
||||
@pytest.mark.parametrize("batch_size", BATCH_SIZES)
|
||||
def test_set_mla_kv_buffer_correctness(dtype, shape, batch_size):
|
||||
nope_dim, rope_dim = shape
|
||||
total_dim = nope_dim + rope_dim
|
||||
|
||||
cache_k_nope = torch.randn((batch_size, 1, nope_dim), dtype=dtype, device=DEVICE)
|
||||
cache_k_rope = torch.randn((batch_size, 1, rope_dim), dtype=dtype, device=DEVICE)
|
||||
kv_buffer = torch.randn((CACHE_SIZE, 1, total_dim), dtype=dtype, device=DEVICE)
|
||||
kv_ref = kv_buffer.clone()
|
||||
|
||||
loc = torch.randperm(CACHE_SIZE, device=DEVICE)[:batch_size]
|
||||
|
||||
set_mla_kv_buffer(kv_buffer, loc, cache_k_nope, cache_k_rope)
|
||||
_ref(kv_ref, loc, cache_k_nope, cache_k_rope)
|
||||
|
||||
assert torch.equal(kv_buffer, kv_ref)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("loc_dtype", [torch.int32, torch.int64])
|
||||
def test_set_mla_kv_buffer_loc_dtypes(loc_dtype):
|
||||
nope_dim, rope_dim = 512, 64
|
||||
batch_size = 128
|
||||
dtype = torch.bfloat16
|
||||
|
||||
cache_k_nope = torch.randn((batch_size, 1, nope_dim), dtype=dtype, device=DEVICE)
|
||||
cache_k_rope = torch.randn((batch_size, 1, rope_dim), dtype=dtype, device=DEVICE)
|
||||
kv_buffer = torch.randn(
|
||||
(CACHE_SIZE, 1, nope_dim + rope_dim), dtype=dtype, device=DEVICE
|
||||
)
|
||||
kv_ref = kv_buffer.clone()
|
||||
|
||||
loc = torch.randperm(CACHE_SIZE, device=DEVICE)[:batch_size].to(loc_dtype)
|
||||
|
||||
set_mla_kv_buffer(kv_buffer, loc, cache_k_nope, cache_k_rope)
|
||||
_ref(kv_ref, loc, cache_k_nope, cache_k_rope)
|
||||
|
||||
assert torch.equal(kv_buffer, kv_ref)
|
||||
|
||||
|
||||
def test_set_mla_kv_buffer_uint8_byte_layout():
|
||||
"""FP8 NSA byte-layout: cache_k_nope is uint8 with [fp8(512) | scales(16)] = 528,
|
||||
cache_k_rope is uint8 [128]; total payload = 656 bytes."""
|
||||
nope_bytes, rope_bytes = 528, 128
|
||||
batch_size = 64
|
||||
dtype = torch.uint8
|
||||
|
||||
cache_k_nope = torch.randint(
|
||||
0, 256, (batch_size, 1, nope_bytes), dtype=dtype, device=DEVICE
|
||||
)
|
||||
cache_k_rope = torch.randint(
|
||||
0, 256, (batch_size, 1, rope_bytes), dtype=dtype, device=DEVICE
|
||||
)
|
||||
kv_buffer = torch.randint(
|
||||
0, 256, (CACHE_SIZE, 1, nope_bytes + rope_bytes), dtype=dtype, device=DEVICE
|
||||
)
|
||||
kv_ref = kv_buffer.clone()
|
||||
|
||||
loc = torch.randperm(CACHE_SIZE, device=DEVICE)[:batch_size]
|
||||
|
||||
set_mla_kv_buffer(kv_buffer, loc, cache_k_nope, cache_k_rope)
|
||||
_ref(kv_ref, loc, cache_k_nope, cache_k_rope)
|
||||
|
||||
assert torch.equal(kv_buffer, kv_ref)
|
||||
|
||||
|
||||
def test_set_mla_kv_buffer_empty_loc():
|
||||
nope_dim, rope_dim = 512, 64
|
||||
dtype = torch.bfloat16
|
||||
cache_k_nope = torch.empty((0, 1, nope_dim), dtype=dtype, device=DEVICE)
|
||||
cache_k_rope = torch.empty((0, 1, rope_dim), dtype=dtype, device=DEVICE)
|
||||
kv_buffer = torch.randn(
|
||||
(CACHE_SIZE, 1, nope_dim + rope_dim), dtype=dtype, device=DEVICE
|
||||
)
|
||||
kv_before = kv_buffer.clone()
|
||||
|
||||
loc = torch.empty((0,), dtype=torch.int64, device=DEVICE)
|
||||
set_mla_kv_buffer(kv_buffer, loc, cache_k_nope, cache_k_rope)
|
||||
|
||||
assert torch.equal(kv_buffer, kv_before)
|
||||
|
||||
|
||||
def test_can_use_set_mla_kv_buffer():
|
||||
assert can_use_set_mla_kv_buffer(1024, 128) # bf16 (512,64)
|
||||
assert can_use_set_mla_kv_buffer(528, 128) # fp8 byte layout
|
||||
assert not can_use_set_mla_kv_buffer(13, 8) # not multiple of 4
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
||||
@@ -91,21 +91,71 @@ def set_mla_kv_buffer_kernel(
|
||||
tl.extra.cuda.gdc_launch_dependents()
|
||||
|
||||
|
||||
# Above this loc count the TMA bulk-store path overtakes the single-CTA-per-loc
|
||||
# Triton kernel. Below it, Triton with BLOCK = next_pow2(total_dim) (one CTA
|
||||
# does the whole row in one tile, no boundary fan-out) is the winning fallback.
|
||||
# Tuned on GB300 with DSv4 row widths.
|
||||
_TMA_BULK_STORE_MIN_LOCS = 768
|
||||
|
||||
|
||||
def set_mla_kv_buffer_triton(
|
||||
kv_buffer: torch.Tensor,
|
||||
loc: torch.Tensor,
|
||||
cache_k_nope: torch.Tensor,
|
||||
cache_k_rope: torch.Tensor,
|
||||
):
|
||||
"""Dispatch MLA paged-KV scatter writes to the fastest available path.
|
||||
|
||||
Two paths, chosen on ``n_loc``:
|
||||
|
||||
- ``n_loc >= 768`` (and SM90+ with TMA-compatible row widths): JIT CUDA
|
||||
kernel where each warp loads one (nope, rope) row into shared memory and
|
||||
issues a single ``cp.async.bulk.global.shared::cta`` store to scatter the
|
||||
row at ``kv_buffer[loc[item]]``. Wins at large bs because it packs 4-8
|
||||
items per CTA, drastically reducing the CTA count vs single-CTA-per-loc.
|
||||
- Otherwise: Triton kernel with ``BLOCK = next_pow2(nope_dim + rope_dim)``,
|
||||
i.e. one CTA per loc covering the entire row in one tile. Wins at small
|
||||
bs because there's no per-loc CTA fan-out (5× fewer CTAs than the old
|
||||
BLOCK=128 dispatch) and the row-spanning block makes the boundary branch
|
||||
a one-shot per CTA. This is also the path for SM<90 and for shapes that
|
||||
violate the TMA 16-byte alignment.
|
||||
|
||||
Speedup vs the legacy BLOCK=128 Triton kernel on GB300 (BF16, nope=512,
|
||||
rope=64): ~1.05× at bs=8, ~1.5× at bs=128, 3.5× at bs=512, **11.7× at
|
||||
bs=16384**.
|
||||
|
||||
Name retained for caller compatibility; the implementation is no longer
|
||||
Triton-only.
|
||||
"""
|
||||
from sglang.jit_kernel.set_mla_kv_buffer import (
|
||||
can_use_set_mla_kv_buffer,
|
||||
)
|
||||
from sglang.jit_kernel.set_mla_kv_buffer import (
|
||||
set_mla_kv_buffer as jit_set_mla_kv_buffer,
|
||||
)
|
||||
|
||||
n_loc = loc.numel()
|
||||
nope_bytes = cache_k_nope.shape[-1] * cache_k_nope.element_size()
|
||||
rope_bytes = cache_k_rope.shape[-1] * cache_k_rope.element_size()
|
||||
if (
|
||||
n_loc >= _TMA_BULK_STORE_MIN_LOCS
|
||||
and is_arch_support_pdl()
|
||||
and can_use_set_mla_kv_buffer(nope_bytes, rope_bytes)
|
||||
):
|
||||
jit_set_mla_kv_buffer(kv_buffer, loc, cache_k_nope, cache_k_rope)
|
||||
return
|
||||
|
||||
# Fallback: Triton with BLOCK = next_pow2(total_dim). One CTA per loc; the
|
||||
# whole row in one tile (the existing 3-way nope/rope/boundary branch in
|
||||
# ``set_mla_kv_buffer_kernel`` handles the over-allocation past total_dim
|
||||
# via the offs<total_dim mask). Beats BLOCK=128 by 60-2700 ns across the
|
||||
# 2 ≤ bs ≤ 512 range on GB300.
|
||||
nope_dim = cache_k_nope.shape[-1]
|
||||
rope_dim = cache_k_rope.shape[-1]
|
||||
total_dim = nope_dim + rope_dim
|
||||
BLOCK = 128
|
||||
n_loc = loc.numel()
|
||||
grid = (n_loc, triton.cdiv(total_dim, BLOCK))
|
||||
|
||||
BLOCK = triton.next_power_of_2(total_dim)
|
||||
grid = (n_loc, 1)
|
||||
pdl_kwargs = {"USE_GDC": True, "launch_pdl": True} if is_arch_support_pdl() else {}
|
||||
|
||||
set_mla_kv_buffer_kernel[grid](
|
||||
kv_buffer,
|
||||
cache_k_nope,
|
||||
|
||||
Reference in New Issue
Block a user