From 81363bf8cb545ff8e1a9838733448c2252dd6cc0 Mon Sep 17 00:00:00 2001 From: DarkSharpness <76582120+DarkSharpness@users.noreply.github.com> Date: Fri, 18 Sep 2026 22:40:54 +0800 Subject: [PATCH] [kernel] Share the warp vectorized copy and enforce its alignment (#36176) Co-authored-by: Claude Opus 5 (1M context) Co-authored-by: BBuf <1182563586@qq.com> --- .claude/skills/add-jit-kernel/SKILL.md | 55 ++- python/sglang/kernels/jit/benchmark/marker.py | 118 +++++-- .../kernels/jit/csrc/deepseek_v4/c_plan.cuh | 14 - .../deepseek_v4/candidate_block_table.cuh | 2 +- .../csrc/deepseek_v4/fused_norm_rope_v2.cuh | 4 +- .../silu_and_mul_masked_post_quant.cuh | 12 +- .../jit/csrc/deepseek_v4/topk_bf16_small.cuh | 6 +- .../jit/csrc/elementwise/concat_mla.cuh | 98 +++--- .../kernels/jit/csrc/elementwise/kvcache.cuh | 312 +++++------------ .../csrc/elementwise/set_mla_kv_buffer.cuh | 216 ++++-------- .../csrc/elementwise/set_mla_kv_concat_q.cuh | 37 +- .../kernels/jit/csrc/kimi_k3/situ_and_mul.cuh | 4 +- .../kernels/jit/csrc/kvcacheio/hicache.cuh | 20 +- .../jit/csrc/minimax/minimax_decode_topk.cuh | 16 +- .../kernels/jit/csrc/moe/route_radix.cuh | 14 +- .../jit/include/sgl_kernel/allocator.h | 40 +++ .../kernels/jit/include/sgl_kernel/atomic.cuh | 130 +++++++ .../kernels/jit/include/sgl_kernel/bits.h | 41 +++ .../sgl_kernel/deepseek_v4/topk_impl.cuh | 20 +- .../jit/include/sgl_kernel/runtime.cuh | 81 +++-- .../kernels/jit/include/sgl_kernel/tensor.h | 23 ++ .../kernels/jit/include/sgl_kernel/tile.cuh | 6 +- .../kernels/jit/include/sgl_kernel/utils.cuh | 100 +++++- .../kernels/jit/include/sgl_kernel/vec.cuh | 17 + .../kernels/jit/include/sgl_kernel/warp.cuh | 321 ++++++++++++------ .../kernels/jit/utils/compile/loader.py | 24 +- .../sglang/kernels/jit/utils/compile/ninja.py | 30 ++ python/sglang/kernels/ops/kvcache/kvcache.py | 56 +-- .../sglang/kernels/ops/kvcache/mla_buffer.py | 29 +- .../kernels/ops/kvcache/set_mla_kv_buffer.py | 22 +- python/sglang/srt/environ.py | 20 ++ .../benchmark/kvcache/bench_hicache.py | 105 ++---- .../kvcache/bench_set_mla_kv_buffer.py | 75 ++-- .../kernels/ops/kvcache/test_store_cache.py | 95 +----- 34 files changed, 1172 insertions(+), 991 deletions(-) create mode 100644 python/sglang/kernels/jit/include/sgl_kernel/allocator.h create mode 100644 python/sglang/kernels/jit/include/sgl_kernel/bits.h diff --git a/.claude/skills/add-jit-kernel/SKILL.md b/.claude/skills/add-jit-kernel/SKILL.md index 9789ce7fd..11fa1d02e 100644 --- a/.claude/skills/add-jit-kernel/SKILL.md +++ b/.claude/skills/add-jit-kernel/SKILL.md @@ -28,12 +28,15 @@ Add a new operation that scales each element of a tensor by a scalar factor: These hold for every step below. - **`namespace sglang` is where JIT code lives.** Open it after the include block and close it at the end of the file, with the device kernels, traits and host wrapper inside. The shared `host::` / `device::` helpers are nested in it too, so they resolve unqualified. `load_jit` emits the `TVM_FFI_DLL_EXPORT_TYPED_FUNC` wrapper inside `namespace sglang` as well, so the `kernel_name` you pass from Python needs no `sglang::` prefix. +- **Tuning policy belongs in Python; C++ gets the answer, not the decision.** For a hyperparameter — split factor, threads per item, block size, vector width, an algorithm variant — prefer (not required) making it a **template parameter** and letting the `@cache_once` module factory choose the value, over a runtime `if`/`switch` in the launcher that picks among pre-instantiated kernels. The kernel then compiles exactly one specialisation and asserts its own preconditions with `static_assert`, while the heuristic that produced the value sits in Python where it is readable, adjustable, and inspectable without recompiling CUDA. `store_cache`'s `num_threads` is the worked example: a `get_kernel(num_split)` ladder guarded by byte-alignment `if constexpr`s became one template argument plus a heuristic in `_jit_kvcache_module`. The exception is a knob that genuinely varies per call with the runtime shape — that has to stay a kernel argument. - **Check where the check is cheapest: `static_assert` > C++ host check > cached Python > per-call Python.** Anything fixed at compile time is a `static_assert`. Anything about the tensors is a `TensorMatcher` / `CHECK_HOST` in the C++ launcher, free next to a kernel launch. A check Python cannot delegate goes inside the `@cache_once` module factory, where it runs once per specialisation. What remains in the per-call entry point costs interpreter time on *every* forward, so it should be nothing but picking the module and allocating `out`. - **Fixed-width integer types.** Prefer `int32_t` / `int64_t` / `uint32_t` / `size_t` over `int`, `long`, or `long long`, so an index has the same width on both sides of the FFI boundary. Bare `int` is fine only where the width plainly cannot matter — an unrolled loop counter over a `constexpr` bound, a template `int` parameter. Shapes arrive as `int64_t` (`SymbolicSize::unwrap()`); narrowing to `uint32_t` for in-kernel indexing is a deliberate act, so write the `static_cast` explicitly and only where the range is known. - **Doxygen comments in C++.** Document exported entities with `///` or `/** ... */` blocks using `\brief`, `\param`, `\tparam`, `\return`, the way `include/sgl_kernel/` does. `python -m sglang.kernels.jit` writes `CommentFormat: Doxygen` into `.clangd` when clangd is 21 or newer, so these render on hover in the editor. Plain `//` remains fine for implementation notes inside a function body. - **ASCII only in C++ and CUDA sources.** Write `--`, `->`, `<=` instead of `—`, `→`, `≤`, including in comments. `grep -nP '[^\x00-\x7F]' ` before committing. +- **`namespace details` is private.** Anything inside one is an implementation detail of its own header, free to change without notice. Do not name `details::` from another module, and never `using namespace details`. If you find yourself wanting something in there, that is the signal to promote it to a documented name instead. - **`const T* __restrict__` for read-only pointers.** This is what `csrc/` does throughout, and it lets the compiler emit non-coherent (`LDG`) loads. -- **Watch the register budget.** For memory-bound kernels, keep to roughly 64 registers per thread so occupancy does not become the limit. Build once with `extra_cuda_cflags=["-Xptxas", "-v"]` to see the actual count, and prefer recomputing a value over letting it spill. +- **Watch the register budget.** For memory-bound kernels, keep to roughly 64 registers per thread so occupancy does not become the limit; prefer recomputing a value over letting it spill. Run `SGLANG_JIT_LOG_RESOURCE_USAGE=1 SGLANG_JIT_FORCE_RECOMPILE=1` to log per-kernel registers, spills and shared memory at INFO. Both halves are load-bearing: a cache hit has no compiler output to report, and passing `-Xptxas -v` by hand does nothing on its own because the build captures compiler output and replays it only on failure. +- **Alignment is a contract you enforce, not a property you hope for.** A vectorized copy derives its access width from the row *size* (`load_bytes` picks 16B for a 1024B row), but the address it lands on is `base + index * row_stride` — and a stride is not constrained by the size. `TensorMatcher` admits any stride unless you say otherwise, so a padded cache row silently produces `misaligned address`. Every kernel that vectorizes is highly recommended to end its validation with `.ensure_alignment(w)`, where `w` is the width the kernel actually uses. --- @@ -62,7 +65,9 @@ These hold for every step below. - **Type aliases**: `fp16_t`, `bf16_t`, `fp32_t`, `fp8_e4m3_t`, `fp8_e5m2_t` and their packed variants `fp16x2_t`, `bf16x2_t`, `fp32x2_t`, etc. - **`SGL_DEVICE`** — Expands to `__forceinline__ __device__`. Use on all device functions. -- **`device::kWarpThreads`** — Constant `32`. +- **`SGL_DEVICE_HOST`** — `__forceinline__ __device__ __host__`. Use it when a `constexpr` helper must be callable from both the kernel and its launcher, so the two cannot drift. +- **`device::kWarpThreads`** — Constant `32` (the codebase's *logical* warp width). `kWarpSize` is an alias. Note this is a group size, not the hardware wave: on gfx950 `warpSize` is 64, and `warp::kFullWidth` reflects that. +- **`device::get_lane_id()`** — this thread's index within its logical group. **Prefer it over `threadIdx.x % kNumThreads` whenever the value feeds an address.** On CUDA the lane register is one read that folds into the address computation, whereas the modulo is a derived value the compiler re-materialises at *every* address scale — 8 instructions and 2 registers on a two-tile warp copy, measured on sm_100a. It picks the cheaper form per platform, so callers never need to care. Only equals the true in-warp lane when `blockDim.x` is a multiple of `kNumThreads`. - **`device::load_as(ptr, offset)`** / **`device::store_as(ptr, val, offset)`** — Type-safe loads/stores from `void*`. - **`device::pointer::offset(ptr, offsets...)`** — Pointer arithmetic on device. - **`host::LaunchKernel(grid, block, device_or_stream [, smem])`** — RAII kernel launcher that: @@ -93,6 +98,7 @@ This is the **primary validation API** for all kernel launchers. Use it to valid - `.with_dtype()` — allow a set of types - `.with_device(device_sym)` — require CUDA and bind the checked device to a `SymbolicDevice` - `.with_strides({strides...})` — validate strides (omit to require contiguous) + - `.ensure_alignment(bytes)` — require the data pointer **and every stride except the innermost** to be a multiple of `bytes` (a power of two). Size-1 dimensions are skipped, since their strides are arbitrary. This is how a vectorized kernel states its precondition; see the alignment convention above. - `.verify(tensor_view)` — execute the check; throws `PanicError` with full context on failure; **chainable** (`verify(a).verify(b)` to check multiple tensors with the same shape) - **`host::is_type(dtype)`** — whether a `DLDataType` denotes the C++ type `T` (e.g. `fp16_t`). @@ -104,6 +110,7 @@ device.set_options(); TensorMatcher({N}) // .with_dtype() .with_device(device) + .ensure_alignment(16) // e.g. for 128-bit vectorized load/store .verify(dst) .verify(src); // same shape, dtype, device as dst const int64_t n = N.unwrap(); @@ -155,6 +162,7 @@ The counterpart to `tensor.h`: that one validates what came in, this one produce - `.store(ptr, offset)` — vectorized store to `ptr[offset]` - `.fill(value)` — fill all N elements with `value` - `operator[](i)` — element access +- **`device::LoadStoreBytes`** — named widths to pass around instead of bare integers: `MAX_GMEM` (arch-dependent, 32 on Blackwell), `MAX_SMEM` (16), `MAX_PORTABLE` (16, safe on CUDA and HIP), `MIN_COALALESCED` (4), and `RAW_1B` .. `RAW_32B`. ### `tile.cuh` — `tile::Memory` (strided memory access pattern) @@ -164,12 +172,22 @@ The counterpart to `tensor.h`: that one validates what came in, this one produce - `tile::Memory` is fundamentally a **1D cooperative accessor** over a contiguous region. - **`device::tile::Memory::cta(blockDim.x)`** — Creates a tile accessor where each thread handles `tid = threadIdx.x` with stride `tsize` (for `cta(blockDim.x)`, this is `blockDim.x`). Common for loops over a 1D array. +- **`device::tile::Memory::warp()`** — the 32-lane flavour, and what `warp::load_bytes` is built on. It takes its lane index from `get_lane_id()` rather than `threadIdx.x % 32`, for the addressing reason described under `utils.cuh`. Use `warp(int n)` for a narrower sub-group; that overload keeps the modulo, since a hardware lane index is the wrong grouping below 32. +- **`device::tile::Memory::thread()`** — single-thread accessor, no cooperation. - **`.load(ptr, offset)`** — loads `ptr[tid + offset * tsize]` - **`.store(ptr, val, offset)`** — stores to `ptr[tid + offset * tsize]` - **`.in_bound(n, offset)`** — boundary check For a **2D tile**, either flatten `(row, col)` into a linear tile index first, or compute the address manually with `ptr[row * stride + col]` using your thread/block coordinates. +### `bits.h` — Compile-time bit helpers (`host::`) + +```cpp +#include +``` + +`constexpr` wrappers over ``, usable in `static_assert` and in template arguments: `host::is_pow2(x)`, `log2_floor(x)`, `log2_ceil(x)` (both `-1` for `x == 0`), `round_up_pow2(x)`, `round_down_pow2(x)`. They live in `host::` but are constant-expression-only, so device-side `static_assert` can call them. + ### `math.cuh` — Device math (`device::math::`) ```cpp @@ -185,8 +203,19 @@ For a **2D tile**, either flatten `(row, col)` into a linear tile index first, o #include ``` -- `device::warp::reduce(value, active_mask)` — generic warp reduction via `__shfl_xor_sync`. `Op` is a `device::ReductionOp` (`SUM`/`MAX`/`MIN`); `kNumThreads` is a power-of-two group size (default 32 = full warp); `kInner=true` (default) reduces within each `kNumThreads`-sized group, `kInner=false` reduces across groups (lanes at the same offset in different groups). -- `device::warp::reduce_sum/reduce_max/reduce_min(value)` — convenience wrappers over `reduce`. Work for any type with a `ReductionTrait`: floats, integers, and packed x2 types. +**Vectorized row copy — reach for this before hand-rolling.** A warp cooperatively moving one contiguous row is the single most common shape in this tree, and it used to be re-implemented per kernel. + +- **`warp::load_bytes(src)`** / **`warp::store_bytes(dst, val)`** — the whole warp moves `kBytes` from/to one contiguous row. `kBytes` is arbitrary, down to 1; the helper picks the widest vector that divides it and handles the ragged tail. The returned value is opaque and encodes the width, so a load and its matching store **must use the same ``**. +- **`warp::LoadStorePattern`** — the `kPattern` values, and the sign carries the meaning. A **negative** one (`WARP_UNIFORM_16B` and friends: `_GMEM`, `_SMEM`, `_4B`/`_8B`/`_32B`) says *the whole warp splits this row*, so the width is derived from the per-lane share (`gcd(kBytes / 32, |kPattern|)`), falling back to `gcd(kBytes, 4)` when the row is too narrow to give every lane 4 bytes. A **positive** one (a bare `LoadStoreBytes` value) makes no warp-splitting assumption and just caps the per-thread vector at `gcd(kBytes, kPattern)`. `WARP_UNIFORM_16B` is the usual choice for a row copy. +- **`LoadStorePattern::get_vec_bytes()`** — the width actually chosen. `SGL_DEVICE_HOST`, so **the launcher must call this to feed `.ensure_alignment(...)`**. Pass it the same value the kernel passes `load_bytes` — for a split row that is the *per-warp share*, not the whole row; the full row can resolve to a different width and would under-constrain the strides. + +**Reductions.** + +- `warp::reduce(value, active_mask)` — `Op` is a `device::ReductionOp` (`SUM`/`MAX`/`MIN`). `kStart` and `kFinish` bound a range of lane-index bits: `` reduces within contiguous groups of `N` lanes, `` reduces *across* groups at the same offset. The operation is **symmetric** — `` and `` reduce the same lane set. +- `warp::reduce_sum/reduce_max/reduce_min(value)` — wrappers. Work for any type with a `ReductionTrait`: floats, integers, packed x2. This is more widely used in practice than generic `reduce`. +- `warp::inclusive_reduce(val, lane_id, mask)` and `inclusive_sum` / `inclusive_max` / `inclusive_min` — segmented inclusive scan: every lane keeps its own running total, unlike `reduce`. Forward when `kStart < kFinish`, backward when `kStart > kFinish`. `lane_id` defaults to `get_lane_id()`, which is correct by construction — only pass it if you already have the **segment-relative** index (`threadIdx.x % kWidth`); a warp-relative one silently corrupts every segment past the first when `kWidth < 32`. +- `warp::broadcast(value, src_lane)` — `src_lane` is **segment-relative**: each `kWidth` segment reads its own lane, not one warp-wide source. +- `warp::get_lane_id`, `warp::elect_one_lane()` (one elected lane, for gating a single-thread TMA issue; CUDA-only), `warp::kFullMask` and `warp::kFullWidth` (32 on CUDA, 64 on HIP — note this differs from `kWarpThreads`). ### `cta.cuh` — CTA-level primitives @@ -203,6 +232,11 @@ For a **2D tile**, either flatten `(row, col)` into a linear tile index first, o ``` - `device::atomic::max(float* addr, float value)` — float atomic max (handles negative values correctly via bit tricks). +- **`device::atomic::Event`** — a cross-CTA arrive/wait counter in one 32-bit word. Producers call `arrive()`; consumers call `wait(num_producers)` (exactly one consumer) or `wait_multi(num_producers, num_consumers, n)` (several). The last released consumer resets the word, so one `Event` is reusable across launches with no host re-zero. Preconditions, all of them load-bearing: + - the word must be **zeroed by the host** before first use, and must live in **global** memory — a shared or local `Event` is an illegal address; + - `arrive()` is a release and the waits are acquires, so a producer's writes before `arrive()` are visible after the wait; + - **generations must be explicitly ordered** — the word has no phase bit, so overlapping generation N+1 with a consumer still in generation N is undefined behaviour (a lagging consumer never exits); + - `wait()` and `wait_multi()` lay the word out incompatibly; mixing them on one `Event` is undefined behaviour. ### `runtime.cuh` — Occupancy and device info @@ -211,8 +245,10 @@ For a **2D tile**, either flatten `(row, col)` into a linear tile index first, o ``` - `host::runtime::get_blocks_per_sm(kernel, block_dim)` — max active blocks per SM (occupancy) -- `host::runtime::get_sm_count(device_id)` — number of SMs on the device -- `host::runtime::get_cc_major(device_id)` — compute capability major version +- `host::runtime::get_sm_count(device_id [, use_cache])` — number of SMs on the device. Memoized per device ordinal, so it is cheap enough to call on a launch path; pass `use_cache=false` to force a driver query. +- `host::runtime::get_cc_major` / `get_cc_minor` / `get_sm_version` — compute capability, same caching. + +**Do not query the architecture at runtime.** The JIT compiles for the exact local GPU, so the arch is a *compile-time* fact: `SGL_CUDA_ARCH` is injected by `load_jit`, and `SGL_ARCH_HOPPER_OR_GREATER` / `SGL_ARCH_BLACKWELL_OR_GREATER` / `device::kMaxVecBytes` are derived from it. Reach for `get_cc_*` only for something the arch genuinely does not determine. SM *count* is the opposite case — it varies within an arch (B200 vs B300 vs a MIG slice), so it has to stay a runtime query. **Persistent kernel pattern** (cap blocks to SM count × occupancy): ```cpp @@ -586,7 +622,7 @@ Benchmarks use the project's own `marker` framework (in `python/sglang/kernels/j - `graph_clone_args` / `graph_clone_kwargs`: which inputs to clone per CUDA-graph iteration to defeat L2 cache reuse. Defaults to `"all"` — pass an iterable of indices/keys to limit to the *read* args (writes don't need cloning). - `use_cuda_graph=False` for kernels that can't be captured. - `metrics=(0.5, "avg")` controls reported quantiles (the first metric becomes the table latency column). - - `disable_log_bandwidth` (defaults from `SGLANG_KERNEL_DISABLE_LOG_BANDWIDTH=1`) skips the bandwidth column entirely. + - `disable_log_bandwidth` (defaults from `SGLANG_JIT_BENCHMARK_DISABLE_LOG_BANDWIDTH=1`) skips the bandwidth column entirely. - **`utils.create_random(*shape)` / `utils.create_empty(*shape)`** — shorthand for `torch.randn` / `torch.empty` with `DEFAULT_DTYPE` (`bfloat16`) and `DEFAULT_DEVICE` (`"cuda"`). Override via the `dtype=` / `device=` kwargs. - **`utils.get_benchmark_range(full_range, ci_range)`** — returns the smaller `ci_range` under CI (`is_in_ci()`), the `full_range` locally. Still available for the `benchmark(...)` column axis (which has no `ci_vals`); for `parametrize` row axes prefer the built-in `ci_vals` argument. @@ -640,7 +676,7 @@ if __name__ == "__main__": - The `line_arg` name passed to `benchmark` (`"impl"` here) must match a parameter on `benchmark(...)`; same for every `parametrize` name (`"size"`). - Stack `@parametrize` once per swept axis. The required `@marker.benchmark` is the **innermost** decorator (bottom of the stack, directly above the function) — `@parametrize` rows go above it. - Prefer `create_random` / `create_empty` from `utils.py` over open-coding `torch.randn(..., dtype=..., device=...)`. -- The GB/s column appears by default (`memory_args="all"` + `memory_output="out"`). For memory-bound kernels it's the most informative number; scope `memory_args` / `memory_output` to the tensors actually touched if the defaults over- or under-count. For compute-bound kernels where bandwidth is misleading, set `SGLANG_KERNEL_DISABLE_LOG_BANDWIDTH=1` (or `disable_log_bandwidth=True`). +- The GB/s column appears by default (`memory_args="all"` + `memory_output="out"`). For memory-bound kernels it's the most informative number; scope `memory_args` / `memory_output` to the tensors actually touched if the defaults over- or under-count. For compute-bound kernels where bandwidth is misleading, set `SGLANG_JIT_BENCHMARK_DISABLE_LOG_BANDWIDTH=1` (or `disable_log_bandwidth=True`). - For in-place kernels (which return `None`), pass the written tensors via `memory_output=(...)` since the `"out"` default would capture nothing. - Tune `graph_clone_args` / `graph_clone_kwargs` to all the arguments that might be read by the kernel. We can only skip cloning for write-only args. For in-place modified args, we still need to clone them to get accurate timing (reusing the same buffer keeps it L2-hot and skews results). - Call `benchmark.run()` (no `print_data=` kwarg — the marker framework prints directly). @@ -665,7 +701,7 @@ cd test && python3 run_suite.py --hw cuda --suite base-b-kernel-benchmark-test-1 - **JIT compilation fails**: ensure the `.cuh` file is under `python/sglang/kernels/jit/csrc/`; reduce template argument combinations - **CUDA crash / illegal memory access**: `CUDA_LAUNCH_BLOCKING=1`; `compute-sanitizer --tool memcheck python ...` - **Unstable benchmark results**: `marker.do_bench` uses CUDA-graph-based timing by default; set `use_cuda_graph=False` only if the kernel can't be captured. `graph_clone_args` defaults to `"all"`; if you narrow it, it must still cover every *read* tensor — reusing a single buffer keeps it L2-hot and skews results. Keep *write* tensors in it too: they are what sets the rotation count, and a shared output buffer stays L2-hot the same way. -- **Missing GB/s column**: the column is on by default; check that `SGLANG_KERNEL_DISABLE_LOG_BANDWIDTH` is not `1` and `disable_log_bandwidth` is not `True`. For in-place kernels (return `None`) the `memory_output="out"` default counts nothing — pass the written tensors via `memory_output=(...)` +- **Missing GB/s column**: the column is on by default; check that `SGLANG_JIT_BENCHMARK_DISABLE_LOG_BANDWIDTH` is not `1` and `disable_log_bandwidth` is not `True`. For in-place kernels (return `None`) the `memory_output="out"` default counts nothing — pass the written tensors via `memory_output=(...)` --- @@ -681,6 +717,7 @@ cd test && python3 run_suite.py --hw cuda --suite base-b-kernel-benchmark-test-1 - `python/sglang/kernels/jit/include/sgl_kernel/utils.cuh` — type aliases, `LaunchKernel`, `SGL_DEVICE` - `python/sglang/kernels/jit/include/sgl_kernel/vec.cuh` — `AlignedVector` - `python/sglang/kernels/jit/include/sgl_kernel/tile.cuh` — `tile::Memory` +- `python/sglang/kernels/jit/include/sgl_kernel/bits.h` — compile-time bit helpers - `python/sglang/kernels/jit/include/sgl_kernel/type.cuh` — `DTypeTrait`, `packed_t`, `device::cast`, `device::unpack`, `ReductionTrait` - `python/sglang/kernels/jit/include/sgl_kernel/math.cuh` — `device::math::` - `python/sglang/kernels/jit/include/sgl_kernel/warp.cuh` — `warp::reduce` and `reduce_sum/max/min` wrappers diff --git a/python/sglang/kernels/jit/benchmark/marker.py b/python/sglang/kernels/jit/benchmark/marker.py index 71bf3fd9b..11d582bbf 100644 --- a/python/sglang/kernels/jit/benchmark/marker.py +++ b/python/sglang/kernels/jit/benchmark/marker.py @@ -1,8 +1,8 @@ +import builtins import contextlib import inspect import itertools import math -import os from typing import ( Any, Callable, @@ -22,6 +22,7 @@ from typing import ( import torch from sglang.kernels.jit.utils import cache_once +from sglang.srt.environ import envs from sglang.utils import is_in_ci F = TypeVar("F", bound=Callable[..., "BenchResult"]) @@ -29,8 +30,9 @@ Metric: TypeAlias = "float | Literal['avg']" BENCH_CONFIG: TypeAlias = "List[Tuple[Tuple[str, ...], List[Tuple[Any, ...]]]]" UNIT_SCALE = {"us": 1e-6, "ms": 1e-3, "s": 1.0} TYPE_LIST = (bool, int, float, str, torch.dtype, torch.device, None.__class__) -DISABLE_LOG_BANDWIDTH = os.environ.get("SGLANG_KERNEL_DISABLE_LOG_BANDWIDTH") == "1" - +DISABLE_LOG_BANDWIDTH = envs.SGLANG_JIT_BENCHMARK_DISABLE_LOG_BANDWIDTH.get() +DISABLE_LOG_FLOPS = envs.SGLANG_JIT_BENCHMARK_DISABLE_LOG_FLOPS.get() +PATTERN: TypeAlias = "Literal['pow2']" __all__ = [ "BenchResult", @@ -40,6 +42,7 @@ __all__ = [ "parametrize", "do_bench", "skip", + "range", ] @@ -153,6 +156,7 @@ class BenchResult(NamedTuple): metrics: Tuple[Metric, ...] times: List[float] # in seconds memory_footprint: Optional[int] + flops: Optional[float] = None class Table: @@ -172,12 +176,7 @@ class Table: def format_latency(r: float) -> str: if math.isnan(r): return "N/A" - length = len(str(int(r))) - if length < 5: - return f"{r:.4f}" - # decrease number of the digits - digits = max(0, 4 - (length - 5)) - return f"{r:.{digits}f}" + return f"{r:.4f}" @staticmethod def format_bandwidth(b: float) -> str: @@ -185,6 +184,12 @@ class Table: return "N/A" return f"{b:.2f}" + @staticmethod + def format_flops(f: float) -> str: + if math.isnan(f): + return "N/A" + return f"{f:.3f}" + def col( self, header: str = "", @@ -205,7 +210,7 @@ class Table: assert len(cells) == len(self._headers) self._rows.append([str(c) for c in cells]) - def print(self) -> None: + def print(self, prefix: Optional[str], suffix: Optional[str]) -> None: widths = [ max(max(len(c) + p for c in [h, *(r[i] for r in self._rows)]), mw) for i, (h, mw, p) in enumerate(zip(self._headers, self._mins, self._pads)) @@ -220,12 +225,18 @@ class Table: parts.append(f"{cell:{a}{w}}") return "".join(parts) + if prefix is not None: + print("=" * total) + print(prefix) print("=" * total) print(fmt(self._headers)) print("-" * total) for r in self._rows: print(fmt(r)) print("=" * total) + if suffix is not None: + print(suffix) + print("=" * total) class Benchmark(Generic[F]): @@ -259,15 +270,18 @@ class Benchmark(Generic[F]): self._seen_args.add(name) self._configs.insert(0, (names, vals)) - def _collect_results(self) -> Tuple[List[List[float]], List[List[float]], bool]: + def _collect_results(self): axis_names = [n for n, _ in self._configs] axis_vals = [v for _, v in self._configs] results: List[List[float]] = [] bandwidth_results: List[List[float]] = [] + flops_results: List[List[float]] = [] should_log_bandwidth = False + should_log_flops = False for system in self._line_vals: latencies: List[float] = [] bandwidths: List[float] = [] + flops: List[float] = [] for combo in itertools.product(*axis_vals): kwargs: Dict[str, Any] = {self._line_arg: system} for names, values in zip(axis_names, combo): @@ -278,6 +292,8 @@ class Benchmark(Generic[F]): latencies.append(float("nan")) if not DISABLE_LOG_BANDWIDTH: bandwidths.append(float("nan")) + if not DISABLE_LOG_FLOPS: + flops.append(float("nan")) continue except BaseException: print(f"Benchmark failed at {system}, kwargs =", kwargs) @@ -288,11 +304,26 @@ class Benchmark(Generic[F]): bandwidths.append( result.memory_footprint / (1024**3) / result.times[0] ) + if not DISABLE_LOG_FLOPS and result.flops is not None: + should_log_flops = True + flops.append(result.flops / (1e12) / result.times[0]) results.append(latencies) bandwidth_results.append(bandwidths) - return results, bandwidth_results, should_log_bandwidth + flops_results.append(flops) + return ( + results, + bandwidth_results, + flops_results, + should_log_bandwidth, + should_log_flops, + ) - def run(self) -> None: + def run( + self, + *, + print_prefix: Optional[str] = None, + print_suffix: Optional[str] = None, + ) -> None: # Pre-check: every required fn param must be covered. flat_names = [n for names, _ in self._configs for n in names] kinds = ( @@ -308,7 +339,9 @@ class Benchmark(Generic[F]): f"parameters not parametrized for {self._fn.__name__}: {sorted(missing)}" ) - results, bandwidths, should_log_bw = self._collect_results() + results, bandwidths, flops, should_log_bw, should_log_flops = ( + self._collect_results() + ) table = Table() table.col(min_width=0, pad=0, align="<") # id column (tight, left-aligned) @@ -317,21 +350,34 @@ class Benchmark(Generic[F]): table.sep() for system in self._line_vals: table.col(f"{system}({self._unit})", min_width=15) + # one entry per row, per system -- guards the skip/append paths above + row_count = math.prod(len(vals) for _, vals in self._configs) if should_log_bw: table.sep() for system in self._line_vals: table.col(f"{system}(GB/s)", min_width=15) + assert all(len(b) == row_count for b in bandwidths) + if should_log_flops: + table.sep() + for system in self._line_vals: + table.col(f"{system}(TFLOPS)", min_width=15) + assert all(len(f) == row_count for f in flops) axis_vals = [v for _, v in self._configs] for row_id, combo in enumerate(itertools.product(*axis_vals)): + # skip entries that are skipped by all systems + if all(math.isnan(r[row_id]) for r in results): + continue cells: List[Any] = [row_id] cells.extend(v for vt in combo for v in vt) cells.extend(table.format_latency(r[row_id]) for r in results) if should_log_bw: cells.extend(table.format_bandwidth(r[row_id]) for r in bandwidths) + if should_log_flops: + cells.extend(table.format_flops(r[row_id]) for r in flops) table.row(*cells) - table.print() + table.print(print_prefix, print_suffix) def benchmark(line_arg: str, line_vals: List[Any], *, unit: str = "us"): @@ -402,14 +448,14 @@ def _do_bench_internal_graph( graph = torch.cuda.CUDAGraph() # NOTE: we rotate the buffer here to avoid L2 cache effect - for i in range(1, rotate_count): + for i in builtins.range(1, rotate_count): input_args_list[i] = tuple( ( _clone_recursive(input_args[j]) if j in graph_clone_args else input_args[j] ) - for j in range(len(input_args)) + for j in builtins.range(len(input_args)) ) input_kwargs_list[i] = dict( (k, (_clone_recursive(v) if k in graph_clone_kwargs else v)) @@ -417,7 +463,7 @@ def _do_bench_internal_graph( ) with graph_context: with torch.cuda.graph(graph, stream=stream): - for i in range(loop_count): + for i in builtins.range(loop_count): args = input_args_list[i % rotate_count] kwargs = input_kwargs_list[i % rotate_count] fn(*args, **kwargs) @@ -427,7 +473,7 @@ def _do_bench_internal_graph( # then replay the graph and measure the time tic = torch.cuda.Event(enable_timing=True) toc = torch.cuda.Event(enable_timing=True) - for _ in range(max(replay_iters // loop_count, 10)): + for _ in builtins.range(max(replay_iters // loop_count, 10)): empty_tensor.zero_() # cold the L2 cache sync_multigpu_fn() # sync GPU before each iteration for precise timing tic.record(stream) @@ -444,7 +490,7 @@ def do_bench( input_args: Tuple[Any, ...] = (), input_kwargs: Dict[str, Any] = {}, use_cuda_graph: bool = True, - warmup_iters: int = 50, + warmup_iters: int = 20, replay_iters: int = 1000, metrics: Tuple[Metric, ...] = (0.5, "avg"), stream: torch.cuda.Stream | None = None, @@ -457,8 +503,10 @@ def do_bench( memory_output: Iterable[Any] | Literal["out"] | None = "out", extra_memory_args: Iterable[Any] | None = None, extra_memory_footprint: int = 0, + flops: Optional[float] = None, graph_context_fn: Optional[Callable[[], ContextManager]] = None, sync_multigpu_fn: Optional[Callable[[], Any]] = None, + estimated_time_ms: Optional[float] = 1.0, ) -> BenchResult: """ Benchmark a function using CUDA graph or naive loop. @@ -484,10 +532,15 @@ def do_bench( :param extra_memory_args: Additional arguments to consider for memory footprint calculation. :param extra_memory_footprint: Additional memory footprint to consider. This is typically used when the load/store bytes is dynamic. + :param flops: The number of floating-point operations performed by the benchmark. + Used for calculating the achieved computation utilization in the profile report. :param graph_context_fn: A callable returning a context manager that wraps the cuda graph capture. :param sync_multigpu_fn: A callable to synchronize multiple GPUs before each iteration. For precise benchmark number in multi-GPU benchmark, it should be some synchronization primitive on GPU side (not on CPU side). + :param estimated_time_ms: Estimated time in milliseconds for the benchmark without CUDA graph. + This is typically used to control the benchmark time. + Ignored if `use_cuda_graph` is True. """ # first warmup the function device_id = torch.cuda.current_device() @@ -499,12 +552,12 @@ def do_bench( with torch.cuda.device(device_id), torch.cuda.stream(stream): stream.wait_stream(old_current_stream) sync_multigpu_fn() - for _ in range(warmup_iters): + for _ in builtins.range(warmup_iters): fn(*input_args, **input_kwargs) if use_cuda_graph: # NOTE: by default, reduce all the CPU-side overhead if graph_clone_args == "all": - graph_clone_args = range(len(input_args)) + graph_clone_args = builtins.range(len(input_args)) elif graph_clone_args is None: graph_clone_args = [] if graph_clone_kwargs == "all": @@ -531,7 +584,18 @@ def do_bench( tic = torch.cuda.Event(enable_timing=True) toc = torch.cuda.Event(enable_timing=True) empty_tensor = _get_flush_l2_buffer() - for _ in range(max(replay_iters, 10)): + if estimated_time_ms is not None: + empty_tensor.zero_() # cold the L2 cache + sync_multigpu_fn() + tic.record(stream) + fn(*input_args, **input_kwargs) + toc.record(stream) + stream.synchronize() + duration_ms = tic.elapsed_time(toc) + estimted_iters = int(estimated_time_ms / duration_ms) + replay_iters = min(replay_iters, estimted_iters) + + for _ in builtins.range(max(replay_iters, 10)): empty_tensor.zero_() # cold the L2 cache sync_multigpu_fn() tic.record(stream) @@ -553,4 +617,10 @@ def do_bench( memory_footprint += _get_nbytes_recursive(memory_args) memory_footprint += _get_nbytes_recursive(memory_output) - return BenchResult(metrics, result, memory_footprint) + return BenchResult(metrics, result, memory_footprint, flops) + + +def range(*args: int, pattern: PATTERN) -> List[int]: + # TODO: support other patterns + assert pattern == "pow2", f"unsupported pattern: {pattern}" + return [2**i for i in builtins.range(*args)] diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/c_plan.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/c_plan.cuh index 45ff0e203..2d1419a94 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/c_plan.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/c_plan.cuh @@ -107,20 +107,6 @@ struct DecodeParamsLegacy { inline constexpr uint32_t kMaxPrefillBatchSize = 1024; -SGL_DEVICE uint32_t warp_inclusive_sum(uint32_t lane_id, uint32_t val) { - static_assert(device::kWarpThreads == 32); -#pragma unroll - for (uint32_t offset = 1; offset < 32; offset *= 2) { -#ifndef USE_ROCM - uint32_t n = __shfl_up_sync(device::kFullMask, val, offset); -#else - uint32_t n = __shfl_up(val, offset, 32); -#endif - if (lane_id >= offset) val += n; - } - return val; -} - __global__ __launch_bounds__(1024, 1) // void plan_compress_prefill_kernel0(const Prefill0Params params) { using namespace device; diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/candidate_block_table.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/candidate_block_table.cuh index 3cffe255c..3a9ca8008 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/candidate_block_table.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/candidate_block_table.cuh @@ -124,7 +124,7 @@ __global__ __launch_bounds__(CandidateBlockTableConfig::kBlockSize, CandidateBlo count[j] = __popc(words[j]); local += count[j]; } - const auto warp_inc = warp::inclusive_sum(lane_id, local); + const auto warp_inc = warp::inclusive_sum(local, lane_id); if (lane_id == kWarpThreads - 1) smem.warp_sum[warp_id] = warp_inc; __syncthreads(); // also: every thread holds its words, the bitmap may become the queue const auto peer_sum = smem.warp_sum[lane_id]; diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh index f05170817..6105fc36d 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh @@ -181,7 +181,7 @@ INDEXER_KERNEL void fused_norm_rope_indexer(const __grid_constant__ FusedNormRop #pragma unroll for (int i = 0; i < kVecSize; ++i) { #ifndef USE_ROCM - const float other = __shfl_xor_sync(kFullMask, data[i], mask, kWarpThreads); + const float other = __shfl_xor_sync(warp::kFullMask, data[i], mask, kWarpThreads); #else const float other = __shfl_xor(data[i], mask, kWarpThreads); #endif @@ -337,7 +337,7 @@ INDEXER_KERNEL void fused_norm_rope_indexer_fp4(const __grid_constant__ FusedNor #pragma unroll for (int i = 0; i < kVecSize; ++i) { #ifndef USE_ROCM - const float other = __shfl_xor_sync(kFullMask, data[i], mask, kWarpThreads); + const float other = __shfl_xor_sync(warp::kFullMask, data[i], mask, kWarpThreads); #else const float other = __shfl_xor(data[i], mask, kWarpThreads); #endif diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/silu_and_mul_masked_post_quant.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/silu_and_mul_masked_post_quant.cuh index fa8550782..1413eb1b0 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/silu_and_mul_masked_post_quant.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/silu_and_mul_masked_post_quant.cuh @@ -39,16 +39,6 @@ struct alignas(16) CTAWork { bool valid; }; -SGL_DEVICE uint32_t warp_inclusive_sum(uint32_t lane_id, uint32_t val) { - static_assert(device::kWarpThreads == 32); -#pragma unroll - for (uint32_t offset = 1; offset < 32; offset *= 2) { - uint32_t n = __shfl_up_sync(0xFFFFFFFF, val, offset); - if (lane_id >= offset) val += n; - } - return val; -} - template SGL_DEVICE fp32x2_t silu_and_mul(DType2 gate, DType2 up, float limit) { using namespace device; @@ -93,7 +83,7 @@ SGL_DEVICE CTAWork get_work(const SiluMulQuantVarlenParams& params) { const uint32_t val = tx < params.num_experts ? params.masked_m[tx] : 0u; // Per-warp inclusive scan of masked_m. - const uint32_t warp_inclusive = warp_inclusive_sum(lane_id, val); + const uint32_t warp_inclusive = warp::inclusive_sum(val, lane_id); const uint32_t warp_exclusive = warp_inclusive - val; // Write each warp total. diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/topk_bf16_small.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/topk_bf16_small.cuh index 37daad49d..cec06c854 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/topk_bf16_small.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/topk_bf16_small.cuh @@ -96,7 +96,7 @@ SGL_DEVICE uint32_t warp_exclusive_suffix_sum(uint32_t x, uint32_t lane_id) { uint32_t inc = x; #pragma unroll for (uint32_t offset = 1; offset < device::kWarpThreads; offset <<= 1) { - const auto t = __shfl_down_sync(device::kFullMask, inc, offset); + const auto t = __shfl_down_sync(device::warp::kFullMask, inc, offset); if (lane_id + offset < device::kWarpThreads) inc += t; } return inc - x; @@ -324,10 +324,10 @@ __global__ __launch_bounds__(TopKBF16Config::kBlockSize, TopKBF16Config::kOccupa // Block-wide exclusive prefix of (gt, eq), packed: one warp scan plus one shared atomic per // warp. Warps land in arrival order, which is fine since the output is unordered. const uint32_t local = cnt_gt << 16 | cnt_eq; - const uint32_t warp_inc = warp::inclusive_sum(lane_id, local); + const uint32_t warp_inc = warp::inclusive_sum(local, lane_id); uint32_t warp_base = 0; if (lane_id == kWarpThreads - 1) warp_base = atomicAdd(&smem.count_gt_eq, warp_inc); - warp_base = __shfl_sync(kFullMask, warp_base, kWarpThreads - 1); + warp_base = __shfl_sync(warp::kFullMask, warp_base, kWarpThreads - 1); const uint32_t before = warp_base + warp_inc - local; // Everything above the pivot is taken, plus `remain` of the elements equal to it. diff --git a/python/sglang/kernels/jit/csrc/elementwise/concat_mla.cuh b/python/sglang/kernels/jit/csrc/elementwise/concat_mla.cuh index 53bd08553..2bba8bf4e 100644 --- a/python/sglang/kernels/jit/csrc/elementwise/concat_mla.cuh +++ b/python/sglang/kernels/jit/csrc/elementwise/concat_mla.cuh @@ -2,9 +2,11 @@ #include #include +#include #include +#include #include #include @@ -138,14 +140,25 @@ struct ConcatMlaKKernel { D_nope.set_value(QK_NOPE_HEAD_DIM); D_rope.set_value(QK_ROPE_HEAD_DIM); + // The widest access is `v2.s32` (8B), and the `>> 2` stride stepping needs + // each row stride to be a whole number of int2's -- which is the same 8B. + constexpr int64_t kAlignNope = 8; + constexpr int64_t kAlignRope = 4; // rope is read with `v1.s32` + // Verify k: [num_tokens, num_heads, k_head_dim] - TensorMatcher({N, H, D}).with_strides({S0_k, S1_k, 1}).with_dtype().with_device(device).verify(k); + TensorMatcher({N, H, D}) + .with_strides({S0_k, S1_k, 1}) + .with_dtype() + .with_device(device) + .ensure_alignment(kAlignNope) + .verify(k); // Verify k_nope: [num_tokens, num_heads, nope_head_dim] TensorMatcher({N, H, D_nope}) .with_strides({S0_k_nope, S1_k_nope, 1}) .with_dtype() .with_device(device) + .ensure_alignment(kAlignNope) .verify(k_nope); // Verify k_rope: [num_tokens, 1, rope_head_dim] @@ -153,13 +166,9 @@ struct ConcatMlaKKernel { .with_strides({S0_k_rope, -1, 1}) .with_dtype() .with_device(device) + .ensure_alignment(kAlignRope) .verify(k_rope); - // Check alignment - RuntimeCheck(reinterpret_cast(k.data_ptr()) % 16 == 0, "Tensor k must be 16-byte aligned"); - RuntimeCheck(reinterpret_cast(k_nope.data_ptr()) % 16 == 0, "Tensor k_nope must be 16-byte aligned"); - RuntimeCheck(reinterpret_cast(k_rope.data_ptr()) % 16 == 0, "Tensor k_rope must be 16-byte aligned"); - const int num_tokens = static_cast(N.unwrap()); constexpr int num_warps_per_block = 32; @@ -188,8 +197,8 @@ constexpr int OUT_LAST_DIM = A_LAST_DIM + B_LAST_DIM; template __global__ void concat_mla_absorb_q_kernel( - bf16_t* a, - bf16_t* b, + const bf16_t* a, + const bf16_t* b, bf16_t* out, const int num_items, const int dim_1, @@ -199,51 +208,27 @@ __global__ void concat_mla_absorb_q_kernel( const int b_stride_1, const int64_t out_stride_0, const int out_stride_1) { - device::PDLWaitPrimary(); + using namespace device; + using enum warp::LoadStorePattern::type; + constexpr int64_t kABytes = A_LAST_DIM * sizeof(bf16_t); + constexpr int64_t kBBytes = B_LAST_DIM * sizeof(bf16_t); - const int flat_warp_id = (blockIdx.x * blockDim.x + threadIdx.x) / 32; - const int lane_id = get_lane_id(); + PDLWaitPrimary(); + + const int flat_warp_id = (blockIdx.x * blockDim.x + threadIdx.x) / kWarpThreads; + if (flat_warp_id >= num_items) return; const int idx_0 = flat_warp_id / dim_1; const int idx_1 = flat_warp_id % dim_1; + const auto out_row = out + idx_0 * out_stride_0 + idx_1 * out_stride_1; - if (flat_warp_id >= num_items) { - return; - } + const auto b_val = warp::load_bytes(b + idx_0 * b_stride_0 + idx_1 * b_stride_1); + const auto a_val = warp::load_bytes(a + idx_0 * a_stride_0 + idx_1 * a_stride_1); - using ABufType = int4; - constexpr int A_NUM_UNROLL = 2; - static_assert(sizeof(ABufType) * A_NUM_UNROLL == A_LAST_DIM * sizeof(a[0]) / 32); - ABufType a_buf[A_NUM_UNROLL]; + PDLTriggerSecondary(); - using BBufType = int; - constexpr int B_NUM_UNROLL = 1; - static_assert(sizeof(BBufType) * B_NUM_UNROLL == B_LAST_DIM * sizeof(b[0]) / 32); - BBufType b_buf; - - { - const BBufType* base_addr = reinterpret_cast(b + idx_0 * b_stride_0 + idx_1 * b_stride_1); - b_buf = *(base_addr + lane_id); - } - -#pragma unroll - for (int i = 0; i < A_NUM_UNROLL; ++i) { - const ABufType* base_addr = reinterpret_cast(a + idx_0 * a_stride_0 + idx_1 * a_stride_1); - a_buf[i] = *(base_addr + i * 32 + lane_id); - } - - device::PDLTriggerSecondary(); - - { - BBufType* base_addr = reinterpret_cast(out + idx_0 * out_stride_0 + idx_1 * out_stride_1 + A_LAST_DIM); - *(base_addr + lane_id) = b_buf; - } - -#pragma unroll - for (int i = 0; i < A_NUM_UNROLL; ++i) { - ABufType* base_addr = reinterpret_cast(out + idx_0 * out_stride_0 + idx_1 * out_stride_1); - *(base_addr + i * 32 + lane_id) = a_buf[i]; - } + warp::store_bytes(out_row + A_LAST_DIM, b_val); + warp::store_bytes(out_row, a_val); } template @@ -273,11 +258,21 @@ struct ConcatMlaAbsorbQKernel { D_b.set_value(B_LAST_DIM); D_out.set_value(OUT_LAST_DIM); + using device::warp::LoadStorePattern; + using enum LoadStorePattern::type; + constexpr int64_t kAlignA = LoadStorePattern::get_vec_bytes(); + constexpr int64_t kAlignB = LoadStorePattern::get_vec_bytes(); + // `out` carries both halves, so it must satisfy the wider one; the B half + // starts A_LAST_DIM in, which must not break B's own alignment. + constexpr int64_t kAlignOut = std::max(kAlignA, kAlignB); + static_assert(A_LAST_DIM * sizeof(bf16_t) % kAlignB == 0, "A width must not misalign the B half"); + // Verify a: [dim_0, dim_1, A_LAST_DIM] TensorMatcher({N0_a, N1_a, D_a}) .with_strides({S0_a, S1_a, 1}) .with_dtype() .with_device(device) + .ensure_alignment(kAlignA) .verify(a); // Verify b: [dim_0, dim_1, B_LAST_DIM] @@ -285,6 +280,7 @@ struct ConcatMlaAbsorbQKernel { .with_strides({S0_b, S1_b, 1}) .with_dtype() .with_device(device) + .ensure_alignment(kAlignB) .verify(b); // Verify out: [dim_0, dim_1, OUT_LAST_DIM] @@ -292,13 +288,9 @@ struct ConcatMlaAbsorbQKernel { .with_strides({S0_out, S1_out, 1}) .with_dtype() .with_device(device) + .ensure_alignment(kAlignOut) .verify(out); - // Check alignment - RuntimeCheck(reinterpret_cast(a.data_ptr()) % 16 == 0, "Tensor a must be 16-byte aligned"); - RuntimeCheck(reinterpret_cast(b.data_ptr()) % 16 == 0, "Tensor b must be 16-byte aligned"); - RuntimeCheck(reinterpret_cast(out.data_ptr()) % 16 == 0, "Tensor out must be 16-byte aligned"); - // Verify dimensions match: a.size(0) * a.size(1) == b.size(0) * b.size(1) RuntimeCheck( N0_a.unwrap() * N1_a.unwrap() == N0_b.unwrap() * N1_b.unwrap(), @@ -315,8 +307,8 @@ struct ConcatMlaAbsorbQKernel { LaunchKernel(grid_size, block_size, device.unwrap()) .enable_pdl(kUsePDL)( concat_mla_absorb_q_kernel, - static_cast(a.data_ptr()), - static_cast(b.data_ptr()), + static_cast(a.data_ptr()), + static_cast(b.data_ptr()), static_cast(out.data_ptr()), num_items, dim_1, diff --git a/python/sglang/kernels/jit/csrc/elementwise/kvcache.cuh b/python/sglang/kernels/jit/csrc/elementwise/kvcache.cuh index 818b010c5..b17168466 100644 --- a/python/sglang/kernels/jit/csrc/elementwise/kvcache.cuh +++ b/python/sglang/kernels/jit/csrc/elementwise/kvcache.cuh @@ -1,9 +1,10 @@ #include #include -#include +#include #include #include +#include #include #include @@ -30,169 +31,30 @@ struct StoreKVCacheParams { int64_t reserved_skip_index; }; -constexpr uint32_t kNumWarps = 4; -constexpr uint32_t kThreadsPerBlock = kNumWarps * device::kWarpThreads; - -/** - * \brief How a warp vectorizes one row of kElementBytes: the widest aligned - * vector type it can use, and how many full loop iterations that takes. - * Shared by the interleaved and single-row copies so the two cannot drift. - * kElementBytes == 0 is a valid (empty) plan, so a zero-width tail can be - * queried before being branched away. - */ -template -struct RowVecPlan { - static constexpr int64_t kAlignment = (kElementBytes % (16 * device::kWarpThreads) == 0) ? 16 - : kElementBytes % (8 * device::kWarpThreads) == 0 ? 8 - : kElementBytes % (4 * device::kWarpThreads) == 0 ? 4 - : kElementBytes % 4 == 0 ? 4 - : 0; - - static_assert(kAlignment > 0, "Element size must be multiple of 4 bytes"); - - using vec_t = device::AlignedStorage; - static constexpr int64_t kLoopBytes = sizeof(vec_t) * device::kWarpThreads; - static constexpr int64_t kLoopCount = kElementBytes / kLoopBytes; - static constexpr int64_t kElementCount = kElementBytes / sizeof(vec_t); - static constexpr bool kHasEpilogue = kLoopCount * kLoopBytes < kElementBytes; -}; - -/** - * \brief Use a single warp to copy key and value data from source to destination. - * Each thread in the warp copies a portion of the data in a coalesced manner. - * Both loads are issued before either store: the two rows live in different - * tensors, and the params' __restrict__ does not survive into the kernel body, - * so the compiler cannot prove k_dst and v_src disjoint and will not sink the - * V load past the K store on its own. - * \tparam kElementBytes The size of each key/value element in bytes. - * \param k_src Pointer to the source key data. - * \param v_src Pointer to the source value data. - * \param k_dst Pointer to the destination key data. - * \param v_dst Pointer to the destination value data. - */ -template -SGL_DEVICE void copy_kv_warp( - const void* __restrict__ k_src, - const void* __restrict__ v_src, - void* __restrict__ k_dst, - void* __restrict__ v_dst) { - using namespace device; - using plan_t = RowVecPlan; - using vec_t = typename plan_t::vec_t; - constexpr auto kLoopCount = plan_t::kLoopCount; - - const auto gmem = tile::Memory::warp(); - -#pragma unroll kLoopCount - for (int64_t i = 0; i < kLoopCount; ++i) { - const auto k = gmem.load(k_src, i); - const auto v = gmem.load(v_src, i); - gmem.store(k_dst, k, i); - gmem.store(v_dst, v, i); - } - - // handle the epilogue if any - if constexpr (plan_t::kHasEpilogue) { - if (gmem.in_bound(plan_t::kElementCount, kLoopCount)) { - const auto k = gmem.load(k_src, kLoopCount); - const auto v = gmem.load(v_src, kLoopCount); - gmem.store(k_dst, k, kLoopCount); - gmem.store(v_dst, v, kLoopCount); - } - } -} - -/** - * \brief Use a single warp to copy one row from source to destination. - * Serves the width by which asymmetric K/V rows differ, which has no counterpart - * row to interleave with. - * \tparam kElementBytes The size of the row in bytes. - * \param src Pointer to the source data. - * \param dst Pointer to the destination data. - */ -template -SGL_DEVICE void copy_row_warp(const void* __restrict__ src, void* __restrict__ dst) { - using namespace device; - using plan_t = RowVecPlan; - using vec_t = typename plan_t::vec_t; - constexpr auto kLoopCount = plan_t::kLoopCount; - - const auto gmem = tile::Memory::warp(); - -#pragma unroll kLoopCount - for (int64_t i = 0; i < kLoopCount; ++i) { - gmem.store(dst, gmem.load(src, i), i); - } - - // handle the epilogue if any - if constexpr (plan_t::kHasEpilogue) { - if (gmem.in_bound(plan_t::kElementCount, kLoopCount)) { - gmem.store(dst, gmem.load(src, kLoopCount), kLoopCount); - } - } -} - -/** - * \brief Copy a K row of kKBytes and a V row of kVBytes with one warp. - * The overlapping prefix goes through the interleaved copy; only the width by - * which the rows differ is left as a serial tail. Equal widths degenerate to a - * single interleaved copy with no tail. - */ -template -SGL_DEVICE void copy_kv_rows_warp( - const void* __restrict__ k_src, - const void* __restrict__ v_src, - void* __restrict__ k_dst, - void* __restrict__ v_dst) { - using namespace device; - constexpr auto kCommon = kKBytes < kVBytes ? kKBytes : kVBytes; - constexpr auto kTail = (kKBytes < kVBytes ? kVBytes : kKBytes) - kCommon; - - // The interleaved copy indexes BOTH rows with kCommon's vector width, so that - // width must divide each row's split offset -- the narrower row's alignment - // does not imply the wider one's (e.g. 512 picks 16B, but 516 is not 16B - // aligned). The tail's own width must likewise divide its kCommon start. - // Whatever these gates admit is alignment-safe for the strides too, since a - // stride is a whole multiple of its split size. - constexpr auto kTailOrCommon = kTail == 0 ? kCommon : kTail; - constexpr auto kCommonAlign = RowVecPlan::kAlignment; - constexpr auto kTailAlign = RowVecPlan::kAlignment; - constexpr bool kCanInterleave = - kKBytes % kCommonAlign == 0 && kVBytes % kCommonAlign == 0 && kCommon % kTailAlign == 0; - - if constexpr (kCanInterleave) { - copy_kv_warp(k_src, v_src, k_dst, v_dst); - if constexpr (kTail > 0) { - if constexpr (kKBytes > kVBytes) { - copy_row_warp(pointer::offset(k_src, kCommon), pointer::offset(k_dst, kCommon)); - } else { - copy_row_warp(pointer::offset(v_src, kCommon), pointer::offset(v_dst, kCommon)); - } - } - } else { - copy_row_warp(k_src, k_dst); - copy_row_warp(v_src, v_dst); - } -} - /** * \brief Kernel to store key-value pairs into the KV cache. * Each element is split into multiple parts to allow parallel memory copy. - * \tparam kKElementBytes The size of each key element in bytes. - * \tparam kVElementBytes The size of each value element in bytes. Differs from - * kKElementBytes for asymmetric KV (head_dim != v_head_dim). - * \tparam kSplit The number of warps that handle each element. + * \tparam kKBytes The size of each key element in bytes. + * \tparam kVBytes The size of each value element in bytes. + * \tparam kNumThreads Threads cooperating on one KV item; a multiple of the + * warp size. The block shape is chosen at launch, independently. * \tparam kUsePDL Whether to use PDL feature. - * \tparam T The data type of the indices (`int32_t` or `int64_t`). + * \tparam TLoc The data type of the indices (`int32_t` or `int64_t`). */ -template -__global__ void store_kvcache(const __grid_constant__ StoreKVCacheParams params) { +template +__global__ void store_kvcache_kernel(const __grid_constant__ StoreKVCacheParams params) { using namespace device; - constexpr auto kKSplitSize = kKElementBytes / kSplit; - constexpr auto kVSplitSize = kVElementBytes / kSplit; - const uint32_t warp_id = blockIdx.x * kNumWarps + threadIdx.x / kWarpThreads; - const uint32_t item_id = warp_id / kSplit; - const uint32_t split_id = warp_id % kSplit; + static_assert(kNumThreads % kWarpThreads == 0, "TODO: support sub-warp copy for small items"); + constexpr uint32_t kNumSplit = kNumThreads / kWarpThreads; + // Integer division below would silently drop the remainder of every row. + static_assert(kKBytes % kNumSplit == 0 && kVBytes % kNumSplit == 0, "the split must divide both rows exactly"); + constexpr uint32_t kKSplitBytes = static_cast(kKBytes) / kNumSplit; + constexpr uint32_t kVSplitBytes = static_cast(kVBytes) / kNumSplit; + + const auto warp_id = blockIdx.x * blockDim.y + threadIdx.y; + const auto item_id = warp_id / kNumSplit; + const auto split_id = warp_id % kNumSplit; + const auto& [ k_input, v_input, k_cache, v_cache, indices, // ptr stride_k, stride_v, stride_k_cache, stride_v_cache, stride_indices, batch_size, // size @@ -200,45 +62,29 @@ __global__ void store_kvcache(const __grid_constant__ StoreKVCacheParams params) ] = params; if (item_id >= batch_size) return; - const auto index_ptr = static_cast(indices) + item_id * stride_indices; PDLWaitPrimary(); + const auto index = static_cast(indices)[item_id * stride_indices]; + const auto k_src = pointer::offset(k_input, item_id * stride_k, split_id * kKSplitBytes); + const auto v_src = pointer::offset(v_input, item_id * stride_v, split_id * kVSplitBytes); - const auto index = *index_ptr; - // A stale/OOB slot id would cause an illegal memory access in the store below; - // fail fast at the culprit instead. always-on (kvcache JIT compiles without NDEBUG). - assert(index >= 0 && index < size_limit); - const auto k_src = pointer::offset(k_input, item_id * stride_k, split_id * kKSplitSize); - const auto v_src = pointer::offset(v_input, item_id * stride_v, split_id * kVSplitSize); - const auto k_dst = pointer::offset(k_cache, index * stride_k_cache, split_id * kKSplitSize); - const auto v_dst = pointer::offset(v_cache, index * stride_v_cache, split_id * kVSplitSize); + using enum warp::LoadStorePattern::type; + const auto k = warp::load_bytes(k_src); + const auto v = warp::load_bytes(v_src); - if (index != reserved_skip_index) { - copy_kv_rows_warp(k_src, v_src, k_dst, v_dst); - } PDLTriggerSecondary(); + assert(index >= 0 && index < size_limit); + if (index != reserved_skip_index) { + const auto k_dst = pointer::offset(k_cache, index * stride_k_cache, split_id * kKSplitBytes); + const auto v_dst = pointer::offset(v_cache, index * stride_v_cache, split_id * kVSplitBytes); + warp::store_bytes(k_dst, k); + warp::store_bytes(v_dst, v); + } } -template +template struct StoreKVCacheKernel { - static_assert(kKElementBytes > 0 && kKElementBytes % 4 == 0); - static_assert(kVElementBytes > 0 && kVElementBytes % 4 == 0); - - template - static constexpr auto store_kernel = store_kvcache; - template - static auto get_kernel(const int num_split) { - using namespace host; - // only apply split optimization when both element sizes are aligned - if constexpr (kKElementBytes % (4 * 128) == 0 && kVElementBytes % (4 * 128) == 0) { - if (num_split == 4) return store_kernel<4, T>; - } - if constexpr (kKElementBytes % (2 * 128) == 0 && kVElementBytes % (2 * 128) == 0) { - if (num_split == 2) return store_kernel<2, T>; - } - if (num_split == 1) return store_kernel<1, T>; - Panic("Unsupported num_split {} for element sizes k={} v={}", num_split, kKElementBytes, kVElementBytes); - } + static constexpr auto store_kernel = store_kvcache_kernel; static void run(const tvm::ffi::TensorView k, @@ -246,53 +92,63 @@ struct StoreKVCacheKernel { const tvm::ffi::TensorView k_cache, const tvm::ffi::TensorView v_cache, const tvm::ffi::TensorView indices, - const int num_split, const int64_t size_limit, const int64_t reserved_skip_index) { using namespace host; auto B = SymbolicSize{"batch_size"}; auto DK = SymbolicSize{"k_element_size"}; auto DV = SymbolicSize{"v_element_size"}; - auto KS = SymbolicSize{"k_stride"}; - auto VS = SymbolicSize{"v_stride"}; - auto SK = SymbolicSize{"k_cache_stride"}; - auto SV = SymbolicSize{"v_cache_stride"}; - auto I = SymbolicSize{"indices_stride"}; auto dtype = SymbolicDType{}; - auto device = SymbolicDevice{}; - auto indice_dtype = SymbolicDType{}; - device.set_options(); + auto device_ = SymbolicDevice{}; + auto idx_dtype = SymbolicDType{}; + device_.set_options(); + + using device::warp::LoadStorePattern; + using enum LoadStorePattern::type; + // Feed get_vec_bytes the SPLIT width, i.e. the exact value the kernel hands + // to load_bytes -- the full row can resolve to a narrower vector and would + // then under-constrain the strides. + constexpr uint32_t kNumSplit = kNumThreads / device::kWarpThreads; + constexpr int64_t kAlignK = LoadStorePattern::get_vec_bytes(); + constexpr int64_t kAlignV = LoadStorePattern::get_vec_bytes(); TensorMatcher({B, DK}) // - .with_strides({KS, 1}) + .with_strides({-1, 1}) .with_dtype(dtype) - .with_device(device) + .with_device(device_) + .ensure_alignment(kAlignK) .verify(k); TensorMatcher({B, DV}) // - .with_strides({VS, 1}) + .with_strides({-1, 1}) .with_dtype(dtype) - .with_device(device) + .with_device(device_) + .ensure_alignment(kAlignV) .verify(v); TensorMatcher({-1, DK}) // - .with_strides({SK, 1}) + .with_strides({-1, 1}) .with_dtype(dtype) - .with_device(device) + .with_device(device_) + .ensure_alignment(kAlignK) .verify(k_cache); TensorMatcher({-1, DV}) // - .with_strides({SV, 1}) + .with_strides({-1, 1}) .with_dtype(dtype) - .with_device(device) + .with_device(device_) + .ensure_alignment(kAlignV) .verify(v_cache); TensorMatcher({B}) // - .with_strides({I}) - .with_dtype(indice_dtype) - .with_device(device) + .with_strides({-1}) + .with_dtype(idx_dtype) + .with_device(device_) .verify(indices); - const int64_t dtype_size = dtype_bytes(dtype.unwrap()); - const uint32_t num_elements = static_cast(B.unwrap()); - RuntimeCheck(kKElementBytes == dtype_size * DK.unwrap()); - RuntimeCheck(kVElementBytes == dtype_size * DV.unwrap()); + const auto dtype_size = static_cast(dtype_bytes(dtype.unwrap())); + const auto batch_size = static_cast(B.unwrap()); + const auto device = device_.unwrap(); + CHECK_HOST(kKBytes == dtype_size * DK.unwrap()); + CHECK_HOST(kVBytes == dtype_size * DV.unwrap()); + + if (batch_size == 0) return; const auto params = StoreKVCacheParams{ .k = k.data_ptr(), @@ -300,20 +156,28 @@ struct StoreKVCacheKernel { .k_cache = k_cache.data_ptr(), .v_cache = v_cache.data_ptr(), .indices = indices.data_ptr(), - .stride_k_bytes = KS.unwrap() * dtype_size, - .stride_v_bytes = VS.unwrap() * dtype_size, - .stride_k_cache_bytes = SK.unwrap() * dtype_size, - .stride_v_cache_bytes = SV.unwrap() * dtype_size, - .stride_indices = I.unwrap(), - .batch_size = static_cast(B.unwrap()), + .stride_k_bytes = k.stride(0) * dtype_size, + .stride_v_bytes = v.stride(0) * dtype_size, + .stride_k_cache_bytes = k_cache.stride(0) * dtype_size, + .stride_v_cache_bytes = v_cache.stride(0) * dtype_size, + .stride_indices = indices.stride(0), + .batch_size = batch_size, .size_limit = size_limit, .reserved_skip_index = reserved_skip_index, }; - // select kernel and update num_split if needed - const auto use_int32 = indice_dtype.is_type(); - const auto kernel = use_int32 ? get_kernel(num_split) : get_kernel(num_split); - const auto num_blocks = div_ceil(num_elements * num_split, kNumWarps); - LaunchKernel(num_blocks, kThreadsPerBlock, device.unwrap()) // + + const auto kernel = idx_dtype.is_type() ? store_kernel : store_kernel; + const auto total_warps = batch_size * kNumSplit; + const auto num_warps = [&] { + const auto sm_count = runtime::get_sm_count(device.device_id); +#pragma unroll + for (uint32_t n : {1, 2, 4}) { + if (total_warps <= sm_count * n) return n; + } + return 8u; + }(); + const auto num_blocks = div_ceil(total_warps, num_warps); + LaunchKernel(num_blocks, {device::kWarpThreads, num_warps}, device) // .enable_pdl(kUsePDL)(kernel, params); } }; diff --git a/python/sglang/kernels/jit/csrc/elementwise/set_mla_kv_buffer.cuh b/python/sglang/kernels/jit/csrc/elementwise/set_mla_kv_buffer.cuh index a80d63562..2fbd99470 100644 --- a/python/sglang/kernels/jit/csrc/elementwise/set_mla_kv_buffer.cuh +++ b/python/sglang/kernels/jit/csrc/elementwise/set_mla_kv_buffer.cuh @@ -1,38 +1,17 @@ -// 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 #include +#include #include #include -#include #include -#include #include #include +#include #include namespace sglang { @@ -49,170 +28,115 @@ struct SetMlaKVBufferParams { int64_t reserved_skip_index; }; -template +template __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; + using enum warp::LoadStorePattern::type; + const auto global_warp_id = threadIdx.y + blockIdx.x * blockDim.y; + const auto input_nope = pointer::offset(params.k_nope, params.stride_nope_bytes * global_warp_id); + const auto input_rope = pointer::offset(params.k_rope, params.stride_rope_bytes * global_warp_id); + if (global_warp_id >= params.batch_size) return; PDLWaitPrimary(); - - const int64_t loc = static_cast(static_cast(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::copy_bytes(nope_src, &smem[warp_in_cta][0]); - warp::copy_bytes(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 && loc != params.reserved_skip_index) { - cuda::ptx::cp_async_bulk( - cuda::ptx::space_global, - cuda::ptx::space_shared, - gmem_dst, - &smem[warp_in_cta][0], - static_cast(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>{}); + const int64_t loc = static_cast(static_cast(params.loc)[global_warp_id]); + const auto nope = warp::load_bytes(input_nope); + const auto rope = warp::load_bytes(input_rope); PDLTriggerSecondary(); + if (loc != params.reserved_skip_index) { + const auto output_nope = pointer::offset(params.kv_buffer, params.stride_buffer_bytes * loc); + const auto output_rope = pointer::offset(output_nope, kNopeBytes); + warp::store_bytes(output_nope, nope); + warp::store_bytes(output_rope, rope); + } } template 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 - static constexpr auto kernel = set_mla_kv_buffer_kernel; + template + static constexpr auto set_kernel = set_mla_kv_buffer_kernel; 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, + int64_t, int64_t reserved_skip_index) { 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(); + auto device_ = SymbolicDevice{}; + device_.set_options(); - TensorMatcher({B, D_nope}) // - .with_strides({S_nope, 1}) + using device::warp::LoadStorePattern; + using enum LoadStorePattern::type; + constexpr int64_t kAlignNope = LoadStorePattern::get_vec_bytes(); + constexpr int64_t kAlignRope = LoadStorePattern::get_vec_bytes(); + // The buffer row carries both halves, so it has to satisfy the WIDER of the + // two -- the narrower one alone would let a nope-misaligned stride through. + constexpr int64_t kAlignBuffer = std::max(kAlignNope, kAlignRope); + // The rope half starts kNopeBytes into the row, so that offset must not + // break the rope alignment the buffer was just checked for. + static_assert(kNopeBytes % kAlignRope == 0, "nope width must not misalign the rope half"); + + TensorMatcher({B, -1}) // + .with_strides({-1, 1}) .with_dtype(dtype) - .with_device(device) + .with_device(device_) + .ensure_alignment(kAlignNope) .verify(k_nope); - TensorMatcher({B, D_rope}) // - .with_strides({S_rope, 1}) + TensorMatcher({B, -1}) // + .with_strides({-1, 1}) .with_dtype(dtype) - .with_device(device) + .with_device(device_) + .ensure_alignment(kAlignRope) .verify(k_rope); - TensorMatcher({-1, D_buf}) // - .with_strides({S_buf, 1}) + TensorMatcher({-1, -1}) // + .with_strides({-1, 1}) .with_dtype(dtype) - .with_device(device) + .with_device(device_) + .ensure_alignment(kAlignBuffer) .verify(kv_buffer); TensorMatcher({B}) // - .with_strides({S_loc}) + .with_strides({-1}) .with_dtype(loc_dtype) - .with_device(device) + .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(B.unwrap()); - if (batch == 0) return; + const auto dtype_size = static_cast(dtype_bytes(dtype.unwrap())); + CHECK_HOST(kv_buffer.size(1) >= k_nope.size(1) + k_rope.size(1)); + CHECK_HOST(k_nope.size(1) * dtype_size == kNopeBytes); + CHECK_HOST(k_rope.size(1) * dtype_size == kRopeBytes); + const auto batch_size = static_cast(B.unwrap()); + if (batch_size == 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, + .stride_nope_bytes = k_nope.stride(0) * dtype_size, + .stride_rope_bytes = k_rope.stride(0) * dtype_size, + .stride_buffer_bytes = kv_buffer.stride(0) * dtype_size, + .batch_size = batch_size, .reserved_skip_index = reserved_skip_index, }; - - const auto use_int32 = loc_dtype.is_type(); - - auto launch = [&]() { - const auto kernel_ptr = use_int32 ? kernel : kernel; - const uint32_t num_blocks = div_ceil(batch, static_cast(kNW)); - const uint32_t threads_per_block = static_cast(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); - } + const auto device = device_.unwrap(); + const auto kernel = loc_dtype.is_type() ? set_kernel : set_kernel; + const auto num_warps = [&] { + const auto sm_count = runtime::get_sm_count(device.device_id); +#pragma unroll + for (uint32_t n : {1, 2, 4}) { + if (batch_size <= sm_count * n) return n; + } + return 8u; + }(); + const auto num_blocks = div_ceil(batch_size, num_warps); + LaunchKernel(num_blocks, {device::kWarpSize, num_warps}, device) // + .enable_pdl(kUsePDL)(kernel, params); } }; diff --git a/python/sglang/kernels/jit/csrc/elementwise/set_mla_kv_concat_q.cuh b/python/sglang/kernels/jit/csrc/elementwise/set_mla_kv_concat_q.cuh index 44af173fa..4bbf24596 100644 --- a/python/sglang/kernels/jit/csrc/elementwise/set_mla_kv_concat_q.cuh +++ b/python/sglang/kernels/jit/csrc/elementwise/set_mla_kv_concat_q.cuh @@ -8,7 +8,7 @@ #include #include // For LaunchKernel, SGL_DEVICE, PDL helpers #include // For AlignedVector -#include // For warp::copy_bytes, elect_one_lane, inclusive_sum +#include // For warp::load_bytes, store_bytes, elect_one_lane #include #include @@ -52,9 +52,6 @@ __global__ void set_mla_kv_concat_q_kernel(const __grid_constant__ SetMlaKVConca constexpr int kQNopeDim = static_cast(kNopeBytes / sizeof(bf16_t)); constexpr int kQRopeDim = static_cast(kRopeBytes / sizeof(bf16_t)); - // Per-warp smem slots for the KV scatter role; concat warps leave theirs idle. - __shared__ alignas(16) uint8_t smem[kNumWarps][kRowBytes]; - const uint32_t warp_in_cta = threadIdx.x / kWarpThreads; const uint32_t lane_id = threadIdx.x % kWarpThreads; const uint32_t flat_warp = blockIdx.x * kNumWarps + warp_in_cta; @@ -64,33 +61,19 @@ __global__ void set_mla_kv_concat_q_kernel(const __grid_constant__ SetMlaKVConca if (flat_warp < params.batch_size) { // --- KV scatter role: one warp per token (smem staging + TMA bulk store) --- const uint32_t item_id = flat_warp; - const int64_t loc = static_cast(static_cast(params.loc)[item_id]); + const auto loc = static_cast(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::copy_bytes(nope_src, &smem[warp_in_cta][0]); - warp::copy_bytes(rope_src, &smem[warp_in_cta][kNopeBytes]); + using enum warp::LoadStorePattern::type; + const auto nope = warp::load_bytes(nope_src); + const auto rope = warp::load_bytes(rope_src); - // TMA reads smem via the async proxy; fence so it can't observe stale sts. - __syncwarp(); - asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); - - // elect.sync rather than `lane_id == 0`: the TMA issue must not sit - // behind a lane-index predicate (see PR review). - if (device::warp::elect_one_lane()) { - cuda::ptx::cp_async_bulk( - cuda::ptx::space_global, - cuda::ptx::space_shared, - gmem_dst, - &smem[warp_in_cta][0], - static_cast(kRowBytes)); - } - - // ``wait_group`` (not ``_read``): waits for gmem commit, not just smem reuse. - cuda::ptx::cp_async_bulk_commit_group(); - cuda::ptx::cp_async_bulk_wait_group(cuda::ptx::n32_t<0>{}); + const auto nope_dst = pointer::offset(params.kv_buffer, loc * params.stride_buffer_bytes); + const auto rope_dst = pointer::offset(nope_dst, kNopeBytes); + warp::store_bytes(nope_dst, nope); + warp::store_bytes(rope_dst, rope); } else if (flat_warp - params.batch_size < params.num_q_items) { // --- Q concat role: one warp per (token, head) row --- const uint32_t q_item = flat_warp - params.batch_size; @@ -223,7 +206,7 @@ struct SetMlaKVConcatQKernel { // Alignment tripwires. The device code does 16-byte vector accesses on the // kv row / nope rows / q rows and 4-byte accesses on the rope rows; the // python-side ``covered()`` mirrors these so uncovered layouts fall back - // instead of faulting (do NOT assume "PyTorch tensors are aligned" — views + // instead of faulting (do NOT assume "PyTorch tensors are aligned" -- views // and odd pool pitches break that). const auto aligned = [](const void* ptr, int64_t align) { return reinterpret_cast(ptr) % static_cast(align) == 0; diff --git a/python/sglang/kernels/jit/csrc/kimi_k3/situ_and_mul.cuh b/python/sglang/kernels/jit/csrc/kimi_k3/situ_and_mul.cuh index 5280de48c..262a27d1b 100644 --- a/python/sglang/kernels/jit/csrc/kimi_k3/situ_and_mul.cuh +++ b/python/sglang/kernels/jit/csrc/kimi_k3/situ_and_mul.cuh @@ -11,7 +11,7 @@ #include // For dtype_trait, bf16_t, fp32_t, cast #include // For LaunchKernel, SGL_DEVICE, PDL helpers #include // For AlignedVector -#include // For warp::copy_bytes, elect_one_lane, inclusive_sum +#include // For warp::inclusive_sum, reduce_sum, reduce_max #include @@ -247,7 +247,7 @@ SGL_DEVICE CTAWork get_work(const SituMulQuantVarlenParams& params) { const uint32_t val = tx < params.num_experts ? params.masked_m[tx] : 0u; // Per-warp inclusive scan of masked_m. - const uint32_t warp_inclusive = device::warp::inclusive_sum(lane_id, val); + const uint32_t warp_inclusive = device::warp::inclusive_sum(val, lane_id); const uint32_t warp_exclusive = warp_inclusive - val; // Write each warp total. diff --git a/python/sglang/kernels/jit/csrc/kvcacheio/hicache.cuh b/python/sglang/kernels/jit/csrc/kvcacheio/hicache.cuh index 3a9a5afa5..2b5228e9b 100644 --- a/python/sglang/kernels/jit/csrc/kvcacheio/hicache.cuh +++ b/python/sglang/kernels/jit/csrc/kvcacheio/hicache.cuh @@ -208,11 +208,16 @@ SGL_HICACHE_KERNEL void hicache_transfer_per_layer(const __grid_constant__ Hicac const auto src_k = pointer::offset(k_cache_src, pos_src * kv_cache_src_stride); const auto dst_k = pointer::offset(k_cache_dst, pos_dst * kv_cache_dst_stride); const auto vec_k = load_vec(src_k); - store_vec(dst_k, vec_k); + // Both loads are issued before either store: the compiler cannot prove + // dst_k and src_v disjoint, so it will not hoist the V load on its own. + std::decay_t vec_v; if constexpr (!kIsMLA) { const auto src_v = pointer::offset(v_cache_src, pos_src * kv_cache_src_stride); + vec_v = load_vec(src_v); + } + store_vec(dst_k, vec_k); + if constexpr (!kIsMLA) { const auto dst_v = pointer::offset(v_cache_dst, pos_dst * kv_cache_dst_stride); - const auto vec_v = load_vec(src_v); store_vec(dst_v, vec_v); } } @@ -253,13 +258,18 @@ SGL_HICACHE_KERNEL void hicache_transfer_all_layer(const __grid_constant__ Hicac const auto src_k = pointer::offset(k_cache_src, pos_src * kv_cache_src_stride); const auto dst_k = pointer::offset(k_cache_dst, pos_dst * kv_cache_dst_stride); const auto vec_k = load_vec(src_k); - store_vec(dst_k, vec_k); + // Both loads are issued before either store: the compiler cannot prove + // dst_k and src_v disjoint, so it will not hoist the V load on its own. + std::decay_t vec_v; if constexpr (!kIsMLA) { const auto v_cache_src = static_cast(v_ptr_src)[layer]; - const auto v_cache_dst = static_cast(v_ptr_dst)[layer]; const auto src_v = pointer::offset(v_cache_src, pos_src * kv_cache_src_stride); + vec_v = load_vec(src_v); + } + store_vec(dst_k, vec_k); + if constexpr (!kIsMLA) { + const auto v_cache_dst = static_cast(v_ptr_dst)[layer]; const auto dst_v = pointer::offset(v_cache_dst, pos_dst * kv_cache_dst_stride); - const auto vec_v = load_vec(src_v); store_vec(dst_v, vec_v); } } diff --git a/python/sglang/kernels/jit/csrc/minimax/minimax_decode_topk.cuh b/python/sglang/kernels/jit/csrc/minimax/minimax_decode_topk.cuh index ee57ed248..a0a429f55 100644 --- a/python/sglang/kernels/jit/csrc/minimax/minimax_decode_topk.cuh +++ b/python/sglang/kernels/jit/csrc/minimax/minimax_decode_topk.cuh @@ -67,20 +67,6 @@ struct TopKTrait { constexpr auto is_greater = [](float x, float y, int32_t delta) { return (x > y) || ((x == y) && delta < 0); // lower block id wins }; - constexpr auto warp_inclusive_sum = [](uint32_t lane_id, uint32_t val) { -#pragma unroll - for (uint32_t offset = 1; offset < device::kWarpThreads; offset *= 2) { - // Width-32 up-shuffle. On wave64 HIP the un-suffixed __shfl_up takes the - // logical-warp width directly; CUDA needs the active mask. -#ifdef USE_ROCM - uint32_t n = __shfl_up(val, offset, device::kWarpThreads); -#else - uint32_t n = __shfl_up_sync(kWarpSyncMask, val, offset, device::kWarpThreads); -#endif - if (lane_id >= offset) val += n; - } - return val; - }; constexpr auto clip_nan = [](float x) { return x != x ? kNegInf : x; }; constexpr auto score_to_key = [](float x) { uint32_t b = __float_as_uint(x); @@ -95,7 +81,7 @@ struct TopKTrait { uint32_t warp_inc = 0; if (tx < kRadixSize) { hist_val = histogram[tx]; - warp_inc = warp_inclusive_sum(lane_id, hist_val); + warp_inc = warp::inclusive_sum(hist_val, lane_id); if (lane_id == kWarpThreads - 1) smem->warp_sum[warp_id] = warp_inc; } __syncthreads(); diff --git a/python/sglang/kernels/jit/csrc/moe/route_radix.cuh b/python/sglang/kernels/jit/csrc/moe/route_radix.cuh index 492d1b980..ac9b45711 100644 --- a/python/sglang/kernels/jit/csrc/moe/route_radix.cuh +++ b/python/sglang/kernels/jit/csrc/moe/route_radix.cuh @@ -10,7 +10,7 @@ #include // For dtype_trait, bf16_t, fp32_t, cast #include // For LaunchKernel, SGL_DEVICE, PDL helpers #include // For AlignedVector -#include // For warp::copy_bytes, elect_one_lane, inclusive_sum +#include // For warp::inclusive_sum, reduce_sum #include @@ -60,7 +60,7 @@ SGL_DEVICE void bar_sync(uint32_t id, uint32_t num_threads) { // smem_warp_sum[kNumWarps]; syncs on entry (so the workspace can be reused // across calls) and before the cross-warp read. SGL_DEVICE uint32_t block_exclusive_sum(uint32_t cnt, uint32_t lane_id, uint32_t warp_id, uint32_t* smem_warp_sum) { - const uint32_t inc = device::warp::inclusive_sum(lane_id, cnt); + const uint32_t inc = device::warp::inclusive_sum(cnt, lane_id); if (lane_id == 31) smem_warp_sum[warp_id] = inc; __syncthreads(); // TODO: replace `__reduce_add_sync` with `warp::reduce_sum` @@ -134,12 +134,12 @@ SGL_DEVICE void route_radix_block(const RouteRadixParams& params, typename Large // ---- Load + key transform: thread tx owns experts [4*tx, 4*tx+4) ---- uint32_t keys[kVecSize]; - float act[kVecSize]; // raw sigmoid (weight source) — never NaN-sanitized + float act[kVecSize]; // raw sigmoid (weight source) -- never NaN-sanitized { const auto scores = static_cast(params.scores) + bx * params.scores_stride; AlignedVector bias_vec; // bf16: 2x bf16x2 (8B row loads); fp32: 2x fp32x2 (16B row loads). The - // radix math below is fp32 either way — only the load width differs. + // radix math below is fp32 either way -- only the load width differs. AlignedVector, kVecSize / 2> scores_vec; // Bias may be produced by a preceding cast or fill kernel (the caller @@ -208,7 +208,7 @@ SGL_DEVICE void route_radix_block(const RouteRadixParams& params, typename Large AlignedVector hist; hist.load(smem.histogram, tx); const auto local_val = hist[0] + hist[1]; - const auto warp_inc = device::warp::inclusive_sum(lane_id, local_val); + const auto warp_inc = device::warp::inclusive_sum(local_val, lane_id); if (lane_id == kWarpThreads - 1) smem.warp_sum[0][warp_id] = warp_inc; moe::radix::bar_sync(BAR_SUM, kRadixLanes); const auto inter = __reduce_add_sync(0xFFFFFFFF, lane_id < warp_id ? smem.warp_sum[0][lane_id] : 0u); @@ -316,7 +316,7 @@ SGL_DEVICE void route_radix_block(const RouteRadixParams& params, typename Large params.out_w[bx * params.out_w_stride + rank] = w; params.out_i[bx * params.out_i_stride + rank] = id; if (params.out_packed != nullptr) { - // (id << 16) | bf16(w) bits — RN float->bf16 matches the triton pack. + // (id << 16) | bf16(w) bits -- RN float->bf16 matches the triton pack. const auto bits = static_cast(__bfloat16_as_ushort(__float2bfloat16_rn(w))); params.out_packed[bx * params.out_packed_stride + rank] = static_cast((static_cast(id) << 16) | bits); @@ -458,7 +458,7 @@ SGL_DEVICE void fgt_select_topk( device::AlignedVector hist; hist.load(smem.histogram, tx); const auto local_val = hist[0] + hist[1]; - const auto warp_inc = device::warp::inclusive_sum(lane_id, local_val); + const auto warp_inc = device::warp::inclusive_sum(local_val, lane_id); if (lane_id == 31) smem.warp_sum[0][warp_id] = warp_inc; moe::radix::bar_sync(BAR_SUM, kRadixLanes); const auto inter = __reduce_add_sync(0xFFFFFFFF, lane_id < warp_id ? smem.warp_sum[0][lane_id] : 0u); diff --git a/python/sglang/kernels/jit/include/sgl_kernel/allocator.h b/python/sglang/kernels/jit/include/sgl_kernel/allocator.h new file mode 100644 index 000000000..9b034f59a --- /dev/null +++ b/python/sglang/kernels/jit/include/sgl_kernel/allocator.h @@ -0,0 +1,40 @@ +#pragma once + +#include +#include +#include +#include + +namespace sglang::host { + +/** + * \brief Allocate only once for a given function. + * \tparam kThreadSafe Whether to make the allocation thread-safe. + * \tparam Salt A salt type to avoid cache collision. + * \param key The key to identify the allocation. It should be unique for each allocation. + * \param callback The callback function to perform the allocation. It should return the allocated value. + * \note The `Fn` type must be unique. It's typically a lambda type that's evaluated only once. + * Otherwise, different call-sites may hit the same cache entry. + * In case where `Fn` is not unique (e.g. std::function), make `Salt` unique to avoid cache collision. + */ +template +inline auto allocate_once(Key&& key, Fn&& callback) -> std::decay_t>& { + using Value = std::decay_t>; + static std::unordered_map, Value> s_map; + const auto alloc = [&]() -> Value& { + const auto iter = s_map.find(key); + if (iter != s_map.end()) return iter->second; + // Evaluate the callback before inserting, so a throwing callback leaves no empty entry behind. + auto value = std::forward(callback)(); + return s_map.emplace(std::forward(key), std::move(value)).first->second; + }; + if constexpr (kThreadSafe) { + static std::mutex s_mutex; + const auto lock = std::lock_guard{s_mutex}; + return alloc(); + } else { + return alloc(); + } +} + +} // namespace sglang::host diff --git a/python/sglang/kernels/jit/include/sgl_kernel/atomic.cuh b/python/sglang/kernels/jit/include/sgl_kernel/atomic.cuh index 56f9b4450..de35bef35 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/atomic.cuh +++ b/python/sglang/kernels/jit/include/sgl_kernel/atomic.cuh @@ -34,6 +34,136 @@ SGL_DEVICE float max(float* addr, float value) { #endif } +namespace ptx { + +SGL_DEVICE void red_release_add_u32(uint32_t* ptr, uint32_t n) { + asm volatile("red.release.gpu.global.add.u32 [%0], %1;" ::"l"(ptr), "r"(n) : "memory"); +} + +SGL_DEVICE void red_relaxed_add_u32(uint32_t* ptr, uint32_t n) { + asm volatile("red.relaxed.gpu.global.add.u32 [%0], %1;" ::"l"(ptr), "r"(n) : "memory"); +} + +SGL_DEVICE uint32_t atom_acquire_cas_b32(uint32_t* addr, uint32_t compare, uint32_t swap) { + uint32_t result; + asm volatile("atom.acquire.gpu.global.cas.b32 %0, [%1], %2, %3;" + : "=r"(result) + : "l"(addr), "r"(compare), "r"(swap) + : "memory"); + return result; +} + +SGL_DEVICE uint32_t load_acquire_u32(uint32_t* addr) { + uint32_t result; + asm volatile("ld.acquire.gpu.global.u32 %0, [%1];" : "=r"(result) : "l"(addr) : "memory"); + return result; +} + +SGL_DEVICE uint32_t atom_acquire_add_u32(uint32_t* addr, uint32_t n) { + uint32_t result; + asm volatile("atom.acquire.gpu.global.add.u32 %0, [%1], %2;" : "=r"(result) : "l"(addr), "r"(n) : "memory"); + return result; +} + +} // namespace ptx + +/** + * \brief Cross-CTA arrive/wait counter packed into one 32-bit word. + * + * Producers call `arrive()`; consumers call `wait()` (exactly one consumer) or + * `wait_multi()` (several) until every producer has. The word is split: the low + * `32 - kConsumerBits` bits count producer arrivals, the high bits count + * consumers that have already been released. The last consumer to be released + * subtracts the whole thing, so one Event is reusable across launches without a + * host-side re-zero. + * + * \note The handle must be ZERO before first use. Nothing constructs it on the + * device, so zero the backing allocation from the host once. + * \note The storage must be GLOBAL memory: the PTX below names the `.global` + * state space, so a shared or local Event is an illegal address. + * \note `arrive()` is a release and the waits are acquires, so a producer's + * writes before `arrive()` are visible to a consumer after the wait. + * \note Generations must be explicitly ordered: every consumer of generation N + * has to be released before any producer arrives for generation N + 1. + * The word carries no phase bit, so overlapping two generations on one + * Event is undefined behavior. + * \note `wait()` and `wait_multi()` lay the word out incompatibly. Mixing them + * on one Event is undefined behavior. + */ +struct Event { + public: + using handle_type = uint32_t; + + Event(const Event&) = delete; + Event& operator=(const Event&) = delete; + + /// \brief DON'T touch unless you know what you're doing. + SGL_DEVICE handle_type& unsafe_get_handle() { + return m_handle; + } + + /** + * \brief Increment the producer count by `n`. + * \param n The number of producers to arrive. Defaults to 1. + * \note This is a release operation, so any writes before `arrive()` are + * visible to a consumer after `wait()`. + */ + SGL_DEVICE void arrive(uint32_t n = 1) { + ptx::red_release_add_u32(&m_handle, n); + } + + /** + * \brief Block until `num_producers` producers have arrived. + * \param num_producers The number of producers to wait for. + * + * Single-consumer: simpler and faster than `wait_multi()`, but exactly one + * thread in the whole grid may call it per generation. + */ + SGL_DEVICE void wait(uint32_t num_producers) { + while (ptx::atom_acquire_cas_b32(&m_handle, num_producers, 0) != num_producers) + ; + } + + /** + * \brief Block until `num_producers` producers have arrived, with several + * consumers sharing the Event. + * \tparam kConsumerBits Bits reserved for the consumer half of the word. + * \param num_producers Must be `< 1 << (32 - kConsumerBits)`. + * \param num_consumers Must be in `[1, 1 << kConsumerBits)`, and the `n` of + * all callers has to sum to exactly this, otherwise the + * Event is never reset. + * \param n How many of `num_consumers` this call stands for. + * Defaults to 1, i.e. one calling thread per consumer. + */ + template + SGL_DEVICE void wait_multi(uint32_t num_producers, uint32_t num_consumers, uint32_t n = 1) { + static_assert(kConsumerBits > 0 && kConsumerBits < 32); + constexpr uint32_t kProducerBits = 32 - kConsumerBits; + constexpr uint32_t kProducerMask = (1u << kProducerBits) - 1; + + __builtin_assume(num_producers < (1u << kProducerBits)); + __builtin_assume(num_consumers > 0 && num_consumers < (1u << kConsumerBits)); + + // Register and observe in the SAME atomic. ticket = consumers ahead of me. + const auto ticket = ptx::atom_acquire_add_u32(&m_handle, n << kProducerBits); + if ((ticket & kProducerMask) != num_producers) { + /// NOTE: when v = 0, a reset has already happened. + while (const auto v = ptx::load_acquire_u32(&m_handle)) { + if ((v & kProducerMask) == num_producers) break; + } + } + + // The last consumer to register should reset the counter to 0 + if ((ticket >> kProducerBits) + n == num_consumers) { + const auto final_value = num_producers | (num_consumers << kProducerBits); + ptx::red_relaxed_add_u32(&m_handle, -final_value); + } + } + + private: + handle_type m_handle; +}; + } // namespace device::atomic } // namespace sglang diff --git a/python/sglang/kernels/jit/include/sgl_kernel/bits.h b/python/sglang/kernels/jit/include/sgl_kernel/bits.h new file mode 100644 index 000000000..c2d7713aa --- /dev/null +++ b/python/sglang/kernels/jit/include/sgl_kernel/bits.h @@ -0,0 +1,41 @@ +#pragma once +#include +#include +#include + +namespace sglang { + +namespace host { + +template +inline constexpr bool is_pow2(T x) { + return std::has_single_bit(x); +} + +/// \brief `floor(log2(x))`; -1 for `x == 0`. +template +inline constexpr int32_t log2_floor(T x) { + if (x == 0) return -1; + return std::bit_width(x) - 1; +} + +/// \brief `ceil(log2(x))`; -1 for `x == 0`. +template +inline constexpr int32_t log2_ceil(T x) { + if (x == 0) return -1; + return std::bit_width(x - 1); +} + +template +inline constexpr T round_up_pow2(T x) { + return std::bit_ceil(x); +} + +template +inline constexpr T round_down_pow2(T x) { + return std::bit_floor(x); +} + +} // namespace host + +} // namespace sglang diff --git a/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/topk_impl.cuh b/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/topk_impl.cuh index 8836389c8..b0b208854 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/topk_impl.cuh +++ b/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/topk_impl.cuh @@ -58,9 +58,6 @@ SGL_DEVICE T broadcast(T value, uint32_t src = 0) { #endif } -/// sgl_kernel names the warp size `kWarpThreads`; alias it locally as `kWarpSize`. -inline constexpr uint32_t kWarpSize = kWarpThreads; - template struct MaxSmem { static constexpr size_t kSize = std::max({sizeof(Smems)...}); @@ -148,19 +145,6 @@ SGL_DEVICE float coarse_bin_lower_bound(uint32_t bin) { return extract_coarse_bin(mid) < bin ? step_up(mid) : mid; } -SGL_DEVICE uint32_t warp_inclusive_sum(uint32_t lane_id, uint32_t val) { -#pragma unroll - for (uint32_t offset = 1; offset < 32; offset *= 2) { -#ifndef USE_ROCM - uint32_t n = __shfl_up_sync(0xFFFFFFFF, val, offset); -#else - uint32_t n = __shfl_up_sync(kFullMask, val, offset, kWarpThreads); -#endif - if (lane_id >= offset) val += n; - } - return val; -} - SGL_DEVICE uint32_t warp_sum_bool(bool pred, uint32_t mask = 0xFFFFFFFF) { #ifdef USE_ROCM // The ballot covers the whole hardware wave, which on wave64 holds two of @@ -377,7 +361,7 @@ struct TopKConfig { uint32_t warp_inc = 0; if (tx < kRadixSize) { hist_val = histogram[tx]; - warp_inc = warp_inclusive_sum(lane_id, hist_val); + warp_inc = warp::inclusive_sum(hist_val, lane_id); if (lane_id == kWarpSize - 1) smem->warp_sum[warp_id] = warp_inc; } __syncthreads(); @@ -516,7 +500,7 @@ struct TopKRadixBase : TopKConfig { const auto local_sum = local_exc_sum[kItems]; const auto lane_id = tx % kWarpSize; const auto warp_id = broadcast(tx / kWarpSize); - const auto warp_inc_sum = warp_inclusive_sum(lane_id, local_sum); + const auto warp_inc_sum = warp::inclusive_sum(local_sum, lane_id); const auto warp_exc_sum = warp_inc_sum - local_sum; if (lane_id == kWarpSize - 1) smem->warp_sum[warp_id] = warp_inc_sum; diff --git a/python/sglang/kernels/jit/include/sgl_kernel/runtime.cuh b/python/sglang/kernels/jit/include/sgl_kernel/runtime.cuh index 9f1b4f1d8..f8affdebe 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/runtime.cuh +++ b/python/sglang/kernels/jit/include/sgl_kernel/runtime.cuh @@ -2,7 +2,7 @@ /// \brief Host-side CUDA runtime query helpers. /// /// Thin wrappers around CUDA occupancy and device-property APIs with -/// automatic error checking via `RuntimeDeviceCheck`. +/// automatic error checking via `CHECK_CUDA`. #pragma once @@ -10,6 +10,7 @@ #include #include +#include #ifndef USE_ROCM #include #else @@ -66,45 +67,87 @@ inline void* get_device_accessible_ptr(const tvm::ffi::TensorView& tensor) { return device_ptr; } +namespace details { + +template +struct DeviceCacheMap { + public: + // Generous bound on the device ordinals one process can see; a larger ordinal + // is not an error, it just falls through to the driver query uncached. + static constexpr uint32_t kNumStaticMaxDevice = 72; + + constexpr DeviceCacheMap() { + for (uint32_t i = 0; i < kNumStaticMaxDevice; ++i) { + m_data[i] = kDefault; + } + } + + template + T get_cached(int32_t device_, bool use_cache, Fn&& fn) { + const auto device = static_cast(device_); + if (use_cache && device < kNumStaticMaxDevice && m_data[device] != kDefault) { + return m_data[device]; + } + const auto value = static_cast(std::forward(fn)(device_)); + if (device < kNumStaticMaxDevice) { + m_data[device] = value; + } + return value; + } + + private: + T m_data[kNumStaticMaxDevice]; +}; + +} // namespace details + // Return the maximum number of active blocks per SM for the given kernel template inline auto get_blocks_per_sm(T&& kernel, int32_t block_dim, std::size_t dynamic_smem = 0) -> uint32_t { int num_blocks_per_sm = 0; - RuntimeDeviceCheck( - cudaOccupancyMaxActiveBlocksPerMultiprocessor(&num_blocks_per_sm, kernel, block_dim, dynamic_smem)); + CHECK_CUDA(cudaOccupancyMaxActiveBlocksPerMultiprocessor(&num_blocks_per_sm, kernel, block_dim, dynamic_smem)); return static_cast(num_blocks_per_sm); } // Return the number of SMs for the given device -inline auto get_sm_count(int device_id) -> uint32_t { - int sm_count; - RuntimeDeviceCheck(cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, device_id)); - return static_cast(sm_count); +inline auto get_sm_count(int device_id, bool use_cache = true) -> uint32_t { + static details::DeviceCacheMap sm_count_cache; + return sm_count_cache.get_cached(device_id, use_cache, [](int32_t device_id) { + int sm_count; + CHECK_CUDA(cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, device_id)); + return sm_count; + }); } // Return the Major compute capability for the given device -inline auto get_cc_major(int device_id) -> int { - int cc_major; - RuntimeDeviceCheck(cudaDeviceGetAttribute(&cc_major, cudaDevAttrComputeCapabilityMajor, device_id)); - return cc_major; +inline auto get_cc_major(int device_id, bool use_cache = true) -> int { + static details::DeviceCacheMap cc_major_cache; + return cc_major_cache.get_cached(device_id, use_cache, [](int32_t device_id) { + int cc_major; + CHECK_CUDA(cudaDeviceGetAttribute(&cc_major, cudaDevAttrComputeCapabilityMajor, device_id)); + return cc_major; + }); } // Return the Minor compute capability for the given device -inline auto get_cc_minor(int device_id) -> int { - int cc_minor; - RuntimeDeviceCheck(cudaDeviceGetAttribute(&cc_minor, cudaDevAttrComputeCapabilityMinor, device_id)); - return cc_minor; +inline auto get_cc_minor(int device_id, bool use_cache = true) -> int { + static details::DeviceCacheMap cc_minor_cache; + return cc_minor_cache.get_cached(device_id, use_cache, [](int32_t device_id) { + int cc_minor; + CHECK_CUDA(cudaDeviceGetAttribute(&cc_minor, cudaDevAttrComputeCapabilityMinor, device_id)); + return cc_minor; + }); } // Return the SM version (major * 10 + minor) for the given device -inline auto get_sm_version(int device_id) -> int { - return get_cc_major(device_id) * 10 + get_cc_minor(device_id); +inline auto get_sm_version(int device_id, bool use_cache = true) -> int { + return get_cc_major(device_id, use_cache) * 10 + get_cc_minor(device_id, use_cache); } // Return the runtime version inline auto get_runtime_version() -> int { int runtime_version; - RuntimeDeviceCheck(cudaRuntimeGetVersion(&runtime_version)); + CHECK_CUDA(cudaRuntimeGetVersion(&runtime_version)); return runtime_version; } @@ -112,7 +155,7 @@ inline auto get_runtime_version() -> int { template inline auto get_available_dynamic_smem_per_block(T&& kernel, int num_blocks, int block_size) -> std::size_t { std::size_t smem_size; - RuntimeDeviceCheck(cudaOccupancyAvailableDynamicSMemPerBlock(&smem_size, kernel, num_blocks, block_size)); + CHECK_CUDA(cudaOccupancyAvailableDynamicSMemPerBlock(&smem_size, kernel, num_blocks, block_size)); return smem_size; } diff --git a/python/sglang/kernels/jit/include/sgl_kernel/tensor.h b/python/sglang/kernels/jit/include/sgl_kernel/tensor.h index 52a0ead7d..196a00a55 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/tensor.h +++ b/python/sglang/kernels/jit/include/sgl_kernel/tensor.h @@ -10,6 +10,7 @@ /// usage examples. #pragma once +#include #include #include @@ -18,6 +19,7 @@ #include #include +#include #include #include #include @@ -527,6 +529,14 @@ struct TensorMatcher { return std::move(*this); } + /// Ensure alignment on all dimensions except for the last dimension. + auto ensure_alignment(int64_t alignment) && -> TensorMatcher&& { + RuntimeCheck(!m_alignment.has_value(), "Alignment already specified"); + RuntimeCheck(is_pow2(alignment), "Alignment must be a power of 2"); + m_alignment = alignment; + return std::move(*this); + } + // once we start verification, we cannot modify anymore auto verify(tvm::ffi::TensorView view, DebugInfo info = {}) const&& -> const TensorMatcher&& { try { @@ -580,6 +590,18 @@ struct TensorMatcher { // since we may double verify, we will force to check m_dtype->verify(view.dtype()); m_device->verify(view.device()); + if (m_alignment.has_value()) { + const auto alignment = *m_alignment; + CHECK_HOST(std::bit_cast(view.data_ptr()) % alignment == 0) + << "Tensor data pointer is not aligned to " << alignment << " bytes"; + if (dim > 0) [[likely]] { + const auto bytes = static_cast(dtype_bytes(view.dtype())); + for (const auto i : irange(dim - 1)) { + CHECK_HOST(view.size(i) == 1 || (view.stride(i) * bytes) % alignment == 0) + << "Tensor stride for dimension " << i << " is not aligned to " << alignment << " bytes"; + } + } + } } auto m_init_dtype() -> void { @@ -602,6 +624,7 @@ struct TensorMatcher { DeviceRef m_device; bool m_has_dtype = false; bool m_has_device = false; + std::optional m_alignment; }; } // namespace host diff --git a/python/sglang/kernels/jit/include/sgl_kernel/tile.cuh b/python/sglang/kernels/jit/include/sgl_kernel/tile.cuh index 35786de13..332e167c1 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/tile.cuh +++ b/python/sglang/kernels/jit/include/sgl_kernel/tile.cuh @@ -34,7 +34,11 @@ struct Memory { return Memory{0, 1}; } /// \brief Create a Memory accessor distributed across warp threads. - SGL_DEVICE static Memory warp(int warp_threads = kWarpThreads) { + SGL_DEVICE static Memory warp() { + return Memory{get_lane_id(), kWarpThreads}; + } + /// \brief Create a Memory accessor over a narrower warp sub-group. + SGL_DEVICE static Memory warp(int warp_threads) { return Memory{static_cast(threadIdx.x % warp_threads), static_cast(warp_threads)}; } /// \brief Create a Memory accessor distributed across all CTA threads. diff --git a/python/sglang/kernels/jit/include/sgl_kernel/utils.cuh b/python/sglang/kernels/jit/include/sgl_kernel/utils.cuh index 53356ff89..505dd6c52 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/utils.cuh +++ b/python/sglang/kernels/jit/include/sgl_kernel/utils.cuh @@ -15,6 +15,7 @@ #pragma once +#include #include #include @@ -22,6 +23,8 @@ #include #include +#include +#include #include #ifndef USE_ROCM #include @@ -108,6 +111,7 @@ namespace device { /// \brief Macro: forced-inline device function qualifier. #define SGL_DEVICE __forceinline__ __device__ +#define SGL_DEVICE_HOST __forceinline__ __device__ __host__ // Architecture detection: SGL_CUDA_ARCH is injected by load_jit() and is // available in both host and device compilation passes, whereas __CUDA_ARCH__ @@ -133,13 +137,42 @@ static_assert( inline constexpr std::size_t kMaxVecBytes = SGL_ARCH_BLACKWELL_OR_GREATER ? 32 : 16; /// \brief Number of threads per warp (always 32 on NVIDIA/AMD GPUs). -inline constexpr auto kWarpThreads = 32u; -/// \brief Full warp active mask (all 32 lanes). +inline constexpr uint32_t kWarpThreads = 32u; +/// \brief Most implementations prefer this name; keep the alias for them. +inline constexpr uint32_t kWarpSize = kWarpThreads; + +/** + * \brief This thread's index within its logical `kNumThreads` group. + * + * \tparam kNumThreads Group width; a power of two, at most 32 on CUDA and at + * most 64 (the wave) on HIP -- so `64` is a HIP-only instantiation. + * + * \note Equals the true in-warp lane only when `blockDim.x` is a multiple of + * `kNumThreads`; every caller in this tree satisfies that. + * \note On CUDA prefer this over `threadIdx.x % kNumThreads` when the value + * feeds an address: `%laneid` is one register read that folds straight into + * `IMAD.WIDE`, while the modulo makes ptxas re-derive the mask at every address + * scale. Worth 8 instructions in a two-tile warp copy, measured on sm_100a. + * That only holds at full width -- a narrower group needs the mask anyway and + * ties with the modulo. + */ +template +SGL_DEVICE uint32_t get_lane_id() { #ifndef USE_ROCM -inline constexpr auto kFullMask = 0xffffffffu; + static_assert(kNumThreads <= 32 && host::is_pow2(kNumThreads)); + uint32_t lane_id; + asm volatile("mov.u32 %0, %%laneid;" : "=r"(lane_id)); + if constexpr (kNumThreads != 32) lane_id %= kNumThreads; + return lane_id; #else -inline constexpr auto kFullMask = 0xffffffffffffffffULL; + static_assert(kNumThreads <= 64 && host::is_pow2(kNumThreads)); + // AMD has no lane-id register: `__lane_id()` is computed from the exec mask as + // a `v_mbcnt_lo`/`v_mbcnt_hi` pair, and the group mask is still needed on top. + // Masking `threadIdx.x` -- already live in v0 -- is 2 instructions cheaper and + // yields the same value (measured on gfx950, hipcc 7.0). + return threadIdx.x % kNumThreads; #endif +} /** * \brief PDL (Programmatic Dependent Launch): wait for the primary kernel. @@ -147,6 +180,14 @@ inline constexpr auto kFullMask = 0xffffffffffffffffULL; * On Hopper (sm_90+), inserts a `griddepcontrol.wait` instruction to * synchronize with a preceding kernel in the same stream. On older * architectures or ROCm this is a no-op. + * + *\note This is the only thing that orders us against the producer. Per the PTX + * ISA, `.wait` makes the executing thread wait until every prerequisite grid in + * flight has COMPLETED and all of its memory operations are performed and made + * visible to this grid -- so it is what a `PDLTriggerSecondary` upstream does + * NOT give us. It acts per thread, so every thread that reads producer data has + * to execute it; put it ahead of the first such load. Stores into our own output + * buffers depend on nothing upstream and may be issued before it. */ template SGL_DEVICE void PDLWaitPrimary() { @@ -162,6 +203,22 @@ SGL_DEVICE void PDLWaitPrimary() { * * On Hopper (sm_90+), inserts a `griddepcontrol.launch_dependents` * instruction. On older architectures or ROCm this is a no-op. + * + * \note Scheduling only: this carries no memory ordering of its own. The + * dependent becomes eligible to launch once every CTA in this grid has issued + * the instruction or has exited, and it may then start before our writes are + * visible -- making them visible is the job of `PDLWaitPrimary` on the dependent + * side, which is why the programming guide requires the dependent to call it. + * + * Granularity is the CTA: the PTX ISA states that repeated invocations by + * threads of the same CTA have no side effect past the first, so one thread + * would do; we call it from all of them because it is free and needs no + * predication. Leaving it out altogether is safe and merely late, since the + * trigger is implied once every CTA exits (SASS code `PREEXIT`) + * + * Placing it early therefore costs nothing and only buys the dependent a head + * start on the work that does not depend on us. Even that is opportunistic: + * concurrent execution is never guaranteed, so nothing may rely on it. */ template SGL_DEVICE void PDLTriggerSecondary() { @@ -229,6 +286,41 @@ SGL_DEVICE void enable_smem_spilling() { #endif } +template +struct DeviceArray { + public: + SGL_DEVICE constexpr static std::size_t size() { + return N; + } + SGL_DEVICE constexpr auto operator[](std::size_t idx) -> T& { + return m_data[idx]; + } + SGL_DEVICE constexpr auto operator[](std::size_t idx) const -> const T& { + return m_data[idx]; + } + SGL_DEVICE constexpr auto data() const -> const T* { + return m_data; + } + SGL_DEVICE constexpr auto data() -> T* { + return m_data; + } + + private: + T m_data[N]; +}; + +/** + * Adapted from + * https://github.com/deepseek-ai/DeepGEMM/blob/559d79fb6994a58b8a15b4b93bf13ccc16edf247/deep_gemm/include/deep_gemm/common/utils.cuh + */ +SGL_DEVICE_HOST constexpr uint32_t get_tmem_cols(uint32_t num_cols) { + if (num_cols <= 32) return 32; + if (num_cols <= 64) return 64; + if (num_cols <= 128) return 128; + if (num_cols <= 256) return 256; + return 512; +} + } // namespace device namespace host { diff --git a/python/sglang/kernels/jit/include/sgl_kernel/vec.cuh b/python/sglang/kernels/jit/include/sgl_kernel/vec.cuh index 4c224d07f..ca8e36e88 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/vec.cuh +++ b/python/sglang/kernels/jit/include/sgl_kernel/vec.cuh @@ -118,6 +118,23 @@ struct AlignedVector { storage_t m_storage; }; +/// \brief Maximum vector width for coalesced memory access on GPU. +struct LoadStoreBytes { + enum type : int64_t { + MAX_GMEM = device::kMaxVecBytes, // architecture-dependent + MAX_SMEM = 16, // smem only support 16B op + MAX_PORTABLE = 16, // general across CUDA/HIP + MIN_COALALESCED = 4, // minimum for coalesced access + // some common vector widths for load/store + RAW_1B = 1, + RAW_2B = 2, + RAW_4B = 4, + RAW_8B = 8, + RAW_16B = 16, + RAW_32B = 32, + }; // namespace device +}; + /// Sum `M` vectors element-wise into one, accumulating in fp32 regardless of /// the packed element type. Used by every collective that reduces peer /// contributions in registers. diff --git a/python/sglang/kernels/jit/include/sgl_kernel/warp.cuh b/python/sglang/kernels/jit/include/sgl_kernel/warp.cuh index 811dc5315..f1a53d78d 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/warp.cuh +++ b/python/sglang/kernels/jit/include/sgl_kernel/warp.cuh @@ -2,55 +2,83 @@ /// \brief Warp-level reduction and cooperative-copy primitives. #pragma once +#include +#include + #include #include #include #include #include +#include #include namespace sglang { namespace device::warp { -/// \brief Full warp active mask. +/// \brief Full warp active mask and lane count: 32 on CUDA, wave64 on HIP. #ifndef USE_ROCM -static constexpr uint32_t kFullMask = 0xffffffffu; +inline constexpr uint32_t kFullMask = 0xffffffffu; +inline constexpr uint32_t kFullWidth = 32u; using mask_t = uint32_t; #else -static constexpr uint64_t kFullMask = 0xffffffffffffffffULL; +inline constexpr uint64_t kFullMask = 0xffffffffffffffffULL; +inline constexpr uint32_t kFullWidth = 64u; using mask_t = uint64_t; #endif +using ::sglang::device::get_lane_id; + +// One elected lane, via elect.sync. Raw PTX rather than cute::elect_one_sync, +// which would drag the whole CuTe include path into elementwise JIT modules; +// cuda::ptx has no elect_sync in CUDA 13.0. Use this to gate a single-thread +// TMA issue instead of a lane-index predicate. +SGL_DEVICE bool elect_one_lane() { + uint32_t pred; + asm volatile( + "{\n" + " .reg .pred p;\n" + " .reg .b32 r;\n" + " elect.sync r|p, 0xFFFFFFFF;\n" + " selp.b32 %0, 1, 0, p;\n" + "}\n" + : "=r"(pred)); + return pred != 0; +} + /** - * \brief Warp-level reduction. + * \brief Warp-level reduction over the lane range spanned by two widths. + * + * `kStart` and `kFinish` bound a range of lane-index bits: lane `i` reduces with + * every lane differing from it only in bits `[log2(lo), log2(hi))`, where `lo` + * and `hi` are the smaller and larger of the two. So `` reduces within + * contiguous groups of `N` lanes (`0..7, 8..15, ...` for `N = 8`), while + * `` reduces across groups at the same offset (`{0, 8, 16, 24}`). * - * On CUDA: uses __shfl_xor_sync with width=32. Full-warp reductions - * use a single `redux.sync` instruction when the target supports it. - * On HIP: uses __shfl_xor with explicit width parameter (supports wave64 sub-groups). * \tparam OP Reduction operation to perform (SUM, MAX, MIN). - * \tparam kNumThreads Number of threads as a group. - * \tparam kInner Whether to perform within a group or not. + * \tparam kStart One end of the reduced lane range; power of two. + * \tparam kFinish The other end of the reduced lane range; power of two. * \tparam T Type of the value to reduce. * * \param value The value to reduce. * \param active_mask The active mask of threads participating in the reduction. * - * \note We will divide into groups of `kNumThreads`. - * e.g. kNumThreads = 8, we have 0..7, 8..15, 16..23, 24..31 as groups. - * By reduction is performed within a group. Inter-group reduction will reduce - * over the same offset in different groups. e.g. {0, 8, 16, 24} in the above example. + * \note Symmetric: `` and `` reduce the same + * lane set, only walking the range in the opposite order. + * \note On CUDA a whole-warp reduction lowers to a single `redux.sync` where the + * target supports it. On HIP the shuffles use `__shfl_xor` with max width 64. */ -template +template SGL_DEVICE T reduce(T value, mask_t active_mask = kFullMask) { - static_assert(kNumThreads >= 1 && kNumThreads <= kWarpThreads); - static_assert(std::has_single_bit(kNumThreads), "must be pow of 2"); + static_assert(host::is_pow2(kStart) && host::is_pow2(kFinish)); + static_assert(kStart <= kFullWidth && kFinish <= kFullWidth); using Trait = ReductionTrait; #ifdef SGL_CUDA_ARCH // CUDA target only - constexpr bool kFullReduction = (kNumThreads == kWarpThreads && kInner) || (kNumThreads == 1 && !kInner); + constexpr bool kFullReduction = (kStart == 1 && kFinish == kFullWidth) || (kStart == kFullWidth && kFinish == 1); if constexpr (kFullReduction) { #if SGL_CUDA_ARCH >= 800 // 32 bit integer reduction @@ -81,24 +109,22 @@ SGL_DEVICE T reduce(T value, mask_t active_mask = kFullMask) { } #endif // redux.sync for CUDA only - if constexpr (kInner) { + if constexpr (kStart > kFinish) { #pragma unroll - for (uint32_t mask = kNumThreads / 2; mask >= 1; mask >>= 1) { + for (uint32_t mask = kStart / 2; mask >= kFinish; mask >>= 1) { #ifndef USE_ROCM - value = Trait::reduce(value, __shfl_xor_sync(active_mask, value, mask, 32)); + value = Trait::reduce(value, __shfl_xor_sync(active_mask, value, mask, kStart)); #else - value = Trait::reduce(value, __shfl_xor(value, mask, kNumThreads)); + value = Trait::reduce(value, __shfl_xor(value, mask, kStart)); #endif } } else { #pragma unroll - for (uint32_t mask = kNumThreads; mask <= kWarpThreads / 2; mask <<= 1) { + for (uint32_t mask = kStart; mask <= kFinish / 2; mask <<= 1) { #ifndef USE_ROCM - value = Trait::reduce(value, __shfl_xor_sync(active_mask, value, mask, 32)); + value = Trait::reduce(value, __shfl_xor_sync(active_mask, value, mask, kFinish)); #else - // Inter-group shuffle crosses kNumThreads-sized sub-groups, so the - // shuffle width must span the whole warp. - value = Trait::reduce(value, __shfl_xor(value, mask, kWarpThreads)); + value = Trait::reduce(value, __shfl_xor(value, mask, kFinish)); #endif } } @@ -106,108 +132,183 @@ SGL_DEVICE T reduce(T value, mask_t active_mask = kFullMask) { } /** \brief Warp-level sum reduction. */ -template +template SGL_DEVICE T reduce_sum(T value, mask_t active_mask = kFullMask) { - return reduce(value, active_mask); + return reduce(value, active_mask); } /** \brief Warp-level max reduction. */ -template +template SGL_DEVICE T reduce_max(T value, mask_t active_mask = kFullMask) { - return reduce(value, active_mask); + return reduce(value, active_mask); } /** \brief Warp-level min reduction. */ -template +template SGL_DEVICE T reduce_min(T value, mask_t active_mask = kFullMask) { - return reduce(value, active_mask); + return reduce(value, active_mask); } -/// \brief Warp-cooperative gmem -> smem copy of a compile-time byte count. -/// -/// Picks the widest vector width that divides both the per-thread share and -/// the byte total. The caller guarantees ``src`` is aligned to the picked -/// width (16B for kBytes % (16*32) == 0, else 8/4) and ``dst`` is the start -/// of a 16B-aligned per-warp smem slot. -// Warp-cooperative byte copy between any two address spaces, vectorised to the -// widest unit `kBytes` allows. Named for what it does rather than where it is -// used: the MLA call sites happen to target shared memory, but nothing here is -// global->shared specific -- no cp.async, no TMA, the payload moves through -// registers. -// -// The strategy was measured against the two async alternatives on B300 (sm_103, -// 148 SMs), copying one MLA row per warp out of a 512 MB pool so every row -// streams from HBM (grid 296, 64 rows/warp, 50 launches): -// -// 1152 B/warp (bf16, nope 1024 + rope 128) 576 B/warp (fp8, 512 + 64) -// this (generic) 47.3 us 3.69 TB/s 40.8 us 2.14 TB/s -// cp.async (ldgsts) 73.9 us 2.36 TB/s 56.4 us 1.55 TB/s -// cp.async.bulk/TMA 50.0 us 3.50 TB/s 43.2 us 2.02 TB/s -// -// The generic path wins at both sizes: a ~1 KB row is too small to amortise -// cp.async's per-lane 16 B issues or TMA's fixed issue plus mbarrier round trip. -// Revisit if a call site ever copies substantially more than one row per warp. -template -SGL_DEVICE void copy_bytes(const void* __restrict__ src, void* __restrict__ dst) { - 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; - 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::warp(); - +/** + * \brief Inclusive scan within each segment of `kWidth` lanes. + * + * Distinct from `reduce` above: every lane keeps its own running total rather + * than the whole-group result. Scans forward when `kStart < kFinish` (lane `p` + * accumulates lanes `<= p`) and backward when `kStart > kFinish` (lane `p` + * accumulates lanes `>= p`). The default `` is a plain + * forward scan over the whole segment. + * + * \tparam OP Reduction operation to combine with (SUM, MAX, MIN). + * \tparam kWidth Segment size; a scan never crosses a segment boundary. + * \tparam kStart First shuffle offset; power of two. `kStart > 1` assumes the + * input is already scanned in blocks of `kStart`, and scans the strided + * subsequences instead. + * \tparam kFinish Exclusive bound on the shuffle offset; power of two. + * \tparam T Type of the value to scan. + * + * \param val The value to scan. + * \param lane_id This thread's index WITHIN its segment, i.e. + * `threadIdx.x % kWidth` -- not `% kWarpThreads`. The shuffles are + * segment-relative but this predicate is not, so a warp-relative id silently + * corrupts every segment past the first whenever `kWidth < kWarpThreads`. + * \param active_mask The active mask of threads participating in the scan. + * + * \note The backward direction accumulates over `[p, kStart)`, so it covers the + * whole segment only when `kStart == kWidth`. + */ +template +SGL_DEVICE T inclusive_reduce(T val, uint32_t lane_id = get_lane_id(), mask_t active_mask = kFullMask) { + static_assert(host::is_pow2(kStart) && host::is_pow2(kFinish)); + static_assert(kStart <= kWidth && kFinish <= kWidth && kWidth <= kFullWidth); + using Trait = ReductionTrait; + if constexpr (kStart < kFinish) { #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); + for (uint32_t offset = kStart; offset < kFinish; offset *= 2) { + const auto n = __shfl_up_sync(active_mask, val, offset, kWidth); + if (lane_id >= offset) val = Trait::reduce(val, n); } - } -} - -/// Inclusive prefix sum across one warp, thread-rank order. Distinct from -/// reduce_sum above: every lane keeps its own running total rather than the -/// whole-warp result. -SGL_DEVICE uint32_t inclusive_sum(uint32_t lane_id, uint32_t val) { - static_assert(kWarpThreads == 32); + } else { #pragma unroll - for (uint32_t offset = 1; offset < 32; offset *= 2) { -#ifndef USE_ROCM - uint32_t n = __shfl_up_sync(0xFFFFFFFF, val, offset); -#else - uint32_t n = __shfl_up_sync(kFullMask, val, offset, kWarpThreads); -#endif - if (lane_id >= offset) val += n; + for (uint32_t offset = kFinish; offset < kStart; offset *= 2) { + const auto n = __shfl_down_sync(active_mask, val, offset, kWidth); + if (lane_id < kStart - offset) val = Trait::reduce(val, n); + } } return val; } -// One elected lane, via elect.sync. Raw PTX rather than cute::elect_one_sync, -// which would drag the whole CuTe include path into elementwise JIT modules; -// cuda::ptx has no elect_sync in CUDA 13.0. Use this to gate a single-thread -// TMA issue instead of a lane-index predicate. -SGL_DEVICE bool elect_one_lane() { - uint32_t pred; - asm volatile( - "{\n" - " .reg .pred p;\n" - " .reg .b32 r;\n" - " elect.sync r|p, 0xFFFFFFFF;\n" - " selp.b32 %0, 1, 0, p;\n" - "}\n" - : "=r"(pred)); - return pred != 0; +template +SGL_DEVICE T inclusive_sum(T val, uint32_t lane_id = get_lane_id(), mask_t active_mask = kFullMask) { + return inclusive_reduce(val, lane_id, active_mask); +} + +template +SGL_DEVICE T inclusive_max(T val, uint32_t lane_id = get_lane_id(), mask_t active_mask = kFullMask) { + return inclusive_reduce(val, lane_id, active_mask); +} + +template +SGL_DEVICE T inclusive_min(T val, uint32_t lane_id = get_lane_id(), mask_t active_mask = kFullMask) { + return inclusive_reduce(val, lane_id, active_mask); +} + +/** + * \brief Broadcast one lane's value to every lane of its `kWidth` segment. + * \param src_lane The source's index WITHIN the segment, i.e. in `[0, kWidth)`; + * each segment reads its own lane `src_lane`, not a single warp-wide one. + */ +template +SGL_DEVICE T broadcast(T value, uint32_t src_lane, mask_t active_mask = kFullMask) { + static_assert(host::is_pow2(kWidth) && kWidth <= kFullWidth); + return __shfl_sync(active_mask, value, src_lane, kWidth); +} + +namespace details { + +// array for load & store operations +template +struct Array : public DeviceArray { + static_assert(alignof(T) == sizeof(T)); + static constexpr int64_t kVecBytes = static_cast(sizeof(T)) * N; +}; + +} // namespace details + +template +struct CopyTrait { + static_assert(kBytes % kVecBytes == 0 && kBytes > 0); + using vec_t = AlignedStorage; + static constexpr int64_t kLoopBytes = sizeof(vec_t) * kWarpThreads; + static constexpr int64_t kLoopCount = kBytes / kLoopBytes; + static constexpr int64_t kTailBytes = kBytes - kLoopCount * kLoopBytes; + static constexpr int64_t kTailVecs = kTailBytes / sizeof(vec_t); + using result_t = details::Array 0 ? 1 : 0), kBytes>; + + template + SGL_DEVICE static void for_each(F&& f) { + const auto mem = tile::Memory::warp(); +#pragma unroll + for (int64_t i = 0; i < kLoopCount; ++i) { + f(mem, i); + } + if constexpr (kTailVecs > 0) { + if (mem.in_bound(kBytes / sizeof(vec_t), kLoopCount)) { + f(mem, kLoopCount); + } + } + } + + SGL_DEVICE static result_t load(const void* src) { + result_t result; + for_each([&](const auto& mem, int64_t i) { result[i] = mem.load(src, i); }); + return result; + } + + SGL_DEVICE static void store(void* dst, const result_t& result) { + for_each([&](const auto& mem, int64_t i) { mem.store(dst, result[i], i); }); + } +}; + +struct LoadStorePattern { + using enum LoadStoreBytes::type; + enum type : int64_t { + WARP_UNIFORM_GMEM = -MAX_GMEM, + WARP_UNIFORM_SMEM = -MAX_SMEM, + WARP_UNIFORM_4B = -4, + WARP_UNIFORM_8B = -8, + WARP_UNIFORM_16B = -16, + WARP_UNIFORM_32B = -32, + }; + + template + SGL_DEVICE_HOST static constexpr int64_t get_vec_bytes() { + if constexpr (kMaxVecBytes < 0) { // best-effort warp uniform load/store + if constexpr (kBytes % (4 * device::kWarpThreads) != 0) { + // at least guarantee 128B coalesced for better performance + return std::gcd(kBytes, 4); + } else { // kBytes is at least 128B coalesced + return std::gcd(kBytes / device::kWarpThreads, -kMaxVecBytes); + } + } else { + return std::gcd(kBytes, kMaxVecBytes); + } + } +}; + +template < + int64_t kBytes, + int64_t kMaxVecBytes = LoadStorePattern::MAX_GMEM, + int64_t kVecBytes = LoadStorePattern::get_vec_bytes()> +SGL_DEVICE auto load_bytes(const void* src) { + return CopyTrait::load(src); +} + +template < + int64_t kBytes, + int64_t kMaxVecBytes = LoadStorePattern::MAX_GMEM, + int64_t kVecBytes = LoadStorePattern::get_vec_bytes()> +SGL_DEVICE void store_bytes(void* dst, const auto& result) { + return CopyTrait::store(dst, result); } } // namespace device::warp diff --git a/python/sglang/kernels/jit/utils/compile/loader.py b/python/sglang/kernels/jit/utils/compile/loader.py index 7b810a445..7dffa00b4 100644 --- a/python/sglang/kernels/jit/utils/compile/loader.py +++ b/python/sglang/kernels/jit/utils/compile/loader.py @@ -86,6 +86,14 @@ def load_jit( if flag not in ("--use_fast_math", "-use_fast_math") ] + if envs.SGLANG_JIT_LOG_RESOURCE_USAGE.get(): + # nvcc reports through ptxas; hipcc through a clang remark pass. + extra_cuda_cflags = list(extra_cuda_cflags or []) + ( + ["-Rpass-analysis=kernel-resource-usage"] + if is_hip_runtime() + else ["-Xptxas=-v"] + ) + includes = list(DEFAULT_INCLUDE) + (extra_include_paths or []) for dep in sorted(set(extra_dependencies or [])): if dep not in REGISTERED_DEPENDENCIES: @@ -113,7 +121,7 @@ def load_jit( build_key = cache.compute_build_key(spec, build_file=build_file) scope = cache.build_key_dir(module_name=spec.module_name, build_key=build_key) - prebuilt = cache.find_prebuilt(scope=scope, module_name=spec.module_name) + prebuilt = _find_prebuilt(spec=spec, scope=scope) if prebuilt is not None: try: return _load(prebuilt) @@ -141,7 +149,7 @@ def load_jit( # published exactly what we were about to build. This is what turns N # tensor-parallel ranks starting together into one compile plus N-1 # cache hits instead of N identical compiles. - prebuilt = cache.find_prebuilt(scope=scope, module_name=spec.module_name) + prebuilt = _find_prebuilt(spec=spec, scope=scope) if prebuilt is not None: try: return _load(prebuilt) @@ -182,6 +190,18 @@ def load_jit( shutil.rmtree(staging, ignore_errors=True) +def _find_prebuilt(*, spec: BuildSpec, scope: pathlib.Path) -> pathlib.Path | None: + """The cached build to reuse, or None when there is nothing to reuse. + + Returns None unconditionally under `SGLANG_JIT_FORCE_RECOMPILE`, which is + what makes the compiler run again. Both lookups go through here, so the flag + cannot take effect on the fast path and not on the one behind the lock. + """ + if envs.SGLANG_JIT_FORCE_RECOMPILE.get(): + return None + return cache.find_prebuilt(scope=scope, module_name=spec.module_name) + + @contextlib.contextmanager def _build_lock(scope: pathlib.Path): """Serialize builds of one module variant across processes. diff --git a/python/sglang/kernels/jit/utils/compile/ninja.py b/python/sglang/kernels/jit/utils/compile/ninja.py index 4045c7831..a54b5bc01 100644 --- a/python/sglang/kernels/jit/utils/compile/ninja.py +++ b/python/sglang/kernels/jit/utils/compile/ninja.py @@ -28,6 +28,7 @@ from typing import List from sglang.kernels.jit.utils.compile import toolchain from sglang.kernels.jit.utils.compile.spec import BuildSpec +from sglang.srt.environ import envs logger = logging.getLogger(__name__) @@ -169,9 +170,38 @@ def build(*, spec: BuildSpec, build_dir: pathlib.Path, build_file: str) -> pathl f"Failed to build JIT module {spec.module_name} in {build_dir}\n" f"stdout:\n{completed.stdout}\nstderr:\n{completed.stderr}" ) + if envs.SGLANG_JIT_LOG_RESOURCE_USAGE.get(): + _log_resource_usage( + spec.module_name, (completed.stdout or "") + (completed.stderr or "") + ) return build_dir / f"{spec.module_name}.so" +# What the two device compilers call their resource report. nvcc routes it +# through ptxas; hipcc emits clang remarks naming the analysis pass. +_RESOURCE_MARKERS = ("ptxas info", "spill", "kernel-resource-usage") + + +def _log_resource_usage(module_name: str, output: str) -> None: + """Replay the compiler's per-kernel resource report. + + The build runs with `capture_output=True` and replays only on failure, so + without this the report is produced and then dropped. Ninja forwards each + subcommand's diagnostics onto its own stdout, so that is where this lands + regardless of which stream the compiler wrote to. + """ + report = [ + line.rstrip() + for line in output.splitlines() + # The echoed compile command also contains the flag that asked for the + # report, so match on the output's own markers, not on the flag. + if not line.startswith("[") and any(m in line for m in _RESOURCE_MARKERS) + ] + if not report: + return + logger.info("JIT resource usage for %s:\n%s", module_name, "\n".join(report)) + + def scan_dependencies(build_dir: pathlib.Path) -> List[pathlib.Path]: """Every file the compiler read, taken from the depfiles the build left. diff --git a/python/sglang/kernels/ops/kvcache/kvcache.py b/python/sglang/kernels/ops/kvcache/kvcache.py index 98a012f6b..bb9814361 100644 --- a/python/sglang/kernels/ops/kvcache/kvcache.py +++ b/python/sglang/kernels/ops/kvcache/kvcache.py @@ -16,10 +16,30 @@ from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: from tvm_ffi.module import Module +logger = logging.getLogger(__name__) + +# Mirrors device::kWarpThreads in include/sgl_kernel/utils.cuh (32 on CUDA and HIP). +_WARP_THREADS = 32 + @cache_once -def _jit_kvcache_module(k_row_bytes: int, v_row_bytes: int) -> Module: - args = make_cpp_args(k_row_bytes, v_row_bytes, is_arch_support_pdl()) +def _jit_kvcache_module(k_row_bytes: int, v_row_bytes: int, num_threads: int) -> Module: + if num_threads == 0: + num_threads = 32 + # rare case. just don't optimize it + if k_row_bytes % num_threads != 0 or v_row_bytes % num_threads != 0: + return _jit_kvcache_module(k_row_bytes, v_row_bytes, num_threads) + k_bytes = k_row_bytes / num_threads + v_bytes = v_row_bytes / num_threads + # increase threads if row is too large + while k_bytes % 8 == 0 and v_bytes % 8 == 0 and (k_bytes + v_bytes) >= 64: + num_threads *= 2 + k_bytes /= 2 + v_bytes /= 2 + logger.debug(f"Heuristic {num_threads = } for {k_row_bytes = }, {v_row_bytes}") + return _jit_kvcache_module(k_row_bytes, v_row_bytes, num_threads) + + args = make_cpp_args(k_row_bytes, v_row_bytes, num_threads, is_arch_support_pdl()) return load_jit( "kvcache", *args, @@ -29,20 +49,14 @@ def _jit_kvcache_module(k_row_bytes: int, v_row_bytes: int) -> Module: @cache_once -def can_use_store_cache(k_row_bytes: int, v_row_bytes: int = 0) -> bool: +def can_use_store_cache( + k_row_bytes: int, v_row_bytes: int = 0, num_threads: int = 0 +) -> bool: """Whether the JIT store_cache kernel can serve these row widths. v_row_bytes=0 means symmetric, i.e. it defaults to k_row_bytes.""" - logger = logging.getLogger(__name__) v_row_bytes = v_row_bytes or k_row_bytes - for name, size in (("k_row_bytes", k_row_bytes), ("v_row_bytes", v_row_bytes)): - if size % 4 != 0: - logger.warning( - f"Unsupported {name}={size} for JIT KV-Cache kernel:" - " must be multiple of 4" - ) - return False try: - _jit_kvcache_module(k_row_bytes, v_row_bytes) + _jit_kvcache_module(k_row_bytes, v_row_bytes, num_threads) return True except Exception as e: logger.warning( @@ -62,7 +76,7 @@ def store_cache( *, row_bytes: int = 0, v_row_bytes: int = 0, - num_split: int = 0, # can be tuned for performance + num_split: int = 0, size_limit: int = 0, reserved_skip_index: int = 0, ) -> None: @@ -77,6 +91,8 @@ def store_cache( row_bytes (int): Key row width in bytes. Inferred from k when 0. v_row_bytes (int): Value row width in bytes; differs from row_bytes for asymmetric KV (head_dim != v_head_dim). Inferred from v when 0. + num_split (int): Warps cooperating on one row. A heuristic picks it + when 0; it is the only knob here that exists purely for tuning. size_limit (int): Valid slot bound (cache row count = real slots + the reserved padding slot); an index outside [0, size_limit) fails fast (device assert) instead of an illegal memory access. Defaults to the @@ -87,15 +103,10 @@ def store_cache( """ row_bytes = row_bytes or k.shape[-1] * k.element_size() v_row_bytes = v_row_bytes or v.shape[-1] * v.element_size() - module = _jit_kvcache_module(row_bytes, v_row_bytes) - if num_split <= 0: - # A split must divide BOTH rows, so require the alignment on each. - if row_bytes % 2048 == 0 and v_row_bytes % 2048 == 0: - num_split = 4 - elif row_bytes % 1024 == 0 and v_row_bytes % 1024 == 0: - num_split = 2 - else: - num_split = 1 + # One warp per split. The knob stays the split count it has always been: + # renaming it changes the registered op schema, and a warm inductor cache + # does not notice that -- it replays generated code carrying the old name. + module = _jit_kvcache_module(row_bytes, v_row_bytes, num_split * _WARP_THREADS) if size_limit <= 0: size_limit = k_cache.shape[0] module.store_cache( @@ -104,7 +115,6 @@ def store_cache( k_cache, v_cache, indices, - num_split, size_limit, reserved_skip_index, ) diff --git a/python/sglang/kernels/ops/kvcache/mla_buffer.py b/python/sglang/kernels/ops/kvcache/mla_buffer.py index 99ff8b8d2..8ba7d287b 100644 --- a/python/sglang/kernels/ops/kvcache/mla_buffer.py +++ b/python/sglang/kernels/ops/kvcache/mla_buffer.py @@ -117,13 +117,6 @@ def set_mla_kv_buffer_kernel_norope( 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_impl( kv_buffer: torch.Tensor, loc: torch.Tensor, @@ -136,19 +129,16 @@ def _set_mla_kv_buffer_impl( ): """Dispatch MLA paged-KV scatter writes to the fastest available path. - Two paths, chosen on ``n_loc``: + Two paths: - - ``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. + - 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]]``. It packs 4-8 items per CTA, so the CTA count + falls well below 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 (5x 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. + i.e. one CTA per loc covering the entire row in one tile. This is 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.05x at bs=8, ~1.5x at bs=128, 3.5x at bs=512, **11.7x at @@ -195,8 +185,7 @@ def _set_mla_kv_buffer_impl( 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() + is_arch_support_pdl() and can_use_set_mla_kv_buffer(nope_bytes, rope_bytes) and dcp_world_size == 1 ): diff --git a/python/sglang/kernels/ops/kvcache/set_mla_kv_buffer.py b/python/sglang/kernels/ops/kvcache/set_mla_kv_buffer.py index e6b557147..7ba095ac0 100644 --- a/python/sglang/kernels/ops/kvcache/set_mla_kv_buffer.py +++ b/python/sglang/kernels/ops/kvcache/set_mla_kv_buffer.py @@ -41,27 +41,7 @@ def set_mla_kv_buffer_module(nope_bytes: int, rope_bytes: int, use_pdl: bool) -> @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, - ) + if (rope_bytes + nope_bytes) % 16 != 0: return False try: set_mla_kv_buffer_module(nope_bytes, rope_bytes, is_arch_support_pdl()) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 3ac57ca59..d13e1b4f6 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -1186,6 +1186,13 @@ class Envs: # SGLANG_CACHE_DIR; set to an empty string to keep compilation # process-local. Must be trusted: cached objects are loaded into the process. SGLANG_CUTE_AOT_CACHE_DIR = EnvStr(lambda: _default_cache_subdir("cute_aot")) + + # =================================================================== + # Kernel development: JIT build cache, diagnostics and benchmarks + # =================================================================== + # Everything here is a developer knob for working ON kernels -- building + # them, inspecting what the compiler produced, and benchmarking them. Flags + # that select a kernel in production live with their own feature instead. # JIT kernel build cache. None = unset, resolving to ~/.cache/sglang/jit; # point it at a persistent mount to share builds across CI jobs. SGLANG_JIT_CACHE_DIR = EnvStr(None) @@ -1195,10 +1202,23 @@ class Envs: # is what makes reverting an edit an instant hit instead of a rebuild; set # it to trade that away for disk (1 keeps only the most recent build). SGLANG_JIT_CACHE_KEEP = EnvInt(None) + # Skip the cache lookup and run the compiler for every module this process + # loads. The result is still published, so the cost is one rebuild per + # module, not one per load. + SGLANG_JIT_FORCE_RECOMPILE = EnvBool(False) # Raise instead of compiling when a module misses the cache, so a # deployment that expects a pre-seeded cache fails loudly at startup # rather than silently eating a cold compile. SGLANG_CRASH_ON_JIT_COMPILE = EnvBool(False) + # Ask the device compiler for per-kernel resource usage (registers, spills, + # shared memory) and log it at INFO. Changes the build flags, so it compiles + # into its own cache entry and leaves the normal one alone -- but that entry + # is a hit on the second run, and a cache hit has nothing to report, so pair + # this with SGLANG_JIT_FORCE_RECOMPILE to see the report every time. + SGLANG_JIT_LOG_RESOURCE_USAGE = EnvBool(False) + # Drop the GB/s and TFLOPS columns from the benchmark marker's table. + SGLANG_JIT_BENCHMARK_DISABLE_LOG_BANDWIDTH = EnvBool(False) + SGLANG_JIT_BENCHMARK_DISABLE_LOG_FLOPS = EnvBool(False) # =================================================================== # Expert-parallel dispatch and MoE execution diff --git a/test/registered/kernels/benchmark/kvcache/bench_hicache.py b/test/registered/kernels/benchmark/kvcache/bench_hicache.py index 054f3e2ea..6d3aec687 100644 --- a/test/registered/kernels/benchmark/kvcache/bench_hicache.py +++ b/test/registered/kernels/benchmark/kvcache/bench_hicache.py @@ -15,19 +15,15 @@ Note: Uses do_bench instead of do_bench_cudagraph since CUDA graph capture doesn't support CPU-GPU memory transfers. """ -import itertools import os from dataclasses import dataclass -from typing import Tuple import torch -import triton -import triton.testing from sgl_kernel import transfer_kv_all_layer, transfer_kv_per_layer -from sglang.kernels.jit.benchmark.utils import DEFAULT_QUANTILES, get_benchmark_range +from sglang.kernels.jit.benchmark import marker +from sglang.kernels.jit.benchmark.utils import get_benchmark_range from sglang.kernels.ops.kvcache.hicache import ( - can_use_hicache_jit_kernel, transfer_hicache_all_layer, transfer_hicache_one_layer, ) @@ -39,7 +35,7 @@ register_cuda_ci( register_amd_ci(est_time=29, stage="jit-kernel-benchmark", runner_config="amd") DISABLE_TORCH = os.environ.get("DISABLE_TORCH", "0") == "1" -PAGE_SIZE = 1 +PAGE_SIZE = int(os.environ.get("PAGE_SIZE", "1")) ENABLE_SORT = True GPU_CACHE_SIZE = 256 * 1024 # 256K tokens on GPU HOST_CACHE_SIZE = 512 * 1024 # 512K tokens on CPU @@ -187,20 +183,14 @@ def pytorch_transfer( # Benchmark configuration -BS_RANGE = get_benchmark_range( - full_range=[2**n for n in range(0, 16)], - ci_range=[16], -) ELEMENT_SIZE_RANGE = get_benchmark_range( full_range=[64, 128, 256, 512, 1024], ci_range=[1024], ) LINE_VALS = ["aot", "jit", "torch"] -LINE_NAMES = ["SGL AOT Kernel", "SGL JIT Kernel", "PyTorch"] -STYLES = [("orange", "-"), ("blue", "--"), ("red", ":")] - -CONFIGS = list(itertools.product(ELEMENT_SIZE_RANGE, BS_RANGE)) +if DISABLE_TORCH: + LINE_VALS.remove("torch") # ============================================================================= @@ -208,22 +198,10 @@ CONFIGS = list(itertools.product(ELEMENT_SIZE_RANGE, BS_RANGE)) # ============================================================================= -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["element_size", "batch_size"], - x_vals=CONFIGS, - line_arg="provider", - line_vals=LINE_VALS, - line_names=LINE_NAMES, - styles=STYLES, - ylabel="us", - plot_name="hicache-one-layer-h2d", - args={}, - ) -) -def benchmark_one_layer_h2d( - element_size: int, batch_size: int, provider: str -) -> Tuple[float, float, float]: +@marker.parametrize("element_size", ELEMENT_SIZE_RANGE) +@marker.parametrize("batch_size", marker.range(14, pattern="pow2"), [16]) +@marker.benchmark("provider", LINE_VALS, unit="ms") +def benchmark_one_layer_h2d(element_size: int, batch_size: int, provider: str): """One Layer: Host (CPU) -> Device (GPU).""" global cache cache_local = cache.get_slice(num_layers=NUM_LAYERS, element_size=element_size) @@ -281,19 +259,10 @@ def benchmark_one_layer_h2d( ], } - if provider == "jit" and not can_use_hicache_jit_kernel(element_size=element_bytes): - return (float("nan"), float("nan"), float("nan")) - - if DISABLE_TORCH and provider in ["torch"]: - return (float("nan"), float("nan"), float("nan")) - - ms, min_ms, max_ms = triton.testing.do_bench( # type: ignore - FN_MAP[provider], quantiles=DEFAULT_QUANTILES, warmup=5, rep=25 - ) - return ( - 1000 * ms / NUM_LAYERS, - 1000 * max_ms / NUM_LAYERS, - 1000 * min_ms / NUM_LAYERS, + return marker.do_bench( + FN_MAP[provider], + use_cuda_graph=False, + extra_memory_footprint=NUM_LAYERS * batch_size * (2 * element_bytes), ) @@ -311,22 +280,10 @@ def _create_ptr_tensor(tensors, device="cuda"): ) -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["element_size", "batch_size"], - x_vals=CONFIGS, - line_arg="provider", - line_vals=LINE_VALS, - line_names=LINE_NAMES, - styles=STYLES, - ylabel="us", - plot_name="hicache-all-layer-d2h", - args={}, - ) -) -def benchmark_all_layer_d2h( - element_size: int, batch_size: int, provider: str -) -> Tuple[float, float, float]: +@marker.parametrize("element_size", ELEMENT_SIZE_RANGE) +@marker.parametrize("batch_size", marker.range(14, pattern="pow2"), [16]) +@marker.benchmark("provider", LINE_VALS, unit="ms") +def benchmark_all_layer_d2h(element_size: int, batch_size: int, provider: str): """All Layer: Device (GPU) -> Host (CPU).""" global cache cache_local = cache.get_slice(num_layers=NUM_LAYERS, element_size=element_size) @@ -385,19 +342,10 @@ def benchmark_all_layer_d2h( ], } - if provider == "jit" and not can_use_hicache_jit_kernel(element_size=element_bytes): - return (float("nan"), float("nan"), float("nan")) - - if DISABLE_TORCH and provider in ["torch"]: - return (float("nan"), float("nan"), float("nan")) - - ms, min_ms, max_ms = triton.testing.do_bench( # type: ignore - FN_MAP[provider], quantiles=DEFAULT_QUANTILES, warmup=5, rep=25 - ) - return ( - 1000 * ms / NUM_LAYERS, - 1000 * max_ms / NUM_LAYERS, - 1000 * min_ms / NUM_LAYERS, + return marker.do_bench( + FN_MAP[provider], + use_cuda_graph=False, + extra_memory_footprint=NUM_LAYERS * batch_size * (2 * element_bytes), ) @@ -413,12 +361,5 @@ if __name__ == "__main__": v_cache_host=torch.empty(HOST_SHAPE, dtype=torch.bfloat16, pin_memory=True), ) - print("=" * 60) - print("One Layer: Host -> Device (CPU -> GPU)") - print("=" * 60) - benchmark_one_layer_h2d.run(print_data=True) - - print("\n" + "=" * 60) - print("All Layer: Device -> Host (GPU -> CPU) [per-layer avg]") - print("=" * 60) - benchmark_all_layer_d2h.run(print_data=True) + benchmark_one_layer_h2d.run(print_prefix="Per Layer: Host -> Device (CPU -> GPU)") + benchmark_all_layer_d2h.run(print_prefix="All Layer: Device -> Host (GPU -> CPU)") diff --git a/test/registered/kernels/benchmark/kvcache/bench_set_mla_kv_buffer.py b/test/registered/kernels/benchmark/kvcache/bench_set_mla_kv_buffer.py index 3af514ef5..e1d31186c 100644 --- a/test/registered/kernels/benchmark/kvcache/bench_set_mla_kv_buffer.py +++ b/test/registered/kernels/benchmark/kvcache/bench_set_mla_kv_buffer.py @@ -7,18 +7,13 @@ Compares three providers across a batch-size sweep: - ``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.kernels.jit.benchmark import marker from sglang.kernels.jit.benchmark.utils import ( DEFAULT_DEVICE, DEFAULT_DTYPE, - DEFAULT_QUANTILES, - get_benchmark_range, ) from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.kernels.ops.kvcache.set_mla_kv_buffer import set_mla_kv_buffer as jit_set @@ -50,62 +45,38 @@ def _triton_baseline(kv_buffer, loc, cache_k_nope, cache_k_rope): cache_k_rope.stride(0), nope_dim, rope_dim, - BLOCK=BLOCK, - DCP_RANK=0, - DCP_WORLD_SIZE=1, - **pdl_kwargs, + BLOCK=BLOCK, # type: ignore + DCP_RANK=0, # type: ignore + DCP_WORLD_SIZE=1, # type: ignore + **pdl_kwargs, # type: ignore ) -NUM_LAYERS = 8 -CACHE_SIZE = (2 * 1024 * 1024) // NUM_LAYERS - +# 2M elements +CACHE_SIZE = 2 * 1024 * 1024 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]: +@marker.parametrize("batch_size", marker.range(15, pattern="pow2"), [1, 128, 8192]) +@marker.benchmark("provider", ["wrapper", "jit_tma", "triton"]) +def benchmark(batch_size: int, provider: str): cache_k_nope = torch.randn( - (NUM_LAYERS, batch_size, 1, NOPE_DIM), + (batch_size, 1, NOPE_DIM), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE, ) cache_k_rope = torch.randn( - (NUM_LAYERS, batch_size, 1, ROPE_DIM), + (batch_size, 1, ROPE_DIM), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE, ) kv_buffer = torch.randn( - (NUM_LAYERS, CACHE_SIZE, 1, NOPE_DIM + ROPE_DIM), + (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, @@ -113,20 +84,14 @@ def benchmark(batch_size: int, provider: str) -> Tuple[float, float, float]: "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, + return marker.do_bench( + FN_MAP[provider], + input_args=(kv_buffer, loc, cache_k_nope, cache_k_rope), + graph_clone_args=(1, 2, 3), + memory_args=(loc, cache_k_nope, cache_k_rope), + memory_output=(cache_k_nope, cache_k_rope), ) if __name__ == "__main__": - benchmark.run(print_data=True) + benchmark.run() diff --git a/test/registered/kernels/ops/kvcache/test_store_cache.py b/test/registered/kernels/ops/kvcache/test_store_cache.py index 6e0be27f3..17b898a2c 100644 --- a/test/registered/kernels/ops/kvcache/test_store_cache.py +++ b/test/registered/kernels/ops/kvcache/test_store_cache.py @@ -17,7 +17,7 @@ BS_LIST = [2**n for n in range(0, 15)] BS_LIST += [x + 1 + i for i, x in enumerate(BS_LIST)] BS_LIST = get_ci_test_range(BS_LIST, [1, 9, 256, 16399]) HIDDEN_DIMS = get_ci_test_range( - [64, 128, 256, 512, 1024, 96, 98, 100], [64, 512, 1024, 98] + [64, 128, 256, 512, 1024, 96, 97, 100], [64, 512, 1024, 97] ) CACHE_SIZE = 1024 * 1024 DTYPE = torch.bfloat16 @@ -35,7 +35,6 @@ def test_store_cache(batch_size: int, element_dim: int) -> None: v_cache = torch.randn((CACHE_SIZE, element_dim), dtype=DTYPE, device=DEVICE) indices = torch.randperm(CACHE_SIZE - 1, device=DEVICE)[:batch_size] + 1 - # AOT store cache store_cache(k, v, k_cache, v_cache, indices) assert torch.all(k_cache[indices] == k) @@ -89,10 +88,7 @@ def test_store_cache_int32_indices(batch_size: int, element_dim: int) -> None: @pytest.mark.parametrize("index_dtype", [torch.int32, torch.int64]) -@pytest.mark.parametrize("num_split", [1, 2, 4]) -def test_store_cache_reserved_skip_index( - index_dtype: torch.dtype, num_split: int -) -> None: +def test_store_cache_reserved_skip_index(index_dtype: torch.dtype) -> None: element_dim = 1024 k = torch.randn((4, element_dim), dtype=DTYPE, device=DEVICE) v = torch.randn((4, element_dim), dtype=DTYPE, device=DEVICE) @@ -112,7 +108,6 @@ def test_store_cache_reserved_skip_index( k_cache, v_cache, indices, - num_split=num_split, ) torch.testing.assert_close(k_cache[0], reserved_k_before, rtol=0.0, atol=0.0) @@ -137,43 +132,6 @@ def test_store_cache_zero_index_can_be_written_when_skip_disabled() -> None: torch.testing.assert_close(v_cache[0], v[0], rtol=0.0, atol=0.0) -def _valid_num_splits(element_dim: int, dtype: torch.dtype) -> list: - """Return the list of valid num_split values for a given element_dim/dtype.""" - row_bytes = element_dim * dtype.itemsize - splits = [1] - if row_bytes % (2 * 128) == 0: - splits.append(2) - if row_bytes % (4 * 128) == 0: - splits.append(4) - return splits - - -_NUM_SPLIT_CASES = [ - (_dim, _ns, _dtype) - for _dtype in [torch.float16, torch.bfloat16, torch.float32] - for _dim in REPR_DIMS - for _ns in _valid_num_splits(_dim, _dtype) -] - - -@pytest.mark.parametrize("element_dim,num_split,dtype", _NUM_SPLIT_CASES) -def test_store_cache_num_split( - element_dim: int, num_split: int, dtype: torch.dtype -) -> None: - batch_size = 128 - k = torch.randn((batch_size, element_dim), dtype=dtype, device=DEVICE) - v = torch.randn((batch_size, element_dim), dtype=dtype, device=DEVICE) - k_cache = torch.randn((SMALL_CACHE, element_dim), dtype=dtype, device=DEVICE) - v_cache = torch.randn((SMALL_CACHE, element_dim), dtype=dtype, device=DEVICE) - indices = torch.randperm(SMALL_CACHE - 1, device=DEVICE)[:batch_size] + 1 - - # Verify each num_split kernel path (1, 2, 4) produces correct results - store_cache(k, v, k_cache, v_cache, indices, num_split=num_split) - - assert torch.all(k_cache[indices] == k) - assert torch.all(v_cache[indices] == v) - - # Asymmetric K/V (head_dim != v_head_dim): different row widths AND cache strides. # MiMoV2 is 192/128. Both orderings, since nothing may assume K is the wider one. ASYM_DIM_PAIRS = get_ci_test_range( @@ -208,55 +166,6 @@ def test_store_cache_asymmetric(k_dim: int, v_dim: int, dtype: torch.dtype) -> N assert torch.all(v_cache[untouched] == v_before[untouched]) -def _valid_asym_num_splits(k_dim: int, v_dim: int, dtype: torch.dtype) -> list: - """num_split values valid for BOTH rows; a split must divide each of them.""" - k_bytes, v_bytes = k_dim * dtype.itemsize, v_dim * dtype.itemsize - splits = [1] - if k_bytes % (2 * 128) == 0 and v_bytes % (2 * 128) == 0: - splits.append(2) - if k_bytes % (4 * 128) == 0 and v_bytes % (4 * 128) == 0: - splits.append(4) - return splits - - -def _default_num_split(k_dim: int, v_dim: int, dtype: torch.dtype) -> int: - """Mirrors the heuristic in store_cache(); the default is already exercised - by test_store_cache_asymmetric, which does not pass num_split.""" - k_bytes, v_bytes = k_dim * dtype.itemsize, v_dim * dtype.itemsize - if k_bytes % 2048 == 0 and v_bytes % 2048 == 0: - return 4 - if k_bytes % 1024 == 0 and v_bytes % 1024 == 0: - return 2 - return 1 - - -# Only splits the default heuristic would NOT pick: the split gate is two-sided -# (K and V must both align), so the off-default branches are what needs pinning. -_ASYM_NUM_SPLIT_CASES = [ - (_k, _v, _ns) - for _k, _v in ASYM_DIM_PAIRS - for _ns in _valid_asym_num_splits(_k, _v, DTYPE) - if _ns != _default_num_split(_k, _v, DTYPE) -] - - -@pytest.mark.parametrize("k_dim,v_dim,num_split", _ASYM_NUM_SPLIT_CASES) -def test_store_cache_asymmetric_num_split( - k_dim: int, v_dim: int, num_split: int -) -> None: - batch_size = 128 - k = torch.randn((batch_size, k_dim), dtype=DTYPE, device=DEVICE) - v = torch.randn((batch_size, v_dim), dtype=DTYPE, device=DEVICE) - k_cache = torch.randn((SMALL_CACHE, k_dim), dtype=DTYPE, device=DEVICE) - v_cache = torch.randn((SMALL_CACHE, v_dim), dtype=DTYPE, device=DEVICE) - indices = torch.randperm(SMALL_CACHE - 1, device=DEVICE)[:batch_size] + 1 - - store_cache(k, v, k_cache, v_cache, indices, num_split=num_split) - - assert torch.all(k_cache[indices] == k) - assert torch.all(v_cache[indices] == v) - - def test_can_use_store_cache() -> None: assert can_use_store_cache(128) assert can_use_store_cache(256)