diff --git a/python/sglang/jit_kernel/benchmark/kv_canary/bench_verify.py b/python/sglang/jit_kernel/benchmark/kv_canary/bench_verify.py new file mode 100644 index 000000000..043e0b65a --- /dev/null +++ b/python/sglang/jit_kernel/benchmark/kv_canary/bench_verify.py @@ -0,0 +1,286 @@ +from __future__ import annotations + +from typing import Tuple + +import torch +import triton +import triton.testing + +from sglang.jit_kernel.benchmark.kv_canary.utils import ( + RING_CAPACITY, + SWA_WINDOW, + BenchCase, + build_fast_matrix_cases, + build_full_matrix_cases, + cases_to_x_vals, + naive_slot_copy_fn, +) +from sglang.jit_kernel.benchmark.utils import ( + DEFAULT_DEVICE, + get_benchmark_range, + run_benchmark, +) +from sglang.jit_kernel.kv_canary import consts +from sglang.jit_kernel.kv_canary.verify import ( + CANARY_SLOT_BYTES, + CanaryLaunchTag, + VerifyOrWriteContext, + VerifyPlan, + launch_canary_verify_kernel, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=900, suite="nightly-kernel-1-gpu", nightly=True) + + +_X_NAMES = [ + "scenario", + "bs", + "prefix_len", + "mode", + "extend_len", + "pool_kind", +] +_X_VALS = cases_to_x_vals( + get_benchmark_range( + full_range=build_full_matrix_cases(), + ci_range=build_fast_matrix_cases(), + ) +) + +_KERNEL_KIND_X_NAMES = ["kernel_kind_name"] +_KERNEL_KIND_X_VALS = [(tag.name,) for tag in CanaryLaunchTag] + + +def _verify_entry_count(case: BenchCase) -> int: + if case.pool_kind == "swa_window_128": + per_req = min(case.prefix_len, SWA_WINDOW) + else: + per_req = case.prefix_len + return case.bs * per_req + + +def _verify_num_slots(case: BenchCase) -> int: + if case.pool_kind == "swa_window_128": + per_req_slots = SWA_WINDOW + else: + per_req_slots = max(1, case.prefix_len) + return max(2, case.bs * per_req_slots + 1) + + +def _build_verify_inputs(case: BenchCase, *, device: torch.device) -> Tuple[ + torch.Tensor, + VerifyPlan, + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, +]: + total_entries = _verify_entry_count(case) + capacity = max(1, total_entries) + num_slots = _verify_num_slots(case) + + canary_buf = torch.zeros( + num_slots, CANARY_SLOT_BYTES, dtype=torch.uint8, device=device + ) + + slot_indices = torch.empty(capacity, dtype=torch.int64, device=device) + positions = torch.empty(capacity, dtype=torch.int64, device=device) + prev_slots = torch.empty(capacity, dtype=torch.int64, device=device) + if total_entries > 0: + flat_idx = torch.arange(total_entries, device=device, dtype=torch.int64) + per_req = total_entries // case.bs if case.bs > 0 else 0 + slot_indices[:total_entries] = (flat_idx % max(num_slots - 1, 1)).to( + torch.int64 + ) + positions[:total_entries] = (flat_idx % max(per_req, 1)).to(torch.int64) + is_head = (flat_idx % max(per_req, 1)) == 0 + prev_seq = (flat_idx - 1) % max(num_slots - 1, 1) + prev_slots[:total_entries] = torch.where( + is_head, torch.full_like(flat_idx, -1), prev_seq + ).to(torch.int64) + if capacity > total_entries: + slot_indices[total_entries:] = 0 + positions[total_entries:] = 0 + prev_slots[total_entries:] = -1 + + num_valid = torch.tensor([total_entries], dtype=torch.int32, device=device) + enable = torch.ones(1, dtype=torch.int32, device=device) + expected_input_ids = torch.full((capacity,), -1, dtype=torch.int64, device=device) + plan = VerifyPlan( + verify_slot_indices=slot_indices, + verify_expected_tokens=expected_input_ids, + verify_expected_positions=positions, + verify_prev_slot_indices=prev_slots, + verify_num_valid=num_valid, + enable=enable, + ) + + violation_ring = torch.zeros( + RING_CAPACITY, consts.VIOLATION_FIELDS, dtype=torch.int64, device=device + ) + violation_write_index = torch.zeros(1, dtype=torch.int32, device=device) + slot_run_counter = torch.zeros(1, dtype=torch.int64, device=device) + kernel_run_counter = torch.zeros(1, dtype=torch.int64, device=device) + enable_chain_position_assert = torch.ones(1, dtype=torch.int32, device=device) + + return ( + canary_buf, + plan, + violation_ring, + violation_write_index, + slot_run_counter, + kernel_run_counter, + enable_chain_position_assert, + ) + + +def _build_context( + *, + canary_buf: torch.Tensor, + violation_ring: torch.Tensor, + violation_write_index: torch.Tensor, + slot_run_counter: torch.Tensor, + kernel_run_counter: torch.Tensor, + enable_chain_position_assert: torch.Tensor, + kernel_kind: CanaryLaunchTag, +) -> VerifyOrWriteContext: + return VerifyOrWriteContext( + canary_buf=canary_buf, + kernel_kind=kernel_kind, + violation_ring=violation_ring, + violation_write_index=violation_write_index, + slot_run_counter=slot_run_counter, + kernel_run_counter=kernel_run_counter, + enable_chain_position_assert=enable_chain_position_assert, + ) + + +@triton.testing.perf_report( + triton.testing.Benchmark( + x_names=_X_NAMES, + x_vals=_X_VALS, + line_arg="provider", + line_vals=["canary", "naive"], + line_names=["canary_verify_step", "naive index_copy_"], + styles=[("blue", "-"), ("red", "--")], + ylabel="us", + plot_name="kv-canary-verify-perf", + args={}, + ) +) +def benchmark( + scenario: str, + bs: int, + prefix_len: int, + mode: str, + extend_len: int, + pool_kind: str, + provider: str, +) -> Tuple[float, float, float]: + case = BenchCase( + scenario=scenario, + bs=bs, + prefix_len=prefix_len, + mode=mode, + extend_len=extend_len, + pool_kind=pool_kind, + ) + device = torch.device(DEFAULT_DEVICE) + + if provider == "canary": + ( + canary_buf, + plan, + violation_ring, + violation_write_index, + slot_run_counter, + kernel_run_counter, + enable_chain_position_assert, + ) = _build_verify_inputs(case, device=device) + context = _build_context( + canary_buf=canary_buf, + violation_ring=violation_ring, + violation_write_index=violation_write_index, + slot_run_counter=slot_run_counter, + kernel_run_counter=kernel_run_counter, + enable_chain_position_assert=enable_chain_position_assert, + kernel_kind=CanaryLaunchTag.HEAD_K_FULL, + ) + + def fn() -> None: + violation_write_index.zero_() + launch_canary_verify_kernel( + context=context, + plan=plan, + check_verify_expected_token=True, + ) + + else: + fn = naive_slot_copy_fn(total=_verify_entry_count(case), device=device) + + return run_benchmark(fn) + + +@triton.testing.perf_report( + triton.testing.Benchmark( + x_names=_KERNEL_KIND_X_NAMES, + x_vals=_KERNEL_KIND_X_VALS, + line_arg="provider", + line_vals=["canary"], + line_names=["canary_verify_step"], + styles=[("blue", "-")], + ylabel="us", + plot_name="kv-canary-verify-kernel-kind-perf", + args={}, + ) +) +def benchmark_kernel_kind( + kernel_kind_name: str, + provider: str, +) -> Tuple[float, float, float]: + case = BenchCase( + scenario="kernel_kind", + bs=32, + prefix_len=4096, + mode="extend", + extend_len=128, + pool_kind="full", + ) + device = torch.device(DEFAULT_DEVICE) + + ( + canary_buf, + plan, + violation_ring, + violation_write_index, + slot_run_counter, + kernel_run_counter, + enable_chain_position_assert, + ) = _build_verify_inputs(case, device=device) + kernel_kind = CanaryLaunchTag[kernel_kind_name] + context = _build_context( + canary_buf=canary_buf, + violation_ring=violation_ring, + violation_write_index=violation_write_index, + slot_run_counter=slot_run_counter, + kernel_run_counter=kernel_run_counter, + enable_chain_position_assert=enable_chain_position_assert, + kernel_kind=kernel_kind, + ) + + def fn() -> None: + violation_write_index.zero_() + launch_canary_verify_kernel( + context=context, + plan=plan, + check_verify_expected_token=True, + ) + + return run_benchmark(fn) + + +if __name__ == "__main__": + benchmark.run(print_data=True) + benchmark_kernel_kind.run(print_data=True) diff --git a/python/sglang/jit_kernel/benchmark/kv_canary/utils.py b/python/sglang/jit_kernel/benchmark/kv_canary/utils.py new file mode 100644 index 000000000..c0e69542c --- /dev/null +++ b/python/sglang/jit_kernel/benchmark/kv_canary/utils.py @@ -0,0 +1,269 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Callable + +import torch + +from sglang.jit_kernel.kv_canary.verify import CANARY_SLOT_BYTES + +BS_AXIS: list[int] = [1, 4, 32, 128, 256, 1024] +PREFIX_AXIS: list[int] = [0, 128, 1024, 4096, 10240, 16384] +EXTEND_LEN_AXIS: list[int] = [128, 512, 4096, 16384] +POOL_AXIS: list[str] = ["full", "swa_window_128"] +SWA_WINDOW: int = 128 +RING_CAPACITY: int = 256 +MAX_EXTEND_TOKENS_PER_FORWARD: int = 4096 + + +@dataclass(frozen=True, slots=True, kw_only=True) +class BenchCase: + scenario: str + bs: int + prefix_len: int + mode: str + extend_len: int + pool_kind: str + + @property + def case_id(self) -> str: + return ( + f"{self.scenario}_bs{self.bs}_prefix{self.prefix_len}_{self.mode}{self.extend_len}" + f"_{self.pool_kind}" + ) + + +def _case( + *, + scenario: str, + bs: int, + prefix_len: int, + mode: str, + extend_len: int, + pool_kind: str, +) -> BenchCase: + return BenchCase( + scenario=scenario, + bs=bs, + prefix_len=prefix_len, + mode=mode, + extend_len=extend_len, + pool_kind=pool_kind, + ) + + +def _is_realistic_extend_case(case: BenchCase) -> bool: + if case.mode != "extend": + return True + return case.bs * case.extend_len <= MAX_EXTEND_TOKENS_PER_FORWARD + + +def _dedupe_cases(cases: list[BenchCase]) -> list[BenchCase]: + seen: set[str] = set() + result: list[BenchCase] = [] + + for case in cases: + if case.case_id in seen: + continue + seen.add(case.case_id) + result.append(case) + + return result + + +def build_fast_matrix_cases() -> list[BenchCase]: + return _dedupe_cases( + [ + _case( + scenario="smoke_decode_empty", + bs=1, + prefix_len=0, + mode="decode", + extend_len=1, + pool_kind="full", + ), + _case( + scenario="small_extend_batch", + bs=32, + prefix_len=4096, + mode="extend", + extend_len=128, + pool_kind="full", + ), + _case( + scenario="e2e_decode_steady", + bs=256, + prefix_len=4096, + mode="decode", + extend_len=1, + pool_kind="full", + ), + _case( + scenario="decode_large_batch_short_prefix", + bs=1024, + prefix_len=1024, + mode="decode", + extend_len=1, + pool_kind="full", + ), + _case( + scenario="e2e_prefill_chunk_first", + bs=1, + prefix_len=0, + mode="extend", + extend_len=4096, + pool_kind="full", + ), + _case( + scenario="e2e_prefill_chunk_mid", + bs=1, + prefix_len=8192, + mode="extend", + extend_len=4096, + pool_kind="full", + ), + _case( + scenario="e2e_prefill_chunk_last", + bs=1, + prefix_len=12288, + mode="extend", + extend_len=4096, + pool_kind="full", + ), + _case( + scenario="e2e_decode_tail", + bs=1, + prefix_len=5120, + mode="decode", + extend_len=1, + pool_kind="full", + ), + _case( + scenario="swa_decode_long_prefix", + bs=128, + prefix_len=10240, + mode="decode", + extend_len=1, + pool_kind="swa_window_128", + ), + _case( + scenario="small_extend_single_req", + bs=1, + prefix_len=128, + mode="extend", + extend_len=128, + pool_kind="full", + ), + _case( + scenario="medium_extend_chunk", + bs=4, + prefix_len=1024, + mode="extend", + extend_len=512, + pool_kind="full", + ), + _case( + scenario="decode_mid_batch", + bs=128, + prefix_len=4096, + mode="decode", + extend_len=1, + pool_kind="full", + ), + _case( + scenario="e2e_prefill_chunk_second", + bs=1, + prefix_len=4096, + mode="extend", + extend_len=4096, + pool_kind="full", + ), + _case( + scenario="swa_decode_short_prefix", + bs=256, + prefix_len=128, + mode="decode", + extend_len=1, + pool_kind="swa_window_128", + ), + _case( + scenario="swa_decode_tail", + bs=4, + prefix_len=10240, + mode="decode", + extend_len=1, + pool_kind="swa_window_128", + ), + ] + ) + + +def build_full_matrix_cases() -> list[BenchCase]: + """Full matrix plus targeted e2e points. + + Extend cases are pruned to a maximum token chunk per forward because the scheduler chunks long + prefills; for example, a 4096-token extend is represented as ``bs=1``, not ``bs=32``. + """ + fast = build_fast_matrix_cases() + fast_keys = {c.case_id for c in fast} + full: list[BenchCase] = list(fast) + + for bs in BS_AXIS: + for prefix_len in PREFIX_AXIS: + for pool_kind in POOL_AXIS: + for mode, extend_len in ( + ("decode", 1), + *(("extend", e) for e in EXTEND_LEN_AXIS), + ): + case = _case( + scenario="matrix", + bs=bs, + prefix_len=prefix_len, + mode=mode, + extend_len=extend_len, + pool_kind=pool_kind, + ) + if not _is_realistic_extend_case(case): + continue + if case.case_id in fast_keys: + continue + full.append(case) + + return full + + +def cases_to_x_vals( + cases: list[BenchCase], +) -> list[tuple[str, int, int, str, int, str]]: + return [ + ( + c.scenario, + c.bs, + c.prefix_len, + c.mode, + c.extend_len, + c.pool_kind, + ) + for c in cases + ] + + +def naive_slot_copy_fn(*, total: int, device: torch.device) -> Callable[[], None]: + n_slots = max(total, 1) + payload = torch.zeros(n_slots, CANARY_SLOT_BYTES, dtype=torch.uint8, device=device) + sink = torch.zeros_like(payload) + indices = torch.arange(n_slots, device=device, dtype=torch.int64) % sink.shape[0] + + def baseline() -> None: + sink.index_copy_(0, indices, payload) + + return baseline + + +def naive_cumsum_fn(*, bs: int, device: torch.device) -> Callable[[], None]: + counts = torch.zeros(max(bs, 1), dtype=torch.int32, device=device) + + def baseline() -> None: + torch.cumsum(counts, dim=0) + + return baseline diff --git a/python/sglang/jit_kernel/csrc/kv_canary/canary_common.cuh b/python/sglang/jit_kernel/csrc/kv_canary/canary_common.cuh new file mode 100644 index 000000000..6de0c44fa --- /dev/null +++ b/python/sglang/jit_kernel/csrc/kv_canary/canary_common.cuh @@ -0,0 +1,78 @@ +#pragma once + +#include // For SGL_DEVICE + +#include "consts.cuh" +#include + +namespace canary { + +SGL_DEVICE uint64_t splitmix64(uint64_t x) { + x = (x ^ (x >> 30)) * 0xBF58476D1CE4E5B9ULL; + x = (x ^ (x >> 27)) * 0x94D049BB133111EBULL; + return x ^ (x >> 31); +} + +SGL_DEVICE uint64_t splitmix64_mix3(uint64_t a, uint64_t b, uint64_t c) { + uint64_t h = splitmix64(a); + h = splitmix64(h ^ b); + h = splitmix64(h ^ c); + return h; +} + +struct ViolationSink { + int64_t* __restrict__ ring; + int32_t* __restrict__ write_index; + int32_t ring_capacity; + int32_t kernel_kind; +}; + +struct ViolationRow { + int64_t slot_idx; + int64_t position; + int64_t stored_token; + int64_t expected_token; + int64_t stored_chain_hash; + int64_t expected_aux; + int64_t fail_reason_bits; +}; + +SGL_DEVICE void record_violation(const ViolationSink& sink, const ViolationRow& row) { + const int32_t seq = atomicAdd(sink.write_index, 1); + if (seq < sink.ring_capacity) { + int64_t* dst = sink.ring + static_cast(seq) * kViolationFields; + dst[kViolationFieldKernelKind] = static_cast(sink.kernel_kind); + dst[kViolationFieldSlotIdx] = row.slot_idx; + dst[kViolationFieldPosition] = row.position; + dst[kViolationFieldStoredToken] = row.stored_token; + dst[kViolationFieldExpectedToken] = row.expected_token; + dst[kViolationFieldStoredChainHash] = row.stored_chain_hash; + dst[kViolationFieldExpectedAux] = row.expected_aux; + dst[kViolationFieldFailReasonBits] = row.fail_reason_bits; + __threadfence_system(); + } +} + +SGL_DEVICE int64_t canary_load_field(const uint8_t* buf, int64_t slot_idx, int64_t slot_stride_bytes, int field) { + const int64_t* p = reinterpret_cast(buf + slot_idx * slot_stride_bytes); + return p[field]; +} + +SGL_DEVICE void +canary_store_field(uint8_t* buf, int64_t slot_idx, int64_t slot_stride_bytes, int field, int64_t value) { + int64_t* p = reinterpret_cast(buf + slot_idx * slot_stride_bytes); + p[field] = value; +} + +SGL_DEVICE uint64_t compute_slot_hash(const uint8_t* canary_buf, int64_t slot_stride_bytes, int64_t source_slot_idx) { + if (source_slot_idx < 0) { + return splitmix64(kCanaryChainAnchor); + } + const int64_t token = canary_load_field(canary_buf, source_slot_idx, slot_stride_bytes, kCanaryFieldToken); + const int64_t position = canary_load_field(canary_buf, source_slot_idx, slot_stride_bytes, kCanaryFieldPosition); + const int64_t prev_hash = canary_load_field(canary_buf, source_slot_idx, slot_stride_bytes, kCanaryFieldPrevHash); + return splitmix64_mix3( + static_cast(prev_hash), static_cast(token), static_cast(position)); +} + +} // namespace canary diff --git a/python/sglang/jit_kernel/csrc/kv_canary/canary_verify.cuh b/python/sglang/jit_kernel/csrc/kv_canary/canary_verify.cuh new file mode 100644 index 000000000..339982543 --- /dev/null +++ b/python/sglang/jit_kernel/csrc/kv_canary/canary_verify.cuh @@ -0,0 +1,208 @@ +#pragma once + +#include // For TensorMatcher, SymbolicSize, SymbolicDevice +#include // For div_ceil, RuntimeCheck + +#include // For LaunchKernel, SGL_DEVICE + +#include +#include + +#include "canary_common.cuh" +#include + +namespace canary { + +namespace { + +constexpr uint32_t kVerifyBlockSize = 512; +constexpr uint32_t kPersistentBlocks = 64; + +struct VerifyKernelParams { + // Canary buffer this launch verifies. Read-only. + const uint8_t* canary_buf; + int64_t slot_stride_bytes; + + // Plan tensors. + const int64_t* verify_slot_indices; + const int64_t* verify_expected_tokens; + const int64_t* verify_expected_positions; + const int64_t* verify_prev_slot_indices; + const int32_t* verify_num_valid; + const int32_t* verify_enable; + int32_t verify_capacity; + + // Violation sink (ring + write_index + capacity + kernel_kind bundled in canary_common.cuh). + ViolationSink violation_sink; + + // Health counters. + int64_t* slot_run_counter; + int64_t* kernel_run_counter; +}; + +template +__global__ void canary_verify_kernel(const VerifyKernelParams __grid_constant__ p) { + const uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x; + const uint32_t stride = gridDim.x * blockDim.x; + + if (tid == 0) { + atomicAdd(reinterpret_cast(p.kernel_run_counter), 1ULL); + } + + if (*p.verify_enable == 0) { + return; + } + + const int32_t active = min(*p.verify_num_valid, p.verify_capacity); + + uint32_t local_active_count = 0; + for (uint32_t entry_idx = tid; entry_idx < static_cast(active); entry_idx += stride) { + ++local_active_count; + + const int64_t slot_idx = p.verify_slot_indices[entry_idx]; + const int64_t expected_position = p.verify_expected_positions[entry_idx]; + const int64_t prev_slot_idx = p.verify_prev_slot_indices[entry_idx]; + int64_t expected_input_id = -1; + if constexpr (CHECK_VERIFY_EXPECTED_TOKEN) { + expected_input_id = p.verify_expected_tokens[entry_idx]; + } + + if (slot_idx == kTokenToKvSlotPadding) { + continue; + } + + const int64_t stored_token = canary_load_field(p.canary_buf, slot_idx, p.slot_stride_bytes, kCanaryFieldToken); + const int64_t stored_position = + canary_load_field(p.canary_buf, slot_idx, p.slot_stride_bytes, kCanaryFieldPosition); + const int64_t stored_chain_hash = + canary_load_field(p.canary_buf, slot_idx, p.slot_stride_bytes, kCanaryFieldPrevHash); + // field[3] (kCanaryFieldRealKvHash) is always 0 in the naive canary; not loaded or checked. + + const bool prev_reachable = (prev_slot_idx != kTokenToKvSlotPadding); + const int64_t expected_chain_hash = + prev_reachable ? static_cast(compute_slot_hash(p.canary_buf, p.slot_stride_bytes, prev_slot_idx)) + : stored_chain_hash; + + FailReason fail_reason_bits{}; + if (prev_reachable && stored_chain_hash != expected_chain_hash) { + fail_reason_bits |= FailReason::kVerifyChainHashMismatch; + } + if constexpr (CHECK_VERIFY_EXPECTED_TOKEN) { + if (expected_input_id != -1 && stored_token != expected_input_id) { + fail_reason_bits |= FailReason::kVerifyTokenMismatch; + } + } + if (stored_position != expected_position) { + fail_reason_bits |= FailReason::kVerifyPositionMismatch; + } + + if (fail_reason_bits != FailReason{}) { + record_violation( + p.violation_sink, + ViolationRow{ + /* slot_idx = */ slot_idx, + /* position = */ stored_position, + /* stored_token = */ stored_token, + /* expected_token = */ expected_input_id, + /* stored_chain_hash = */ stored_chain_hash, + /* expected_aux = */ expected_chain_hash, + /* fail_reason_bits = */ static_cast(fail_reason_bits), + }); + } + } + + uint32_t warp_active_count = local_active_count; + for (int offset = 16; offset > 0; offset >>= 1) { + warp_active_count += __shfl_down_sync(0xFFFFFFFFu, warp_active_count, offset); + } + if ((threadIdx.x & 31u) == 0u && warp_active_count != 0u) { + atomicAdd( + reinterpret_cast(p.slot_run_counter), static_cast(warp_active_count)); + } +} + +} // namespace + +// API source of truth: docstring of canary_verify_step in python/sglang/jit_kernel/kv_canary/verify.py. +template +struct CanaryVerifyKernel { + static void + run(tvm::ffi::TensorView canary_buf, + tvm::ffi::TensorView verify_slot_indices, + tvm::ffi::TensorView verify_expected_tokens, + tvm::ffi::TensorView verify_expected_positions, + tvm::ffi::TensorView verify_prev_slot_indices, + tvm::ffi::TensorView verify_num_valid, + tvm::ffi::TensorView verify_enable, + int64_t kernel_kind, + tvm::ffi::TensorView violation_ring, + tvm::ffi::TensorView violation_write_index, + tvm::ffi::TensorView slot_run_counter, + tvm::ffi::TensorView kernel_run_counter) { + using namespace host; + + SymbolicSize N_slots = {"num_canary_slots"}; + SymbolicSize N_stride = {"slot_stride_bytes"}; + SymbolicSize N_verify = {"verify_capacity"}; + SymbolicDevice device_; + device_.set_options(); + + TensorMatcher({N_slots, N_stride}).with_dtype().with_device(device_).verify(canary_buf); + + TensorMatcher({N_verify}) + .with_dtype() + .with_device(device_) + .verify(verify_slot_indices) + .verify(verify_expected_tokens) + .verify(verify_expected_positions) + .verify(verify_prev_slot_indices); + TensorMatcher({1}).with_dtype().with_device(device_).verify(verify_num_valid); + TensorMatcher({1}).with_dtype().with_device(device_).verify(verify_enable); + + TensorMatcher({1}).with_dtype().with_device(device_).verify(violation_write_index); + SymbolicSize N_ring = {"ring_capacity"}; + TensorMatcher({N_ring, static_cast(kViolationFields)}) + .with_dtype() + .with_device(device_) + .verify(violation_ring); + TensorMatcher({1}) + .with_dtype() + .with_device(device_) + .verify(slot_run_counter) + .verify(kernel_run_counter); + + const int64_t slot_stride_bytes = N_stride.unwrap(); + const int32_t verify_capacity = static_cast(N_verify.unwrap()); + const int32_t ring_capacity = static_cast(N_ring.unwrap()); + const DLDevice device = device_.unwrap(); + + RuntimeCheck( + slot_stride_bytes >= static_cast(kCanaryFieldsPerSlot * sizeof(int64_t)), + "canary_verify: slot_stride_bytes must hold at least ", + static_cast(kCanaryFieldsPerSlot * sizeof(int64_t)), + " bytes per slot, got ", + slot_stride_bytes); + + VerifyKernelParams p{}; + p.canary_buf = static_cast(canary_buf.data_ptr()); + p.slot_stride_bytes = slot_stride_bytes; + p.verify_slot_indices = static_cast(verify_slot_indices.data_ptr()); + p.verify_expected_tokens = static_cast(verify_expected_tokens.data_ptr()); + p.verify_expected_positions = static_cast(verify_expected_positions.data_ptr()); + p.verify_prev_slot_indices = static_cast(verify_prev_slot_indices.data_ptr()); + p.verify_num_valid = static_cast(verify_num_valid.data_ptr()); + p.verify_enable = static_cast(verify_enable.data_ptr()); + p.verify_capacity = verify_capacity; + p.violation_sink.ring = static_cast(violation_ring.data_ptr()); + p.violation_sink.write_index = static_cast(violation_write_index.data_ptr()); + p.violation_sink.ring_capacity = ring_capacity; + p.violation_sink.kernel_kind = static_cast(kernel_kind); + p.slot_run_counter = static_cast(slot_run_counter.data_ptr()); + p.kernel_run_counter = static_cast(kernel_run_counter.data_ptr()); + + const uint32_t grid = kPersistentBlocks; + LaunchKernel(grid, kVerifyBlockSize, device)(canary_verify_kernel, p); + } +}; + +} // namespace canary diff --git a/python/sglang/jit_kernel/csrc/kv_canary/consts.cuh b/python/sglang/jit_kernel/csrc/kv_canary/consts.cuh new file mode 100644 index 000000000..d36577c57 --- /dev/null +++ b/python/sglang/jit_kernel/csrc/kv_canary/consts.cuh @@ -0,0 +1,49 @@ +#pragma once + +#include + +namespace canary { + +constexpr uint64_t kCanaryChainAnchor = 0xC0FFEE1234567890ULL; + +// Mirrors SGLang's TokenToKVPoolAllocator contract: token-to-KV slot 0 is reserved for padded-token dummy +// writes. Since req_to_token stores token-to-KV slot ids and is zero-initialized, canary slot 0 is skipped +// instead of treating unfilled entries as real KV slots. +constexpr int64_t kTokenToKvSlotPadding = 0; +constexpr int64_t kReqPoolIdxPadding = 0; + +constexpr int kCanaryFieldsPerSlot = 4; +constexpr int kCanaryFieldToken = 0; +constexpr int kCanaryFieldPosition = 1; +constexpr int kCanaryFieldPrevHash = 2; +constexpr int kCanaryFieldRealKvHash = 3; + +constexpr int kViolationFields = 8; +constexpr int kViolationFieldKernelKind = 0; +constexpr int kViolationFieldSlotIdx = 1; +constexpr int kViolationFieldPosition = 2; +constexpr int kViolationFieldStoredToken = 3; +constexpr int kViolationFieldExpectedToken = 4; +constexpr int kViolationFieldStoredChainHash = 5; +constexpr int kViolationFieldExpectedAux = 6; +constexpr int kViolationFieldFailReasonBits = 7; + +enum class FailReason : int64_t { + kVerifyChainHashMismatch = 1LL << 0, + kVerifyPositionMismatch = 1LL << 1, + kVerifyRealKvHashMismatch = 1LL << 2, + kWriteTokenMismatch = 1LL << 3, + kWritePositionMismatch = 1LL << 4, + kVerifyTokenMismatch = 1LL << 5, +}; + +constexpr FailReason operator|(FailReason a, FailReason b) { + return static_cast(static_cast(a) | static_cast(b)); +} + +constexpr FailReason& operator|=(FailReason& a, FailReason b) { + a = a | b; + return a; +} + +} // namespace canary diff --git a/python/sglang/jit_kernel/kv_canary/__init__.py b/python/sglang/jit_kernel/kv_canary/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/sglang/jit_kernel/kv_canary/consts.py b/python/sglang/jit_kernel/kv_canary/consts.py new file mode 100644 index 000000000..cde98eb8a --- /dev/null +++ b/python/sglang/jit_kernel/kv_canary/consts.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +from enum import IntFlag +from typing import Final + +CANARY_CHAIN_ANCHOR: Final[int] = 0xC0FFEE1234567890 + +# Mirrors SGLang's ReqToTokenPool contract: req_pool_idx 0 is the CUDA-graph padding row, while real +# request rows start at 1. +REQ_POOL_IDX_PADDING: Final[int] = 0 + +# Mirrors SGLang's TokenToKVPoolAllocator contract: token-to-KV slot 0 is reserved for padded-token dummy +# writes. Since req_to_token stores token-to-KV slot ids and is zero-initialized, canary slot 0 is skipped +# instead of treating unfilled entries as real KV slots. +TOKEN_TO_KV_SLOT_PADDING: Final[int] = 0 + +CANARY_FIELDS_PER_SLOT: Final[int] = 4 +CANARY_FIELD_TOKEN: Final[int] = 0 +CANARY_FIELD_POSITION: Final[int] = 1 +CANARY_FIELD_PREV_HASH: Final[int] = 2 +CANARY_FIELD_REAL_KV_HASH: Final[int] = 3 + +VIOLATION_FIELDS: Final[int] = 8 +VIOLATION_FIELD_KERNEL_KIND: Final[int] = 0 +VIOLATION_FIELD_SLOT_IDX: Final[int] = 1 +VIOLATION_FIELD_POSITION: Final[int] = 2 +VIOLATION_FIELD_STORED_TOKEN: Final[int] = 3 +VIOLATION_FIELD_EXPECTED_TOKEN: Final[int] = 4 +VIOLATION_FIELD_STORED_CHAIN_HASH: Final[int] = 5 +VIOLATION_FIELD_EXPECTED_AUX: Final[int] = 6 +VIOLATION_FIELD_FAIL_REASON_BITS: Final[int] = 7 + + +class FailReason(IntFlag): + VERIFY_CHAIN_HASH_MISMATCH = 1 << 0 + VERIFY_POSITION_MISMATCH = 1 << 1 + VERIFY_REAL_KV_HASH_MISMATCH = 1 << 2 + WRITE_TOKEN_MISMATCH = 1 << 3 + WRITE_POSITION_MISMATCH = 1 << 4 + VERIFY_TOKEN_MISMATCH = 1 << 5 + + +_U64_MASK: int = (1 << 64) - 1 + + +def splitmix64(value: int) -> int: + x = value & _U64_MASK + x = ((x ^ (x >> 30)) * 0xBF58476D1CE4E5B9) & _U64_MASK + x = ((x ^ (x >> 27)) * 0x94D049BB133111EB) & _U64_MASK + return (x ^ (x >> 31)) & _U64_MASK + + +def splitmix64_mix3(a: int, b: int, c: int) -> int: + h = splitmix64(a & _U64_MASK) + h = splitmix64(h ^ (b & _U64_MASK)) + h = splitmix64(h ^ (c & _U64_MASK)) + return h diff --git a/python/sglang/jit_kernel/kv_canary/verify.py b/python/sglang/jit_kernel/kv_canary/verify.py new file mode 100644 index 000000000..978665440 --- /dev/null +++ b/python/sglang/jit_kernel/kv_canary/verify.py @@ -0,0 +1,254 @@ +from __future__ import annotations + +from dataclasses import dataclass +from enum import IntEnum +from typing import TYPE_CHECKING, Final + +import torch + +from sglang.jit_kernel.kv_canary import consts +from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args + +if TYPE_CHECKING: + from tvm_ffi.module import Module + +# Bytes per canary slot = CANARY_FIELDS_PER_SLOT * 8. +CANARY_SLOT_BYTES: Final[int] = consts.CANARY_FIELDS_PER_SLOT * 8 + + +class CanaryLaunchTag(IntEnum): + """Unique tag per (head | tail) × (K | V) × (FULL | SWA) launch.""" + + HEAD_K_FULL = 0 + HEAD_V_FULL = 1 + TAIL_K_FULL = 2 + TAIL_V_FULL = 3 + HEAD_K_SWA = 4 + HEAD_V_SWA = 5 + TAIL_K_SWA = 6 + TAIL_V_SWA = 7 + + +def _assert_contiguous(tensor: torch.Tensor, name: str) -> None: + if not tensor.is_contiguous(): + raise ValueError(f"kv-canary: {name} must be contiguous") + + +@dataclass(frozen=True, slots=True, kw_only=True) +class VerifyOrWriteContext: + """Shared launch context for canary verify/write kernels. + + Fields: + canary_buf: Canary buffer this launch verifies or writes, shape [num_slots, slot_stride_bytes], uint8. + slot_stride_bytes is read from canary_buf.shape[1]. + kernel_kind: CanaryLaunchTag identifying which launch fired. Stamped (as int) into every violation row + so host can attribute a violation back to its source launch. + violation_ring: Global append-only sink, shape [ring_capacity, VIOLATION_FIELDS], int64. Shared across + all canary launches; fill-once. + violation_write_index: Global monotonic violation counter, shape [1], int32. + slot_run_counter: Health counter, shape [1], int64. Verify increments by active entries processed; + write increments by write entries processed. + kernel_run_counter: Health counter, shape [1], int64. Incremented by 1 per call. + enable_chain_position_assert: int32 [1] device flag gating the write kernel's chain-step + write_position assert. 0 during warmup / cuda-graph capture; flipped to 1 in + CanaryManager.mark_init_finished(). + """ + + canary_buf: torch.Tensor + kernel_kind: CanaryLaunchTag + violation_ring: torch.Tensor + violation_write_index: torch.Tensor + slot_run_counter: torch.Tensor + kernel_run_counter: torch.Tensor + enable_chain_position_assert: torch.Tensor + + +@dataclass(frozen=True, slots=True, kw_only=True) +class VerifyPlan: + """Flat verify entries consumed by launch_canary_verify_kernel. + + Each row is a self-contained (slot_idx, position, prev_slot_idx) triple, so the verify kernel makes no + assumption about the entry's source. prev_slot_idx == -1 flags a chain-seed entry (kernel + anchors on the hardcoded CANARY_CHAIN_ANCHOR constant instead of reading a predecessor). + + Sized to a cuda-graph-captured capacity; active prefix is verify_num_valid[0]. Padding tail entries are + unspecified — kernel skips tid >= verify_num_valid[0]. + + Fields: + verify_slot_indices: Canary slot index per entry, shape [verify_capacity], int64. Already SWA-translated + for the SWA group. + verify_expected_tokens: Source-of-truth token id per entry, shape [verify_capacity], int64. + The plan-side entries kernel gathers from + ``CanaryDeviceState.req_to_verify_expected_tokens[rp, position + kv_token_id_vs_position_offset]``; + entries that fall outside the pool's row (e.g. EAGLE draft's last slot rotating in a bonus + token, or padding beyond the per-req length) get the ``-1`` sentinel. The verify kernel + compares against the stored canary token and skips when this is ``-1``. + verify_expected_positions: Expected sequence position per entry, shape [verify_capacity], int64. + verify_prev_slot_indices: Chain predecessor slot per entry, shape [verify_capacity], int64. -1 = chain + head (anchor on CANARY_CHAIN_ANCHOR). Explicit (not derived from verify_slot_indices[i-1]) + because chain heads, SWA window starts, cross-req boundaries, and radix-orphan extras break the + "predecessor == previous array entry" assumption. + verify_num_valid: Active entry count, shape [1], int32. Clamped by the plan kernel to + min(total_requested, verify_capacity) so the verify kernel grid never reads past the buffer. + enable: Run-this-step flag, shape [1], int32. 1 = verify kernel runs as usual; 0 = the plan kernel + detected overflow (requested > verify_capacity) and the entire verify launch is skipped this step. + Allocated as 1 by default; the plan kernel rewrites it every step. + """ + + verify_slot_indices: torch.Tensor + verify_expected_tokens: torch.Tensor + verify_expected_positions: torch.Tensor + verify_prev_slot_indices: torch.Tensor + verify_num_valid: torch.Tensor + enable: torch.Tensor + + @classmethod + def allocate(cls, *, verify_capacity: int, device: torch.device) -> "VerifyPlan": + if verify_capacity <= 0: + raise ValueError( + f"kv-canary: VerifyPlan verify_capacity must be positive, got {verify_capacity}" + ) + return cls( + verify_slot_indices=torch.empty( + verify_capacity, dtype=torch.int64, device=device + ), + verify_expected_tokens=torch.empty( + verify_capacity, dtype=torch.int64, device=device + ), + verify_expected_positions=torch.empty( + verify_capacity, dtype=torch.int64, device=device + ), + verify_prev_slot_indices=torch.empty( + verify_capacity, dtype=torch.int64, device=device + ), + verify_num_valid=torch.empty(1, dtype=torch.int32, device=device), + # enable defaults to 1 ("run verify") so test helpers that build a VerifyPlan + # directly (no plan kernel) don't have to remember to set it. Plan kernel always + # overwrites this so the default is observable only when no plan kernel runs. + enable=torch.ones(1, dtype=torch.int32, device=device), + ) + + def zero_for_testing_(self) -> "VerifyPlan": + """WARN: ONLY use it when testing plan kernel. Do not use it when testing verify or + write kernel to avoid hiding bugs.""" + self.verify_slot_indices.zero_() + # Test helpers expect the "skip token check" sentinel after zero-out, matching + # the verify-kernel contract. + self.verify_expected_tokens.fill_(-1) + self.verify_expected_positions.zero_() + self.verify_prev_slot_indices.zero_() + self.verify_num_valid.zero_() + self.enable.zero_() + return self + + +def launch_canary_verify_kernel( + *, + context: VerifyOrWriteContext, + plan: VerifyPlan, + check_verify_expected_token: bool, +) -> None: + """Verify one canary buffer against a VerifyPlan. + + A fixed persistent grid of `kPersistentBlocks * kVerifyBlockSize` CUDA threads grid-strides over active + verify entries. Each thread reads the slot's 4 stored int64 fields (token_id, position, prev_hash, + real_kv_hash), recomputes the expected prev_hash from the predecessor slot (or from + splitmix64(CANARY_CHAIN_ANCHOR) for chain heads, signaled by prev_slot_idx == -1), and atomically appends + any mismatch (chain hash / position) to violation_ring. Read-only on canary_buf. + + Canary slot layout: each slot is canary_buf.shape[1] bytes holding 4 int64 fields (token_id, position, + prev_hash, real_kv_hash). The real_kv_hash field is always 0. Chain link: next.prev_hash == + splitmix64_mix3(this.prev_hash, this.token_id, + this.position), where splitmix64_mix3 folds each input into a running accumulator + via ``acc = splitmix64(acc ^ next)`` starting from ``splitmix64(prev_hash)``. ``real_kv_hash`` is NOT + folded into the chain (see ``compute_slot_hash`` rationale: keeps chain content-only and immune to + legitimate radix prefix folding). Chain head anchors on + splitmix64(CANARY_CHAIN_ANCHOR), where CANARY_CHAIN_ANCHOR is a hardcoded module-level constant (no + runtime seed parameter — the canary is for bug detection, not adversarial security, so a fixed anchor + is sufficient). + + Args: + context: Shared verify/write launch context, including canary buffer, launch tag, violation sink, + and health counters. + plan: Pre-allocated VerifyPlan; addresses baked into cuda-graph capture. + + Token-to-KV slot 0 is unconditionally skipped by the verify kernel: SGLang's TokenToKVPoolAllocator + reserves it for padded-token dummy writes, and zero-initialized req_to_token entries therefore point to + a non-real KV slot. Canary-attached pools mirror that contract by reserving canary slot 0. + + Implementation: + - CUDA __global__ `canary_verify_kernel`: fixed 1-D grid `(kPersistentBlocks=64, 1, 1)` blocks × + `(kVerifyBlockSize=512, 1, 1)` threads (= 32768 threads total). Each thread grid-strides over + verify entries `entry_idx ∈ [tid, tid + grid_threads, ...)` until + `min(plan.verify_num_valid[0], verify_capacity)`. + - Per thread, gather: + (a) self_slot fields: 4 separate ``canary_load_field`` int64 loads from + canary_buf[plan.verify_slot_indices[tid]] for (token, position, prev_hash, real_kv_hash). + (b) expected_prev_hash = compute_slot_hash(canary_buf, slot_stride_bytes, prev_slot_idx), which + folds only (token, position, prev_hash) from canary_buf[plan.verify_prev_slot_indices[tid]]; + prev_slot_idx == -1 anchors at splitmix64(CANARY_CHAIN_ANCHOR). + - Compare expected vs stored (chain hash, position) and accumulate fail_reason bits; if + non-zero → record_violation(). + - record_violation(): idx = atomicAdd(violation_write_index, 1); if idx < ring_capacity, atomic-write + the 8 int64 fields to violation_ring[idx] (kernel_kind, slot_idx, position, stored vs expected + fields, fail_reason). + - Counters: each thread maintains a local count of active entries it processed, warp-reduces via + ``__shfl_down_sync`` (offsets 16..1), then the warp leader (lane 0) does a single atomicAdd of the + warp's summed count into slot_run_counter. kernel_run_counter += 1: single thread (tid == 0) does an + atomicAdd once per launch. + + Calling contract: + - Pure side-effect; never raises. Host polls violation_write_index[0] > 0 for is_errored and + violation_ring[0] for the first violation. + - kernel_run_counter is bumped every call (canary-ran health signal). + - Safe in cuda-graph capture; caller refills plan in-place before replay. + + Pinned by torch reference + :func:`sglang.jit_kernel.kv_canary.verify_ref.launch_canary_verify_kernel_torch_reference`; CUDA must match + byte-for-byte. + """ + canary_buf = context.canary_buf + + _assert_contiguous(canary_buf, "canary_buf") + _assert_contiguous(plan.verify_slot_indices, "plan.verify_slot_indices") + _assert_contiguous(plan.verify_expected_tokens, "plan.verify_expected_tokens") + _assert_contiguous(plan.verify_expected_positions, "plan.verify_expected_positions") + _assert_contiguous(plan.verify_prev_slot_indices, "plan.verify_prev_slot_indices") + _assert_contiguous(plan.verify_num_valid, "plan.verify_num_valid") + _assert_contiguous(plan.enable, "plan.enable") + _assert_contiguous(context.violation_ring, "violation_ring") + _assert_contiguous(context.violation_write_index, "violation_write_index") + _assert_contiguous(context.slot_run_counter, "slot_run_counter") + _assert_contiguous(context.kernel_run_counter, "kernel_run_counter") + + module = _jit_canary_verify_module(check_verify_expected_token) + module.canary_verify_step_cuda( + canary_buf, + plan.verify_slot_indices, + plan.verify_expected_tokens, + plan.verify_expected_positions, + plan.verify_prev_slot_indices, + plan.verify_num_valid, + plan.enable, + int(context.kernel_kind), + context.violation_ring, + context.violation_write_index, + context.slot_run_counter, + context.kernel_run_counter, + ) + + +@cache_once +def _jit_canary_verify_module(check_verify_expected_token: bool) -> "Module": + args = make_cpp_args(check_verify_expected_token) + return load_jit( + "kv_canary_verify", + *args, + cuda_files=["kv_canary/canary_verify.cuh"], + cuda_wrappers=[ + ( + "canary_verify_step_cuda", + f"canary::CanaryVerifyKernel<{args}>::run", + ), + ], + ) diff --git a/python/sglang/jit_kernel/kv_canary/verify_ref.py b/python/sglang/jit_kernel/kv_canary/verify_ref.py new file mode 100644 index 000000000..50c054e74 --- /dev/null +++ b/python/sglang/jit_kernel/kv_canary/verify_ref.py @@ -0,0 +1,182 @@ +from __future__ import annotations + +import torch + +from sglang.jit_kernel.kv_canary import consts +from sglang.jit_kernel.kv_canary.consts import splitmix64, splitmix64_mix3 +from sglang.jit_kernel.kv_canary.verify import ( + VerifyOrWriteContext, + VerifyPlan, +) + +_U64_MASK: int = (1 << 64) - 1 +_I64_SIGN_BIT: int = 1 << 63 + + +def launch_canary_verify_kernel_torch_reference( + *, + context: VerifyOrWriteContext, + plan: VerifyPlan, + check_verify_expected_token: bool, +) -> None: + canary_buf = context.canary_buf + kernel_kind = context.kernel_kind + violation_ring = context.violation_ring + violation_write_index = context.violation_write_index + slot_run_counter = context.slot_run_counter + kernel_run_counter = context.kernel_run_counter + + work_device = torch.device("cpu") + + kernel_run_counter.add_(1) + + enable = int(plan.enable.detach().to("cpu").item()) + if enable == 0: + return + + num_valid = int( + plan.verify_slot_indices.new_empty(()).copy_(plan.verify_num_valid[0]).item() + ) + capacity = int(plan.verify_slot_indices.shape[0]) + active = max(0, min(num_valid, capacity)) + if active <= 0: + return + + slot_indices_host = plan.verify_slot_indices[:active].to( + device=work_device, dtype=torch.int64 + ) + if check_verify_expected_token: + expected_input_ids_host = plan.verify_expected_tokens[:active].to( + device=work_device, dtype=torch.int64 + ) + else: + expected_input_ids_host = torch.full( + (active,), -1, dtype=torch.int64, device=work_device + ) + expected_positions_host = plan.verify_expected_positions[:active].to( + device=work_device, dtype=torch.int64 + ) + prev_slot_indices_host = plan.verify_prev_slot_indices[:active].to( + device=work_device, dtype=torch.int64 + ) + + slot_run_counter.add_(active) + + kept_slots: list[int] = [] + kept_expected_positions: list[int] = [] + kept_expected_input_ids: list[int] = [] + kept_prev_slots: list[int] = [] + for k in range(active): + s = int(slot_indices_host[k].item()) + # Skip SGLang's padded-token dummy KV slot so unfilled req_to_token entries (zero-initialized) do not + # produce spurious chain_hash / position violations. + if s != consts.TOKEN_TO_KV_SLOT_PADDING: + kept_slots.append(s) + kept_expected_positions.append(int(expected_positions_host[k].item())) + kept_expected_input_ids.append(int(expected_input_ids_host[k].item())) + kept_prev_slots.append(int(prev_slot_indices_host[k].item())) + active = len(kept_slots) + if active <= 0: + return + slot_indices_list: list[int] = kept_slots + expected_positions_list: list[int] = kept_expected_positions + expected_input_ids_list: list[int] = kept_expected_input_ids + prev_slot_indices_list: list[int] = kept_prev_slots + + buf_i64 = canary_buf.detach().to(device=work_device).contiguous().view(torch.int64) + slot_stride_i64 = int(buf_i64.shape[1]) + if slot_stride_i64 < 4: + raise ValueError( + f"kv-canary: canary_buf slot stride must hold at least 4 int64 fields, got {slot_stride_i64}" + ) + + violation_rows: list[list[int]] = [] + + for k in range(active): + slot_idx = slot_indices_list[k] + expected_position = expected_positions_list[k] + expected_input_id = expected_input_ids_list[k] + prev_slot = prev_slot_indices_list[k] + + stored_token = int(buf_i64[slot_idx, consts.CANARY_FIELD_TOKEN].item()) + stored_position = int(buf_i64[slot_idx, consts.CANARY_FIELD_POSITION].item()) + stored_chain_hash = int(buf_i64[slot_idx, consts.CANARY_FIELD_PREV_HASH].item()) + stored_real_kv_hash = int( + buf_i64[slot_idx, consts.CANARY_FIELD_REAL_KV_HASH].item() + ) + + prev_reachable = prev_slot != consts.TOKEN_TO_KV_SLOT_PADDING + if prev_reachable: + expected_chain_hash = _to_signed_int64( + compute_slot_hash(buf_i64, prev_slot) + ) + else: + expected_chain_hash = stored_chain_hash + + # real_kv_hash is always 0 in the naive canary; field[3] is stored as 0 by write. + expected_real_kv_hash = 0 + + fail_reason = consts.FailReason(0) + if prev_reachable and stored_chain_hash != expected_chain_hash: + fail_reason |= consts.FailReason.VERIFY_CHAIN_HASH_MISMATCH + if check_verify_expected_token: + if expected_input_id != -1 and stored_token != expected_input_id: + fail_reason |= consts.FailReason.VERIFY_TOKEN_MISMATCH + if stored_position != expected_position: + fail_reason |= consts.FailReason.VERIFY_POSITION_MISMATCH + if stored_real_kv_hash != expected_real_kv_hash: + fail_reason |= consts.FailReason.VERIFY_REAL_KV_HASH_MISMATCH + + if fail_reason != consts.FailReason(0): + row = [0] * consts.VIOLATION_FIELDS + row[consts.VIOLATION_FIELD_KERNEL_KIND] = int(kernel_kind) + row[consts.VIOLATION_FIELD_SLOT_IDX] = slot_idx + row[consts.VIOLATION_FIELD_POSITION] = stored_position + row[consts.VIOLATION_FIELD_STORED_TOKEN] = stored_token + row[consts.VIOLATION_FIELD_EXPECTED_TOKEN] = expected_input_id + row[consts.VIOLATION_FIELD_STORED_CHAIN_HASH] = stored_chain_hash + row[consts.VIOLATION_FIELD_EXPECTED_AUX] = expected_chain_hash + row[consts.VIOLATION_FIELD_FAIL_REASON_BITS] = int(fail_reason) + violation_rows.append(row) + + if len(violation_rows) == 0: + return + + num_new_violations = len(violation_rows) + base_idx = int( + violation_write_index.new_empty(()).copy_(violation_write_index[0]).item() + ) + ring_capacity = int(violation_ring.shape[0]) + + new_rows = torch.zeros( + (num_new_violations, consts.VIOLATION_FIELDS), dtype=torch.int64 + ) + for v, row in enumerate(violation_rows): + for f in range(consts.VIOLATION_FIELDS): + new_rows[v, f] = row[f] + + write_count_in_ring = max(0, min(num_new_violations, ring_capacity - base_idx)) + if write_count_in_ring > 0: + ring_host = violation_ring.detach().to(device=work_device) + ring_host[base_idx : base_idx + write_count_in_ring, :] = new_rows[ + :write_count_in_ring, : + ] + violation_ring.copy_(ring_host.to(violation_ring.device)) + + violation_write_index[0] = violation_write_index[0] + num_new_violations + + +def _to_signed_int64(value: int) -> int: + value &= _U64_MASK + if value >= _I64_SIGN_BIT: + value -= 1 << 64 + return value + + +def compute_slot_hash(buf_i64: torch.Tensor, source_slot_idx: int) -> int: + if source_slot_idx < 0: + return splitmix64(consts.CANARY_CHAIN_ANCHOR) + token = int(buf_i64[source_slot_idx, consts.CANARY_FIELD_TOKEN].item()) + position = int(buf_i64[source_slot_idx, consts.CANARY_FIELD_POSITION].item()) + prev_hash = int(buf_i64[source_slot_idx, consts.CANARY_FIELD_PREV_HASH].item()) + return splitmix64_mix3(prev_hash, token, position) diff --git a/python/sglang/jit_kernel/tests/kv_canary/__init__.py b/python/sglang/jit_kernel/tests/kv_canary/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/sglang/jit_kernel/tests/kv_canary/_constants.py b/python/sglang/jit_kernel/tests/kv_canary/_constants.py new file mode 100644 index 000000000..5fb8874be --- /dev/null +++ b/python/sglang/jit_kernel/tests/kv_canary/_constants.py @@ -0,0 +1,12 @@ +from __future__ import annotations + +from sglang.jit_kernel.kv_canary.verify import CANARY_SLOT_BYTES + +# Default fixture sizes — small enough for fast tests, large enough that ring overflow / multi-req cases +# stay realistic without bloating the assertion surface. +DEFAULT_RING_CAPACITY: int = 64 +DEFAULT_NUM_SLOTS: int = 32 +DEFAULT_SLOT_STRIDE_BYTES: int = CANARY_SLOT_BYTES + +_U64_MASK: int = (1 << 64) - 1 +_I64_SIGN_BIT: int = 1 << 63 diff --git a/python/sglang/jit_kernel/tests/kv_canary/test_const_sync.py b/python/sglang/jit_kernel/tests/kv_canary/test_const_sync.py new file mode 100644 index 000000000..ce8d6957a --- /dev/null +++ b/python/sglang/jit_kernel/tests/kv_canary/test_const_sync.py @@ -0,0 +1,70 @@ +from __future__ import annotations + +import re +from pathlib import Path + +from sglang.jit_kernel.kv_canary import consts +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=5, suite="base-b-kernel-unit-1-gpu-large") + + +_CONSTS_CUH: Path = ( + Path(__file__).resolve().parents[2] / "csrc" / "kv_canary" / "consts.cuh" +) + + +def _camel_to_upper_snake(name: str) -> str: + return re.sub(r"([A-Z])", r"_\1", name).lstrip("_").upper() + + +def _decode(expr: str) -> int: + expr = expr.strip().rstrip("UuLl") + if "<<" in expr: + return 1 << int(expr.split("<<")[1].strip()) + return int(expr, 0) + + +def _parse_constexpr_ints(source: str) -> dict[str, int]: + pattern = re.compile(r"constexpr\s+(?:[\w:]+)\s+(k[A-Za-z]\w*)\s*=\s*([^;]+);") + return {name: _decode(rhs) for name, rhs in pattern.findall(source)} + + +def _parse_enum_class(source: str, enum_name: str) -> dict[str, int]: + pattern = re.compile( + r"enum\s+class\s+" + re.escape(enum_name) + r"\s*:\s*[^\{]+\{([^}]+)\}" + ) + body = pattern.search(source).group(1) + member_re = re.compile(r"(k[A-Za-z]\w*)\s*=\s*([^,]+)") + return {name: _decode(rhs) for name, rhs in member_re.findall(body)} + + +def test_int_consts_sync() -> None: + cpp = _parse_constexpr_ints(_CONSTS_CUH.read_text(encoding="utf-8")) + cpp_normalized = {_camel_to_upper_snake(n[1:]): v for n, v in cpp.items()} + py = { + n: v + for n, v in vars(consts).items() + if isinstance(v, int) and not isinstance(v, bool) and not n.startswith("_") + } + assert cpp_normalized == py + + +def test_enums_sync() -> None: + cuh = _CONSTS_CUH.read_text(encoding="utf-8") + for enum_name in ("FailReason",): + cpp_members = _parse_enum_class(cuh, enum_name) + py_enum = getattr(consts, enum_name) + cpp_normalized = { + _camel_to_upper_snake(n[1:]): v for n, v in cpp_members.items() + } + py_normalized = {m.name: int(m.value) for m in py_enum} + assert cpp_normalized == py_normalized + + +if __name__ == "__main__": + import sys + + import pytest + + sys.exit(pytest.main([__file__, "-v"])) diff --git a/python/sglang/jit_kernel/tests/kv_canary/test_utils.py b/python/sglang/jit_kernel/tests/kv_canary/test_utils.py new file mode 100644 index 000000000..429cf2d26 --- /dev/null +++ b/python/sglang/jit_kernel/tests/kv_canary/test_utils.py @@ -0,0 +1,62 @@ +from __future__ import annotations + +from sglang.jit_kernel.benchmark.kv_canary.utils import ( + MAX_EXTEND_TOKENS_PER_FORWARD, + build_fast_matrix_cases, + build_full_matrix_cases, + cases_to_x_vals, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=3, suite="base-a-test-cpu") + + +def test_fast_matrix_cases_include_e2e_decode_and_chunked_prefill_scenarios() -> None: + cases = build_fast_matrix_cases() + scenarios = {case.scenario for case in cases} + + assert { + "e2e_decode_steady", + "e2e_decode_tail", + "e2e_prefill_chunk_first", + "e2e_prefill_chunk_second", + "e2e_prefill_chunk_mid", + "e2e_prefill_chunk_last", + } <= scenarios + + +def test_extend_cases_are_bounded_to_scheduler_chunk_size() -> None: + cases = build_full_matrix_cases() + bad_cases = [ + case + for case in cases + if case.mode == "extend" + and case.bs * case.extend_len > MAX_EXTEND_TOKENS_PER_FORWARD + ] + + assert bad_cases == [] + + +def test_cases_to_x_vals_includes_scenario_axis() -> None: + case = build_fast_matrix_cases()[0] + + x_vals = cases_to_x_vals([case]) + + assert x_vals == [ + ( + case.scenario, + case.bs, + case.prefix_len, + case.mode, + case.extend_len, + case.pool_kind, + ) + ] + + +if __name__ == "__main__": + import sys + + import pytest + + sys.exit(pytest.main([__file__, "-v"]))