diff --git a/python/sglang/jit_kernel/benchmark/kv_canary/bench_verify.py b/python/sglang/jit_kernel/benchmark/kv_canary/bench_verify.py index 043e0b65a..a027d2a98 100644 --- a/python/sglang/jit_kernel/benchmark/kv_canary/bench_verify.py +++ b/python/sglang/jit_kernel/benchmark/kv_canary/bench_verify.py @@ -13,6 +13,7 @@ from sglang.jit_kernel.benchmark.kv_canary.utils import ( build_fast_matrix_cases, build_full_matrix_cases, cases_to_x_vals, + make_real_kv_sources, naive_slot_copy_fn, ) from sglang.jit_kernel.benchmark.utils import ( @@ -24,6 +25,7 @@ from sglang.jit_kernel.kv_canary import consts from sglang.jit_kernel.kv_canary.verify import ( CANARY_SLOT_BYTES, CanaryLaunchTag, + RealKvSource, VerifyOrWriteContext, VerifyPlan, launch_canary_verify_kernel, @@ -40,6 +42,8 @@ _X_NAMES = [ "mode", "extend_len", "pool_kind", + "real_kv_kind", + "hash_mode", ] _X_VALS = cases_to_x_vals( get_benchmark_range( @@ -76,6 +80,7 @@ def _build_verify_inputs(case: BenchCase, *, device: torch.device) -> Tuple[ torch.Tensor, torch.Tensor, torch.Tensor, + tuple[RealKvSource, ...], ]: total_entries = _verify_entry_count(case) capacity = max(1, total_entries) @@ -125,6 +130,10 @@ def _build_verify_inputs(case: BenchCase, *, device: torch.device) -> Tuple[ kernel_run_counter = torch.zeros(1, dtype=torch.int64, device=device) enable_chain_position_assert = torch.ones(1, dtype=torch.int32, device=device) + real_kv_sources = make_real_kv_sources( + kind=case.real_kv_kind, num_slots=num_slots, device=device + ) + return ( canary_buf, plan, @@ -133,6 +142,7 @@ def _build_verify_inputs(case: BenchCase, *, device: torch.device) -> Tuple[ slot_run_counter, kernel_run_counter, enable_chain_position_assert, + real_kv_sources, ) @@ -144,7 +154,9 @@ def _build_context( slot_run_counter: torch.Tensor, kernel_run_counter: torch.Tensor, enable_chain_position_assert: torch.Tensor, + real_kv_sources: tuple[RealKvSource, ...], kernel_kind: CanaryLaunchTag, + hash_mode: consts.RealKvHashMode, ) -> VerifyOrWriteContext: return VerifyOrWriteContext( canary_buf=canary_buf, @@ -153,6 +165,8 @@ def _build_context( violation_write_index=violation_write_index, slot_run_counter=slot_run_counter, kernel_run_counter=kernel_run_counter, + real_kv_sources=real_kv_sources, + real_kv_hash_mode=hash_mode, enable_chain_position_assert=enable_chain_position_assert, ) @@ -177,6 +191,8 @@ def benchmark( mode: str, extend_len: int, pool_kind: str, + real_kv_kind: str, + hash_mode: str, provider: str, ) -> Tuple[float, float, float]: case = BenchCase( @@ -186,6 +202,8 @@ def benchmark( mode=mode, extend_len=extend_len, pool_kind=pool_kind, + real_kv_kind=real_kv_kind, + hash_mode=hash_mode, ) device = torch.device(DEFAULT_DEVICE) @@ -198,7 +216,9 @@ def benchmark( slot_run_counter, kernel_run_counter, enable_chain_position_assert, + real_kv_sources, ) = _build_verify_inputs(case, device=device) + hash_mode_enum = consts.RealKvHashMode[case.hash_mode.upper()] context = _build_context( canary_buf=canary_buf, violation_ring=violation_ring, @@ -206,7 +226,9 @@ def benchmark( slot_run_counter=slot_run_counter, kernel_run_counter=kernel_run_counter, enable_chain_position_assert=enable_chain_position_assert, + real_kv_sources=real_kv_sources, kernel_kind=CanaryLaunchTag.HEAD_K_FULL, + hash_mode=hash_mode_enum, ) def fn() -> None: @@ -247,6 +269,8 @@ def benchmark_kernel_kind( mode="extend", extend_len=128, pool_kind="full", + real_kv_kind="none", + hash_mode="none", ) device = torch.device(DEFAULT_DEVICE) @@ -258,8 +282,10 @@ def benchmark_kernel_kind( slot_run_counter, kernel_run_counter, enable_chain_position_assert, + real_kv_sources, ) = _build_verify_inputs(case, device=device) kernel_kind = CanaryLaunchTag[kernel_kind_name] + hash_mode_enum = consts.RealKvHashMode[case.hash_mode.upper()] context = _build_context( canary_buf=canary_buf, violation_ring=violation_ring, @@ -267,7 +293,9 @@ def benchmark_kernel_kind( slot_run_counter=slot_run_counter, kernel_run_counter=kernel_run_counter, enable_chain_position_assert=enable_chain_position_assert, + real_kv_sources=real_kv_sources, kernel_kind=kernel_kind, + hash_mode=hash_mode_enum, ) def fn() -> None: diff --git a/python/sglang/jit_kernel/benchmark/kv_canary/bench_write.py b/python/sglang/jit_kernel/benchmark/kv_canary/bench_write.py index 1721f4a3b..8dfd28d09 100644 --- a/python/sglang/jit_kernel/benchmark/kv_canary/bench_write.py +++ b/python/sglang/jit_kernel/benchmark/kv_canary/bench_write.py @@ -13,6 +13,7 @@ from sglang.jit_kernel.benchmark.kv_canary.utils import ( build_fast_matrix_cases, build_full_matrix_cases, cases_to_x_vals, + make_real_kv_sources, naive_slot_copy_fn, ) from sglang.jit_kernel.benchmark.utils import ( @@ -39,6 +40,8 @@ _X_NAMES = [ "mode", "extend_len", "pool_kind", + "real_kv_kind", + "hash_mode", ] _X_VALS = cases_to_x_vals( get_benchmark_range( @@ -142,6 +145,10 @@ def _build_write_inputs( kernel_run_counter = torch.zeros(1, dtype=torch.int64, device=device) enable_chain_position_assert = torch.ones(1, dtype=torch.int32, device=device) + real_kv_sources = make_real_kv_sources( + kind=case.real_kv_kind, num_slots=num_slots, device=device + ) + return dict( canary_buf=canary_buf, plan=plan, @@ -155,6 +162,7 @@ def _build_write_inputs( slot_run_counter=slot_run_counter, kernel_run_counter=kernel_run_counter, enable_chain_position_assert=enable_chain_position_assert, + real_kv_sources=real_kv_sources, ) @@ -162,6 +170,7 @@ def _build_context( *, inputs: dict, kernel_kind: CanaryLaunchTag, + hash_mode: consts.RealKvHashMode, ) -> VerifyOrWriteContext: return VerifyOrWriteContext( canary_buf=inputs["canary_buf"], @@ -170,6 +179,8 @@ def _build_context( violation_write_index=inputs["violation_write_index"], slot_run_counter=inputs["slot_run_counter"], kernel_run_counter=inputs["kernel_run_counter"], + real_kv_sources=inputs["real_kv_sources"], + real_kv_hash_mode=hash_mode, enable_chain_position_assert=inputs["enable_chain_position_assert"], ) @@ -194,6 +205,8 @@ def benchmark( mode: str, extend_len: int, pool_kind: str, + real_kv_kind: str, + hash_mode: str, provider: str, ) -> Tuple[float, float, float]: case = BenchCase( @@ -203,14 +216,18 @@ def benchmark( mode=mode, extend_len=extend_len, pool_kind=pool_kind, + real_kv_kind=real_kv_kind, + hash_mode=hash_mode, ) device = torch.device(DEFAULT_DEVICE) if provider == "canary": inputs = _build_write_inputs(case, device=device) + hash_mode_enum = consts.RealKvHashMode[case.hash_mode.upper()] context = _build_context( inputs=inputs, kernel_kind=CanaryLaunchTag.HEAD_K_FULL, + hash_mode=hash_mode_enum, ) def fn() -> None: @@ -256,6 +273,8 @@ def benchmark_kernel_kind( mode="extend", extend_len=128, pool_kind="full", + real_kv_kind="none", + hash_mode="none", ) device = torch.device(DEFAULT_DEVICE) @@ -264,9 +283,11 @@ def benchmark_kernel_kind( case, device=device, mirror_expected_inputs=enable_write_verify_inputs ) kernel_kind = CanaryLaunchTag[kernel_kind_name] + hash_mode_enum = consts.RealKvHashMode[case.hash_mode.upper()] context = _build_context( inputs=inputs, kernel_kind=kernel_kind, + hash_mode=hash_mode_enum, ) def fn() -> None: diff --git a/python/sglang/jit_kernel/benchmark/kv_canary/utils.py b/python/sglang/jit_kernel/benchmark/kv_canary/utils.py index c0e69542c..80b76170a 100644 --- a/python/sglang/jit_kernel/benchmark/kv_canary/utils.py +++ b/python/sglang/jit_kernel/benchmark/kv_canary/utils.py @@ -5,12 +5,14 @@ from typing import Callable import torch -from sglang.jit_kernel.kv_canary.verify import CANARY_SLOT_BYTES +from sglang.jit_kernel.kv_canary.verify import CANARY_SLOT_BYTES, RealKvSource 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"] +REAL_KV_AXIS: list[str] = ["none", "small_1src", "med_2src", "max_4src"] +HASH_MODE_AXIS: list[str] = ["none", "partial", "all"] SWA_WINDOW: int = 128 RING_CAPACITY: int = 256 MAX_EXTEND_TOKENS_PER_FORWARD: int = 4096 @@ -24,12 +26,14 @@ class BenchCase: mode: str extend_len: int pool_kind: str + real_kv_kind: str + hash_mode: 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}" + f"_{self.pool_kind}_rkv{self.real_kv_kind}_hash{self.hash_mode}" ) @@ -41,6 +45,8 @@ def _case( mode: str, extend_len: int, pool_kind: str, + real_kv_kind: str = "none", + hash_mode: str = "none", ) -> BenchCase: return BenchCase( scenario=scenario, @@ -49,6 +55,8 @@ def _case( mode=mode, extend_len=extend_len, pool_kind=pool_kind, + real_kv_kind=real_kv_kind, + hash_mode=hash_mode, ) @@ -194,6 +202,56 @@ def build_fast_matrix_cases() -> list[BenchCase]: extend_len=1, pool_kind="swa_window_128", ), + _case( + scenario="small_extend_batch_hash", + bs=32, + prefix_len=4096, + mode="extend", + extend_len=128, + pool_kind="full", + real_kv_kind="small_1src", + hash_mode="partial", + ), + _case( + scenario="e2e_prefill_chunk_hash", + bs=1, + prefix_len=12288, + mode="extend", + extend_len=4096, + pool_kind="full", + real_kv_kind="med_2src", + hash_mode="all", + ), + _case( + scenario="e2e_decode_steady_hash", + bs=256, + prefix_len=4096, + mode="decode", + extend_len=1, + pool_kind="full", + real_kv_kind="max_4src", + hash_mode="all", + ), + _case( + scenario="swa_decode_long_prefix_hash", + bs=128, + prefix_len=10240, + mode="decode", + extend_len=1, + pool_kind="swa_window_128", + real_kv_kind="med_2src", + hash_mode="partial", + ), + _case( + scenario="smoke_decode_empty_hash", + bs=1, + prefix_len=0, + mode="decode", + extend_len=1, + pool_kind="full", + real_kv_kind="small_1src", + hash_mode="all", + ), ] ) @@ -229,12 +287,41 @@ def build_full_matrix_cases() -> list[BenchCase]: continue full.append(case) + fast_base_points = [ + (c.bs, c.prefix_len, c.mode, c.extend_len, c.pool_kind) + for c in fast + if c.real_kv_kind == "none" and c.hash_mode == "none" + ] + for bs, prefix_len, mode, extend_len, pool_kind in fast_base_points: + for hash_mode in HASH_MODE_AXIS: + if hash_mode == "none": + continue + for real_kv_kind in REAL_KV_AXIS: + if real_kv_kind == "none": + continue + case = _case( + scenario="fold_matrix", + bs=bs, + prefix_len=prefix_len, + mode=mode, + extend_len=extend_len, + pool_kind=pool_kind, + real_kv_kind=real_kv_kind, + hash_mode=hash_mode, + ) + if not _is_realistic_extend_case(case): + continue + if case.case_id in fast_keys: + continue + full.append(case) + fast_keys.add(case.case_id) + return full def cases_to_x_vals( cases: list[BenchCase], -) -> list[tuple[str, int, int, str, int, str]]: +) -> list[tuple[str, int, int, str, int, str, str, str]]: return [ ( c.scenario, @@ -243,11 +330,59 @@ def cases_to_x_vals( c.mode, c.extend_len, c.pool_kind, + c.real_kv_kind, + c.hash_mode, ) for c in cases ] +def _one_real_kv_source( + *, num_slots: int, num_bytes: int, read_bytes: int, device: torch.device +) -> RealKvSource: + tensor = torch.zeros(max(1, num_slots), num_bytes, dtype=torch.uint8, device=device) + return RealKvSource( + tensor=tensor, + page_size=1, + num_bytes_per_token=num_bytes, + read_bytes=read_bytes, + ) + + +def make_real_kv_sources( + *, kind: str, num_slots: int, device: torch.device +) -> tuple[RealKvSource, ...]: + """Map a ``real_kv_kind`` axis label to a tuple of ``RealKvSource`` configs. + + Byte-volume ladder (none -> small_1src -> med_2src -> max_4src) so the bench exposes the + ``real_kv_fold_sources`` PARTIAL/ALL cost gradient. ``max_4src`` hits the + ``consts.MAX_REAL_KV_SOURCES = 4`` ABI ceiling. + """ + if kind == "none": + return () + if kind == "small_1src": + return ( + _one_real_kv_source( + num_slots=num_slots, num_bytes=16, read_bytes=16, device=device + ), + ) + if kind == "med_2src": + return tuple( + _one_real_kv_source( + num_slots=num_slots, num_bytes=32, read_bytes=16, device=device + ) + for _ in range(2) + ) + if kind == "max_4src": + return tuple( + _one_real_kv_source( + num_slots=num_slots, num_bytes=64, read_bytes=32, device=device + ) + for _ in range(4) + ) + raise ValueError(f"kv-canary bench: unknown real_kv_kind {kind!r}") + + 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) diff --git a/python/sglang/jit_kernel/csrc/kv_canary/canary_common.cuh b/python/sglang/jit_kernel/csrc/kv_canary/canary_common.cuh index 6de0c44fa..138e1cfab 100644 --- a/python/sglang/jit_kernel/csrc/kv_canary/canary_common.cuh +++ b/python/sglang/jit_kernel/csrc/kv_canary/canary_common.cuh @@ -7,6 +7,15 @@ namespace canary { +// Device-side handle for one real-KV source. +struct RealKvSourceHandle { + const uint8_t* tensor; // raw uint8 byte pointer to the source tensor + int32_t row_stride_bytes; // tensor.shape[1] in bytes (may exceed page_size * num_bytes_per_token) + int32_t page_size; + int32_t num_bytes_per_token; + int32_t read_bytes; +}; + SGL_DEVICE uint64_t splitmix64(uint64_t x) { x = (x ^ (x >> 30)) * 0xBF58476D1CE4E5B9ULL; x = (x ^ (x >> 27)) * 0x94D049BB133111EBULL; @@ -20,6 +29,55 @@ SGL_DEVICE uint64_t splitmix64_mix3(uint64_t a, uint64_t b, uint64_t c) { return h; } +// Read 16 aligned bytes from a source as two uint64 little-endian words, following the RealKvSource access +// invariant. The invariant (from kv_canary/verify.py docstring) is: +// +// tensor[slot_idx // page_size, +// (slot_idx % page_size) * num_bytes_per_token + byte_offset] +// +// row_stride_bytes is the dim-1 size of the underlying tensor in bytes (which may exceed +// page_size * num_bytes_per_token; trailing bytes are skipped). +SGL_DEVICE void real_kv_load_uint4( + const RealKvSourceHandle& src, int64_t slot_idx, int64_t byte_offset, uint64_t& word_lo, uint64_t& word_hi) { + const int64_t row = slot_idx / src.page_size; + const int64_t col_within_page = slot_idx % src.page_size; + const int64_t col = col_within_page * src.num_bytes_per_token + byte_offset; + const int64_t flat_index = row * static_cast(src.row_stride_bytes) + col; + const uint4 vec = *reinterpret_cast(src.tensor + flat_index); + word_lo = static_cast(vec.x) | (static_cast(vec.y) << 32); + word_hi = static_cast(vec.z) | (static_cast(vec.w) << 32); +} + +SGL_DEVICE uint64_t real_kv_fold_one_source(const RealKvSourceHandle& src, int64_t slot_idx, RealKvHashMode mode) { + const int64_t effective_read_bytes = (mode == RealKvHashMode::kPartial) ? static_cast(16) : src.read_bytes; + uint64_t acc = 0ULL; + for (int64_t byte_offset = 0; byte_offset < effective_read_bytes; byte_offset += 16) { + uint64_t word_lo; + uint64_t word_hi; + real_kv_load_uint4(src, slot_idx, byte_offset, word_lo, word_hi); + acc = splitmix64(acc ^ word_lo); + acc = splitmix64(acc ^ word_hi); + } + return acc; +} + +// Fold all configured real-KV sources for a given slot. Iterates sequentially and combines each source's +// contribution via acc = splitmix64(acc XOR source_hash); matches _compute_real_kv_hash_scalar in +// kv_canary/verify_ref.py. In OFF mode the function returns 0 unconditionally (the running real_kv_hash +// field is always 0). +SGL_DEVICE uint64_t +real_kv_fold_sources(const RealKvSourceHandle* sources, int num_sources, int64_t slot_idx, RealKvHashMode mode) { + if (mode == RealKvHashMode::kNone || num_sources <= 0) { + return 0ULL; + } + uint64_t acc = 0ULL; + for (int s = 0; s < num_sources; ++s) { + const uint64_t source_hash = real_kv_fold_one_source(sources[s], slot_idx, mode); + acc = splitmix64(acc ^ source_hash); + } + return acc; +} + struct ViolationSink { int64_t* __restrict__ ring; int32_t* __restrict__ write_index; diff --git a/python/sglang/jit_kernel/csrc/kv_canary/canary_verify.cuh b/python/sglang/jit_kernel/csrc/kv_canary/canary_verify.cuh index 339982543..7c63016da 100644 --- a/python/sglang/jit_kernel/csrc/kv_canary/canary_verify.cuh +++ b/python/sglang/jit_kernel/csrc/kv_canary/canary_verify.cuh @@ -38,6 +38,11 @@ struct VerifyKernelParams { // Health counters. int64_t* slot_run_counter; int64_t* kernel_run_counter; + + // Real-KV sources (fixed-size ABI; padding slots have read_bytes = 0). + RealKvSourceHandle sources[kMaxRealKvSources]; + int32_t num_sources; + RealKvHashMode real_kv_hash_mode; }; template @@ -76,13 +81,18 @@ __global__ void canary_verify_kernel(const VerifyKernelParams __grid_constant__ 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 int64_t stored_real_kv_hash = + canary_load_field(p.canary_buf, slot_idx, p.slot_stride_bytes, kCanaryFieldRealKvHash); 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; + const uint64_t expected_real_kv_hash_u64 = + real_kv_fold_sources(p.sources, p.num_sources, slot_idx, p.real_kv_hash_mode); + const int64_t expected_real_kv_hash = static_cast(expected_real_kv_hash_u64); + FailReason fail_reason_bits{}; if (prev_reachable && stored_chain_hash != expected_chain_hash) { fail_reason_bits |= FailReason::kVerifyChainHashMismatch; @@ -95,6 +105,9 @@ __global__ void canary_verify_kernel(const VerifyKernelParams __grid_constant__ if (stored_position != expected_position) { fail_reason_bits |= FailReason::kVerifyPositionMismatch; } + if (stored_real_kv_hash != expected_real_kv_hash) { + fail_reason_bits |= FailReason::kVerifyRealKvHashMismatch; + } if (fail_reason_bits != FailReason{}) { record_violation( @@ -124,6 +137,14 @@ __global__ void canary_verify_kernel(const VerifyKernelParams __grid_constant__ } // namespace // API source of truth: docstring of canary_verify_step in python/sglang/jit_kernel/kv_canary/verify.py. +// +// ABI notes: +// - real_kv_buf_0 .. real_kv_buf_3 are 4 fixed uint8 tensor slots. Unused slots are dummy 1-byte tensors; +// the host wrapper sets real_kv_source_params[s, 2] (read_bytes) = 0 for them. +// - real_kv_source_params is a CPU int32 tensor of shape [kMaxRealKvSources, 3]: per-source (page_size, +// num_bytes_per_token, read_bytes). Lives on CPU because the host wrapper materializes it from a tuple +// of Python dataclasses. +// - num_sources is passed as int64 (tvm-ffi scalar convention) but always in [0, kMaxRealKvSources]. template struct CanaryVerifyKernel { static void @@ -138,7 +159,14 @@ struct CanaryVerifyKernel { tvm::ffi::TensorView violation_ring, tvm::ffi::TensorView violation_write_index, tvm::ffi::TensorView slot_run_counter, - tvm::ffi::TensorView kernel_run_counter) { + tvm::ffi::TensorView kernel_run_counter, + tvm::ffi::TensorView real_kv_buf_0, + tvm::ffi::TensorView real_kv_buf_1, + tvm::ffi::TensorView real_kv_buf_2, + tvm::ffi::TensorView real_kv_buf_3, + tvm::ffi::TensorView real_kv_source_params, + int64_t num_sources, + int64_t real_kv_hash_mode) { using namespace host; SymbolicSize N_slots = {"num_canary_slots"}; @@ -171,6 +199,37 @@ struct CanaryVerifyKernel { .verify(slot_run_counter) .verify(kernel_run_counter); + // Real-KV source buffers are 2-D uint8 (any shape); dim 1 carries the row-stride in bytes. They live on + // CUDA. real_kv_source_params is a small CPU int32 table of length kMaxRealKvSources * 3. + SymbolicSize N_real_kv_rows_0 = {"real_kv_rows_0"}; + SymbolicSize N_real_kv_cols_0 = {"real_kv_cols_0"}; + TensorMatcher({N_real_kv_rows_0, N_real_kv_cols_0}) + .with_dtype() + .with_device(device_) + .verify(real_kv_buf_0); + SymbolicSize N_real_kv_rows_1 = {"real_kv_rows_1"}; + SymbolicSize N_real_kv_cols_1 = {"real_kv_cols_1"}; + TensorMatcher({N_real_kv_rows_1, N_real_kv_cols_1}) + .with_dtype() + .with_device(device_) + .verify(real_kv_buf_1); + SymbolicSize N_real_kv_rows_2 = {"real_kv_rows_2"}; + SymbolicSize N_real_kv_cols_2 = {"real_kv_cols_2"}; + TensorMatcher({N_real_kv_rows_2, N_real_kv_cols_2}) + .with_dtype() + .with_device(device_) + .verify(real_kv_buf_2); + SymbolicSize N_real_kv_rows_3 = {"real_kv_rows_3"}; + SymbolicSize N_real_kv_cols_3 = {"real_kv_cols_3"}; + TensorMatcher({N_real_kv_rows_3, N_real_kv_cols_3}) + .with_dtype() + .with_device(device_) + .verify(real_kv_buf_3); + TensorMatcher({static_cast(kMaxRealKvSources), static_cast(kRealKvSourceFieldsPerEntry)}) + .with_dtype() + .with_device() + .verify(real_kv_source_params); + 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()); @@ -182,6 +241,12 @@ struct CanaryVerifyKernel { static_cast(kCanaryFieldsPerSlot * sizeof(int64_t)), " bytes per slot, got ", slot_stride_bytes); + RuntimeCheck( + num_sources >= 0 && num_sources <= static_cast(kMaxRealKvSources), + "canary_verify: num_sources must be in [0, ", + static_cast(kMaxRealKvSources), + "], got ", + num_sources); VerifyKernelParams p{}; p.canary_buf = static_cast(canary_buf.data_ptr()); @@ -200,6 +265,19 @@ struct CanaryVerifyKernel { p.slot_run_counter = static_cast(slot_run_counter.data_ptr()); p.kernel_run_counter = static_cast(kernel_run_counter.data_ptr()); + // Materialize the source handle array on the host. The CPU param tensor carries the per-source ints. + const int32_t* params = static_cast(real_kv_source_params.data_ptr()); + tvm::ffi::TensorView source_bufs[kMaxRealKvSources] = {real_kv_buf_0, real_kv_buf_1, real_kv_buf_2, real_kv_buf_3}; + for (int s = 0; s < kMaxRealKvSources; ++s) { + p.sources[s].tensor = static_cast(source_bufs[s].data_ptr()); + p.sources[s].row_stride_bytes = static_cast(source_bufs[s].size(1)); + p.sources[s].page_size = params[s * kRealKvSourceFieldsPerEntry + kRealKvSourceFieldPageSize]; + p.sources[s].num_bytes_per_token = params[s * kRealKvSourceFieldsPerEntry + kRealKvSourceFieldNumBytesPerToken]; + p.sources[s].read_bytes = params[s * kRealKvSourceFieldsPerEntry + kRealKvSourceFieldReadBytes]; + } + p.num_sources = static_cast(num_sources); + p.real_kv_hash_mode = static_cast(real_kv_hash_mode); + const uint32_t grid = kPersistentBlocks; LaunchKernel(grid, kVerifyBlockSize, device)(canary_verify_kernel, p); } diff --git a/python/sglang/jit_kernel/csrc/kv_canary/canary_write.cuh b/python/sglang/jit_kernel/csrc/kv_canary/canary_write.cuh index 1401e11ff..1ce2d0918 100644 --- a/python/sglang/jit_kernel/csrc/kv_canary/canary_write.cuh +++ b/python/sglang/jit_kernel/csrc/kv_canary/canary_write.cuh @@ -49,6 +49,11 @@ struct WriteKernelParams { // Gates the chain-step position assert below. Default-on (1); CanaryManager zeros during the // warmup window and flips back in mark_init_finished(). const int32_t* enable_chain_position_assert; + + // Real-KV sources. + RealKvSourceHandle sources[kMaxRealKvSources]; + int32_t num_sources; + RealKvHashMode real_kv_hash_mode; }; __global__ void canary_write_kernel(const WriteKernelParams __grid_constant__ p) { @@ -98,8 +103,8 @@ __global__ void canary_write_kernel(const WriteKernelParams __grid_constant__ p) const int64_t token = p.input_ids[entry_idx]; const int64_t position = p.positions[entry_idx]; - // field[3] (kCanaryFieldRealKvHash) is always 0 in the naive canary. - const int64_t real_kv_hash = 0; + const uint64_t real_kv_hash_u64 = real_kv_fold_sources(p.sources, p.num_sources, slot, p.real_kv_hash_mode); + const int64_t real_kv_hash = static_cast(real_kv_hash_u64); if (p.enable_write_input_assert) { const int64_t expected_token = p.expected_input_tokens[entry_idx]; @@ -163,7 +168,10 @@ __global__ void canary_write_kernel(const WriteKernelParams __grid_constant__ p) // API source of truth: docstring of canary_write_step in python/sglang/jit_kernel/kv_canary/write.py. // -// ABI notes: +// ABI notes (same as verify): +// - real_kv_buf_0 .. real_kv_buf_3 are 4 fixed uint8 tensor slots. +// - real_kv_source_params is a CPU int32 [kMaxRealKvSources, 3] table of (page_size, num_bytes_per_token, +// read_bytes) triplets. // - out_cache_loc is caller-pre-translated for SWA groups; -1 entries mark skip. The kernel does not // consult any LUT. inline void canary_write_step_cuda( @@ -182,7 +190,14 @@ inline void canary_write_step_cuda( tvm::ffi::TensorView violation_write_index, tvm::ffi::TensorView slot_run_counter, tvm::ffi::TensorView kernel_run_counter, - tvm::ffi::TensorView enable_chain_position_assert) { + tvm::ffi::TensorView enable_chain_position_assert, + tvm::ffi::TensorView real_kv_buf_0, + tvm::ffi::TensorView real_kv_buf_1, + tvm::ffi::TensorView real_kv_buf_2, + tvm::ffi::TensorView real_kv_buf_3, + tvm::ffi::TensorView real_kv_source_params, + int64_t num_sources, + int64_t real_kv_hash_mode) { using namespace host; SymbolicSize N_slots = {"num_canary_slots"}; @@ -235,6 +250,35 @@ inline void canary_write_step_cuda( .verify(kernel_run_counter); TensorMatcher({1}).with_dtype().with_device(device_).verify(enable_chain_position_assert); + SymbolicSize N_real_kv_rows_0 = {"real_kv_rows_0"}; + SymbolicSize N_real_kv_cols_0 = {"real_kv_cols_0"}; + TensorMatcher({N_real_kv_rows_0, N_real_kv_cols_0}) + .with_dtype() + .with_device(device_) + .verify(real_kv_buf_0); + SymbolicSize N_real_kv_rows_1 = {"real_kv_rows_1"}; + SymbolicSize N_real_kv_cols_1 = {"real_kv_cols_1"}; + TensorMatcher({N_real_kv_rows_1, N_real_kv_cols_1}) + .with_dtype() + .with_device(device_) + .verify(real_kv_buf_1); + SymbolicSize N_real_kv_rows_2 = {"real_kv_rows_2"}; + SymbolicSize N_real_kv_cols_2 = {"real_kv_cols_2"}; + TensorMatcher({N_real_kv_rows_2, N_real_kv_cols_2}) + .with_dtype() + .with_device(device_) + .verify(real_kv_buf_2); + SymbolicSize N_real_kv_rows_3 = {"real_kv_rows_3"}; + SymbolicSize N_real_kv_cols_3 = {"real_kv_cols_3"}; + TensorMatcher({N_real_kv_rows_3, N_real_kv_cols_3}) + .with_dtype() + .with_device(device_) + .verify(real_kv_buf_3); + TensorMatcher({static_cast(kMaxRealKvSources), static_cast(kRealKvSourceFieldsPerEntry)}) + .with_dtype() + .with_device() + .verify(real_kv_source_params); + const int64_t slot_stride_bytes = N_stride.unwrap(); const int32_t write_req_capacity = static_cast(N_write_reqs.unwrap()); const int32_t ring_capacity = static_cast(N_ring.unwrap()); @@ -252,6 +296,12 @@ inline void canary_write_step_cuda( static_cast(kCanaryFieldsPerSlot * sizeof(int64_t)), " bytes per slot, got ", slot_stride_bytes); + RuntimeCheck( + num_sources >= 0 && num_sources <= static_cast(kMaxRealKvSources), + "canary_write: num_sources must be in [0, ", + static_cast(kMaxRealKvSources), + "], got ", + num_sources); WriteKernelParams p{}; p.canary_buf = static_cast(canary_buf.data_ptr()); @@ -277,6 +327,18 @@ inline void canary_write_step_cuda( p.kernel_run_counter = static_cast(kernel_run_counter.data_ptr()); p.enable_chain_position_assert = static_cast(enable_chain_position_assert.data_ptr()); + const int32_t* params = static_cast(real_kv_source_params.data_ptr()); + tvm::ffi::TensorView source_bufs[kMaxRealKvSources] = {real_kv_buf_0, real_kv_buf_1, real_kv_buf_2, real_kv_buf_3}; + for (int s = 0; s < kMaxRealKvSources; ++s) { + p.sources[s].tensor = static_cast(source_bufs[s].data_ptr()); + p.sources[s].row_stride_bytes = static_cast(source_bufs[s].size(1)); + p.sources[s].page_size = params[s * kRealKvSourceFieldsPerEntry + kRealKvSourceFieldPageSize]; + p.sources[s].num_bytes_per_token = params[s * kRealKvSourceFieldsPerEntry + kRealKvSourceFieldNumBytesPerToken]; + p.sources[s].read_bytes = params[s * kRealKvSourceFieldsPerEntry + kRealKvSourceFieldReadBytes]; + } + p.num_sources = static_cast(num_sources); + p.real_kv_hash_mode = static_cast(real_kv_hash_mode); + // Grid: one block per write req capacity slot; the kernel early-exits on r >= write_num_valid_reqs[0]. // Always launch at least one block so the unconditional kernel_run_counter bump runs even when capacity // == 0. diff --git a/python/sglang/jit_kernel/csrc/kv_canary/consts.cuh b/python/sglang/jit_kernel/csrc/kv_canary/consts.cuh index d36577c57..49e66a2f7 100644 --- a/python/sglang/jit_kernel/csrc/kv_canary/consts.cuh +++ b/python/sglang/jit_kernel/csrc/kv_canary/consts.cuh @@ -46,4 +46,17 @@ constexpr FailReason& operator|=(FailReason& a, FailReason b) { return a; } +enum class RealKvHashMode : int32_t { + kNone = 0, + kPartial = 1, + kAll = 2, +}; + +constexpr int kMaxRealKvSources = 4; + +constexpr int kRealKvSourceFieldsPerEntry = 3; +constexpr int kRealKvSourceFieldPageSize = 0; +constexpr int kRealKvSourceFieldNumBytesPerToken = 1; +constexpr int kRealKvSourceFieldReadBytes = 2; + } // namespace canary diff --git a/python/sglang/jit_kernel/kv_canary/consts.py b/python/sglang/jit_kernel/kv_canary/consts.py index cde98eb8a..aed64508f 100644 --- a/python/sglang/jit_kernel/kv_canary/consts.py +++ b/python/sglang/jit_kernel/kv_canary/consts.py @@ -1,6 +1,6 @@ from __future__ import annotations -from enum import IntFlag +from enum import IntEnum, IntFlag from typing import Final CANARY_CHAIN_ANCHOR: Final[int] = 0xC0FFEE1234567890 @@ -40,6 +40,20 @@ class FailReason(IntFlag): VERIFY_TOKEN_MISMATCH = 1 << 5 +MAX_REAL_KV_SOURCES: Final[int] = 4 + +REAL_KV_SOURCE_FIELDS_PER_ENTRY: Final[int] = 3 +REAL_KV_SOURCE_FIELD_PAGE_SIZE: Final[int] = 0 +REAL_KV_SOURCE_FIELD_NUM_BYTES_PER_TOKEN: Final[int] = 1 +REAL_KV_SOURCE_FIELD_READ_BYTES: Final[int] = 2 + + +class RealKvHashMode(IntEnum): + NONE = 0 + PARTIAL = 1 + ALL = 2 + + _U64_MASK: int = (1 << 64) - 1 diff --git a/python/sglang/jit_kernel/kv_canary/verify.py b/python/sglang/jit_kernel/kv_canary/verify.py index e0121bd40..2f43f0253 100644 --- a/python/sglang/jit_kernel/kv_canary/verify.py +++ b/python/sglang/jit_kernel/kv_canary/verify.py @@ -38,6 +38,90 @@ def _assert_contiguous(tensor: torch.Tensor, name: str) -> None: raise ValueError(f"kv-canary: {name} must be contiguous") +@dataclass(frozen=True, slots=True, kw_only=True) +class RealKvSource: + """One piece of real KV the canary folds into its fingerprint. + + Slot access invariant (must hold for every source, regardless of underlying layout) — for a given slot_idx, + the canary reads exactly these bytes: + + tensor[ + slot_idx // page_size, + (slot_idx % page_size) * num_bytes_per_token + : ((slot_idx % page_size) + 1) * num_bytes_per_token + ] + + Note that ``tensor`` may have "holes" in dim 1 — ``tensor.shape[1]`` can exceed ``page_size * + num_bytes_per_token``. Trailing bytes of each row are ignored by the canary; this is exactly how the + abstraction accommodates pools whose per-row layout interleaves canary-relevant bytes with other metadata + (layer-split storage, K/V interleaving, ...). When ``page_size == 1`` the pattern + collapses to the simple ``tensor[slot_idx, :num_bytes_per_token]`` case. + + A pool may expose multiple RealKvSource instances per (canary buffer × K/V half) — the launch wrappers + iterate the source list and fold each into the running real_kv_hash via splitmix64 (one int64 fingerprint + per slot, regardless of source count). + + Pool patchers construct sources by: + - viewing / reshaping the underlying KV layer into the canonical [num_rows, dim1_bytes] form (no stage-copy + needed when the underlying storage is already row-major contiguous on dim 0), + - choosing ``page_size`` and ``num_bytes_per_token`` so that the access pattern above lands on the bytes + the canary should fingerprint, + - leaving any per-row padding / non-canary bytes in the trailing portion of each row (they will simply be + skipped). + + 16-byte alignment precondition: the CUDA fold kernel issues 128-bit aligned loads, so ``read_bytes``, + ``num_bytes_per_token``, and the row stride (``tensor.shape[1]`` in bytes) must all be positive + multiples of 16. There is no "skip this source" sentinel — callers omit the source from their + ``real_kv_sources`` tuple entirely (factory helpers return an empty tuple in that case). + + Fields: + tensor: The source tensor, any shape such that the access pattern above yields ``num_bytes_per_token`` + uint8 bytes per slot. Dtype is whatever the underlying pool uses; the canary views the relevant + bytes via ``.view(torch.uint8)``. + page_size: Number of slots packed into one row of dim 0. ``>= 1``. + num_bytes_per_token: Bytes per slot in the dim-1 strip the canary reads. Must be a positive + multiple of 16. + read_bytes: Leading bytes (out of ``num_bytes_per_token``) per slot folded into the fingerprint. + Must be a positive multiple of 16, ``<= num_bytes_per_token``. + """ + + tensor: torch.Tensor + page_size: int + num_bytes_per_token: int + read_bytes: int + + def __post_init__(self) -> None: + if self.page_size < 1: + raise ValueError( + f"kv-canary: RealKvSource.page_size must be >= 1, got {self.page_size}" + ) + if self.num_bytes_per_token <= 0 or self.num_bytes_per_token % 16 != 0: + raise ValueError( + f"kv-canary: RealKvSource.num_bytes_per_token must be a positive multiple of 16, " + f"got {self.num_bytes_per_token}" + ) + if ( + self.read_bytes <= 0 + or self.read_bytes > self.num_bytes_per_token + or self.read_bytes % 16 != 0 + ): + raise ValueError( + f"kv-canary: RealKvSource.read_bytes must be a positive multiple of 16 in " + f"(0, num_bytes_per_token={self.num_bytes_per_token}], got {self.read_bytes}" + ) + if self.tensor.ndim < 2: + raise ValueError( + f"kv-canary: RealKvSource.tensor must be at least 2-D, got shape {tuple(self.tensor.shape)}" + ) + row_stride_bytes = int(self.tensor.shape[1]) * self.tensor.element_size() + if row_stride_bytes % 16 != 0: + raise ValueError( + f"kv-canary: RealKvSource.tensor dim-1 byte width must be a multiple of 16, " + f"got {row_stride_bytes} bytes (shape={tuple(self.tensor.shape)}, " + f"dtype={self.tensor.dtype})" + ) + + @dataclass(frozen=True, slots=True, kw_only=True) class VerifyOrWriteContext: """Shared launch context for canary verify/write kernels. @@ -53,6 +137,10 @@ class VerifyOrWriteContext: 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. + real_kv_sources: Real KV pieces folded into each slot's real_kv_hash, as a tuple of RealKvSource. Empty + tuple disables the mixin. Multiple sources are folded sequentially via splitmix64 to produce one + int64 fingerprint per slot. + real_kv_hash_mode: RealKvHashMode (NONE / PARTIAL / ALL). Applies uniformly across all sources. 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(). @@ -64,6 +152,8 @@ class VerifyOrWriteContext: violation_write_index: torch.Tensor slot_run_counter: torch.Tensor kernel_run_counter: torch.Tensor + real_kv_sources: tuple[RealKvSource, ...] + real_kv_hash_mode: consts.RealKvHashMode enable_chain_position_assert: torch.Tensor @@ -159,11 +249,10 @@ def launch_canary_verify_kernel( 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. + any mismatch (chain hash / position / real_kv_hash) 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, + prev_hash, real_kv_hash). 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 @@ -174,7 +263,7 @@ def launch_canary_verify_kernel( Args: context: Shared verify/write launch context, including canary buffer, launch tag, violation sink, - and health counters. + health counters, and real KV fingerprint sources. 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 @@ -192,7 +281,9 @@ def launch_canary_verify_kernel( (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 + (c) For each src in real_kv_sources: read src.read_bytes leading bytes from src.tensor[...] (per the + RealKvSource access invariant) and splitmix64-fold into running_real_kv_hash. + - Compare expected vs stored (chain hash, position, real_kv_hash) 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 @@ -213,6 +304,12 @@ def launch_canary_verify_kernel( byte-for-byte. """ canary_buf = context.canary_buf + real_kv_sources = context.real_kv_sources + if len(real_kv_sources) > consts.MAX_REAL_KV_SOURCES: + raise ValueError( + f"kv-canary: at most {consts.MAX_REAL_KV_SOURCES} RealKvSource entries supported by the CUDA ABI, " + f"got {len(real_kv_sources)}" + ) _assert_contiguous(canary_buf, "canary_buf") _assert_contiguous(plan.verify_slot_indices, "plan.verify_slot_indices") @@ -226,6 +323,10 @@ def launch_canary_verify_kernel( _assert_contiguous(context.slot_run_counter, "slot_run_counter") _assert_contiguous(context.kernel_run_counter, "kernel_run_counter") + padded_bufs, source_params = _build_real_kv_source_abi( + real_kv_sources=real_kv_sources, device=canary_buf.device + ) + module = _jit_canary_verify_module(check_verify_expected_token) module.canary_verify_step_cuda( canary_buf, @@ -240,6 +341,13 @@ def launch_canary_verify_kernel( context.violation_write_index, context.slot_run_counter, context.kernel_run_counter, + padded_bufs[0], + padded_bufs[1], + padded_bufs[2], + padded_bufs[3], + source_params, + len(real_kv_sources), + int(context.real_kv_hash_mode), ) @@ -257,3 +365,38 @@ def _jit_canary_verify_module(check_verify_expected_token: bool) -> "Module": ), ], ) + + +def _build_real_kv_source_abi( + *, + real_kv_sources: tuple[RealKvSource, ...], + device: torch.device, +) -> tuple[list[torch.Tensor], torch.Tensor]: + padded_bufs: list[torch.Tensor] = [] + params = torch.zeros( + (consts.MAX_REAL_KV_SOURCES, consts.REAL_KV_SOURCE_FIELDS_PER_ENTRY), + dtype=torch.int32, + device="cpu", + ) + + for i, source in enumerate(real_kv_sources): + _assert_contiguous(source.tensor, f"real_kv_sources[{i}].tensor") + source_u8 = source.tensor.view(torch.uint8) + if source_u8.dim() != 2: + raise ValueError( + f"kv-canary: real_kv_sources[{i}].tensor (viewed as uint8) must be 2-D, " + f"got {source_u8.dim()}-D" + ) + padded_bufs.append(source_u8) + params[i, consts.REAL_KV_SOURCE_FIELD_PAGE_SIZE] = source.page_size + params[i, consts.REAL_KV_SOURCE_FIELD_NUM_BYTES_PER_TOKEN] = ( + source.num_bytes_per_token + ) + params[i, consts.REAL_KV_SOURCE_FIELD_READ_BYTES] = source.read_bytes + + # Pad bufs (never read by the kernel — num_sources bounds the iteration); params already zero. + dummy = torch.empty((1, 1), dtype=torch.uint8, device=device) + for _ in range(len(real_kv_sources), consts.MAX_REAL_KV_SOURCES): + padded_bufs.append(dummy) + + return padded_bufs, params diff --git a/python/sglang/jit_kernel/kv_canary/verify_ref.py b/python/sglang/jit_kernel/kv_canary/verify_ref.py index 50c054e74..952018309 100644 --- a/python/sglang/jit_kernel/kv_canary/verify_ref.py +++ b/python/sglang/jit_kernel/kv_canary/verify_ref.py @@ -5,6 +5,7 @@ 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 ( + RealKvSource, VerifyOrWriteContext, VerifyPlan, ) @@ -25,6 +26,8 @@ def launch_canary_verify_kernel_torch_reference( violation_write_index = context.violation_write_index slot_run_counter = context.slot_run_counter kernel_run_counter = context.kernel_run_counter + real_kv_sources = context.real_kv_sources + real_kv_hash_mode = context.real_kv_hash_mode work_device = torch.device("cpu") @@ -113,8 +116,13 @@ def launch_canary_verify_kernel_torch_reference( 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 + expected_real_kv_hash_u64 = _compute_real_kv_hash_scalar( + slot_idx=slot_idx, + real_kv_sources=real_kv_sources, + real_kv_hash_mode=real_kv_hash_mode, + work_device=work_device, + ) + expected_real_kv_hash = _to_signed_int64(expected_real_kv_hash_u64) fail_reason = consts.FailReason(0) if prev_reachable and stored_chain_hash != expected_chain_hash: @@ -180,3 +188,60 @@ def compute_slot_hash(buf_i64: torch.Tensor, source_slot_idx: int) -> int: 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) + + +def _compute_real_kv_hash_scalar( + *, + slot_idx: int, + real_kv_sources: tuple[RealKvSource, ...], + real_kv_hash_mode: consts.RealKvHashMode, + work_device: torch.device, +) -> int: + mode = int(real_kv_hash_mode) + if mode == int(consts.RealKvHashMode.NONE) or len(real_kv_sources) == 0: + return 0 + + acc: int = 0 + + for source in real_kv_sources: + page_size = source.page_size + num_bytes_per_token = source.num_bytes_per_token + read_bytes = source.read_bytes + tensor_u8 = ( + source.tensor.detach().to(device=work_device).contiguous().view(torch.uint8) + ) + + row = slot_idx // page_size + col_within_page = slot_idx % page_size + col_start = col_within_page * num_bytes_per_token + + effective_read_bytes = ( + 16 if mode == int(consts.RealKvHashMode.PARTIAL) else read_bytes + ) + raw_bytes: list[int] = [] + for b in range(effective_read_bytes): + raw_bytes.append(int(tensor_u8[row, col_start + b].item())) + + source_hash = _splitmix64_fold_bytes_scalar(raw_bytes=raw_bytes) + + combined = acc ^ source_hash + acc = splitmix64(combined) + + return acc + + +def _splitmix64_fold_bytes_scalar(*, raw_bytes: list[int]) -> int: + read_bytes = len(raw_bytes) + pad = (8 - read_bytes % 8) % 8 + padded = raw_bytes + [0] * pad + num_words = len(padded) // 8 + + acc: int = 0 + for w in range(num_words): + word: int = 0 + for k in range(8): + word |= padded[w * 8 + k] << (8 * k) + word &= _U64_MASK + acc = splitmix64(acc ^ word) + + return acc diff --git a/python/sglang/jit_kernel/kv_canary/write.py b/python/sglang/jit_kernel/kv_canary/write.py index 80e6c89eb..d3d0c4382 100644 --- a/python/sglang/jit_kernel/kv_canary/write.py +++ b/python/sglang/jit_kernel/kv_canary/write.py @@ -5,9 +5,11 @@ from typing import TYPE_CHECKING import torch +from sglang.jit_kernel.kv_canary import consts from sglang.jit_kernel.kv_canary.verify import ( VerifyOrWriteContext, _assert_contiguous, + _build_real_kv_source_abi, ) from sglang.jit_kernel.utils import cache_once, load_jit @@ -92,7 +94,7 @@ def launch_canary_write_kernel( - ``slot`` = ``out_cache_loc[i]`` (caller-pre-translated for SWA groups; entries set to -1 are skipped). - ``token / position`` = ``input_ids[i] / positions[i]``. - - ``real_kv_hash`` is always 0 (the naive canary writes a constant 0 into field[3]). + - ``real_kv_hash`` = ``real_kv_fold_sources(real_kv_sources, slot)`` if ``real_kv_hash_mode != NONE`` else 0. - Store 4 int64s ``(token, position, running_prev_hash, real_kv_hash)`` into ``canary_buf[slot]``. - Advance ``running_prev_hash = splitmix64_mix3(prev, token, position)``, where splitmix64_mix3 folds each input via ``acc = splitmix64(acc ^ next)`` starting from ``splitmix64(prev)``. @@ -121,7 +123,7 @@ def launch_canary_write_kernel( Args: context: Shared verify/write launch context, including canary buffer, launch tag, violation sink, - and health counters. + health counters, and real KV fingerprint sources. plan: Pre-allocated WritePlan. input_ids: ForwardBatch.input_ids; token ids being written, shape [num_tokens_padded], int64. Flattened across reqs in plan.write_offsets order; tail beyond @@ -157,7 +159,8 @@ def launch_canary_write_kernel( slot = out_cache_loc[i]; // caller-pre-translated; the kernel never consults a LUT if (slot < 0) continue; // -1 sentinel = skip (SWA out-of-window or padding) token = input_ids[i]; position = positions[i]; - real_kv_hash = 0; // naive canary always stores 0 into field[3] + real_kv_hash = (real_kv_hash_mode == NONE) ? 0 : real_kv_fold_sources(real_kv_sources, slot); + // applies RealKvSource access invariant if enable_write_input_assert: if token != expected_input_tokens[i] or position != expected_input_positions[i]: record_violation(); // chain still advances on the ACTUAL (token, position) below @@ -181,6 +184,12 @@ def launch_canary_write_kernel( byte-for-byte. """ canary_buf = context.canary_buf + real_kv_sources = context.real_kv_sources + if len(real_kv_sources) > consts.MAX_REAL_KV_SOURCES: + raise ValueError( + f"kv-canary: at most {consts.MAX_REAL_KV_SOURCES} RealKvSource entries supported by the CUDA ABI, " + f"got {len(real_kv_sources)}" + ) _assert_contiguous(canary_buf, "canary_buf") _assert_contiguous(plan.write_offsets, "plan.write_offsets") @@ -209,6 +218,10 @@ def launch_canary_write_kernel( context.enable_chain_position_assert, "enable_chain_position_assert" ) + padded_bufs, source_params = _build_real_kv_source_abi( + real_kv_sources=real_kv_sources, device=canary_buf.device + ) + module = _jit_canary_write_module() module.canary_write_step_cuda( canary_buf, @@ -227,6 +240,13 @@ def launch_canary_write_kernel( context.slot_run_counter, context.kernel_run_counter, context.enable_chain_position_assert, + padded_bufs[0], + padded_bufs[1], + padded_bufs[2], + padded_bufs[3], + source_params, + len(real_kv_sources), + int(context.real_kv_hash_mode), ) diff --git a/python/sglang/jit_kernel/kv_canary/write_ref.py b/python/sglang/jit_kernel/kv_canary/write_ref.py index 21f9a1442..6165e9bd9 100644 --- a/python/sglang/jit_kernel/kv_canary/write_ref.py +++ b/python/sglang/jit_kernel/kv_canary/write_ref.py @@ -7,6 +7,7 @@ from sglang.jit_kernel.kv_canary.verify import ( VerifyOrWriteContext, ) from sglang.jit_kernel.kv_canary.verify_ref import ( + _compute_real_kv_hash_scalar, _to_signed_int64, compute_slot_hash, splitmix64_mix3, @@ -31,6 +32,8 @@ def launch_canary_write_kernel_torch_reference( violation_write_index = context.violation_write_index slot_run_counter = context.slot_run_counter kernel_run_counter = context.kernel_run_counter + real_kv_sources = context.real_kv_sources + real_kv_hash_mode = context.real_kv_hash_mode enable_chain_position_assert_value = int( context.enable_chain_position_assert.detach().to("cpu").item() ) @@ -124,6 +127,13 @@ def launch_canary_write_kernel_torch_reference( token = int(input_ids_host[entry_idx].item()) position = int(positions_host[entry_idx].item()) + real_kv_hash_u64 = _compute_real_kv_hash_scalar( + slot_idx=slot, + real_kv_sources=real_kv_sources, + real_kv_hash_mode=real_kv_hash_mode, + work_device=work_device, + ) + if enable_write_input_assert: assert expected_input_tokens_host is not None assert expected_input_positions_host is not None @@ -172,7 +182,9 @@ def launch_canary_write_kernel_torch_reference( buf_i64[slot, consts.CANARY_FIELD_PREV_HASH] = _to_signed_int64( running_prev_hash ) - buf_i64[slot, consts.CANARY_FIELD_REAL_KV_HASH] = 0 + buf_i64[slot, consts.CANARY_FIELD_REAL_KV_HASH] = _to_signed_int64( + real_kv_hash_u64 + ) running_prev_hash = splitmix64_mix3(running_prev_hash, token, position) diff --git a/python/sglang/jit_kernel/tests/kv_canary/_canary_helpers.py b/python/sglang/jit_kernel/tests/kv_canary/_canary_helpers.py index 42366f55c..dd1ce578c 100644 --- a/python/sglang/jit_kernel/tests/kv_canary/_canary_helpers.py +++ b/python/sglang/jit_kernel/tests/kv_canary/_canary_helpers.py @@ -16,6 +16,10 @@ from sglang.jit_kernel.tests.kv_canary._constants import ( DEFAULT_RING_CAPACITY, DEFAULT_SLOT_STRIDE_BYTES, ) +from sglang.jit_kernel.tests.kv_canary._fixtures import ( + make_real_kv_source, + make_real_kv_sources, +) __all__ = [ "FakeViolationLog", @@ -26,6 +30,8 @@ __all__ = [ "make_canary_buf", "make_canary_buf_pair", "make_log_pair", + "make_real_kv_source", + "make_real_kv_sources", "make_verify_plan", "make_verify_plan_pair", "make_write_plan", @@ -111,7 +117,7 @@ def make_verify_plan( ``expected_input_ids`` defaults to ``[-1] * n_active`` (the verify-kernel "skip token check" sentinel) so existing tests that only exercise the - chain / position paths keep working unchanged. + chain / position / real-kv-hash paths keep working unchanged. """ n_active = len(slot_indices) if not (len(positions) == n_active and len(prev_slot_indices) == n_active): @@ -246,13 +252,13 @@ def write_slot_fields( token: int, position: int, prev_hash: int, + real_kv_hash: int, ) -> None: view = canary_buf.view(torch.int64) view[slot_idx, 0] = token view[slot_idx, 1] = position view[slot_idx, 2] = prev_hash - # field[3] (real_kv_hash) is always 0 in the naive canary. - view[slot_idx, 3] = 0 + view[slot_idx, 3] = real_kv_hash def stamp_pair( @@ -262,6 +268,7 @@ def stamp_pair( token: int, position: int, prev_hash: int, + real_kv_hash: int = 0, ) -> None: """Stamp the same slot fields into both (cuda, ref) canary buffers.""" for buf in buf_pair: @@ -271,6 +278,7 @@ def stamp_pair( token=token, position=position, prev_hash=prev_hash, + real_kv_hash=real_kv_hash, ) @@ -288,10 +296,15 @@ def stamp_clean_chain( slot_indices: list[int], tokens: list[int], positions: list[int], + real_kv_hashes: Optional[list[int]] = None, ) -> list[int]: + n = len(tokens) + real_kv_hashes = real_kv_hashes if real_kv_hashes is not None else [0] * n running_prev_hash = splitmix64(consts.CANARY_CHAIN_ANCHOR) stored_prev_hashes: list[int] = [] - for slot_idx, token, position in zip(slot_indices, tokens, positions): + for slot_idx, token, position, real_kv_hash in zip( + slot_indices, tokens, positions, real_kv_hashes + ): signed_prev = to_signed_int64(running_prev_hash) for buf in (cuda_buf, ref_buf): write_slot_fields( @@ -300,6 +313,7 @@ def stamp_clean_chain( token=token, position=position, prev_hash=signed_prev, + real_kv_hash=to_signed_int64(real_kv_hash), ) stored_prev_hashes.append(signed_prev) running_prev_hash = splitmix64_mix3(running_prev_hash, token, position) diff --git a/python/sglang/jit_kernel/tests/kv_canary/_differential.py b/python/sglang/jit_kernel/tests/kv_canary/_differential.py index b05619777..8c7c422fc 100644 --- a/python/sglang/jit_kernel/tests/kv_canary/_differential.py +++ b/python/sglang/jit_kernel/tests/kv_canary/_differential.py @@ -5,12 +5,14 @@ from typing import Any, Callable, Iterator, Optional import torch +from sglang.jit_kernel.kv_canary import consts from sglang.jit_kernel.kv_canary.plan import launch_canary_plan_kernels from sglang.jit_kernel.kv_canary.plan_ref import ( launch_canary_plan_kernels_torch_reference, ) from sglang.jit_kernel.kv_canary.verify import ( CanaryLaunchTag, + RealKvSource, VerifyOrWriteContext, VerifyPlan, launch_canary_verify_kernel, @@ -186,6 +188,9 @@ def _run_both_verify( plan_ref, cuda_log: FakeViolationLog, ref_log: FakeViolationLog, + real_kv_sources_cuda: tuple[RealKvSource, ...], + real_kv_sources_ref: tuple[RealKvSource, ...], + real_kv_hash_mode: consts.RealKvHashMode, kernel_kind: CanaryLaunchTag = CanaryLaunchTag.HEAD_K_FULL, assert_equal: bool = True, check_verify_expected_token: bool = True, @@ -199,6 +204,8 @@ def _run_both_verify( slot_run_counter=cuda_log.slot_run_counter, kernel_run_counter=cuda_log.kernel_run_counter, enable_chain_position_assert=cuda_log.enable_chain_position_assert, + real_kv_sources=real_kv_sources_cuda, + real_kv_hash_mode=real_kv_hash_mode, ), plan=plan_cuda, check_verify_expected_token=check_verify_expected_token, @@ -212,6 +219,8 @@ def _run_both_verify( slot_run_counter=ref_log.slot_run_counter, kernel_run_counter=ref_log.kernel_run_counter, enable_chain_position_assert=ref_log.enable_chain_position_assert, + real_kv_sources=real_kv_sources_ref, + real_kv_hash_mode=real_kv_hash_mode, ), plan=plan_ref, check_verify_expected_token=check_verify_expected_token, @@ -236,6 +245,9 @@ def _run_both_write( expected_input_positions: torch.Tensor, cuda_log: FakeViolationLog, ref_log: FakeViolationLog, + real_kv_sources_cuda: tuple[RealKvSource, ...], + real_kv_sources_ref: tuple[RealKvSource, ...], + real_kv_hash_mode: consts.RealKvHashMode, kernel_kind: CanaryLaunchTag = CanaryLaunchTag.HEAD_K_FULL, assert_equal: bool = True, ) -> None: @@ -254,6 +266,8 @@ def _run_both_write( slot_run_counter=cuda_log.slot_run_counter, kernel_run_counter=cuda_log.kernel_run_counter, enable_chain_position_assert=cuda_log.enable_chain_position_assert, + real_kv_sources=real_kv_sources_cuda, + real_kv_hash_mode=real_kv_hash_mode, ), plan=plan_cuda, input_ids=input_ids, @@ -272,6 +286,8 @@ def _run_both_write( slot_run_counter=ref_log.slot_run_counter, kernel_run_counter=ref_log.kernel_run_counter, enable_chain_position_assert=ref_log.enable_chain_position_assert, + real_kv_sources=real_kv_sources_ref, + real_kv_hash_mode=real_kv_hash_mode, ), plan=plan_ref, input_ids=input_ids, @@ -381,6 +397,20 @@ def _yield_simpler(inputs: Any) -> Iterator[tuple[str, Any]]: if fields["extras_count"] > 0: yield from emit("extras_zero", extras_count=0) + if "real_kv_hash_mode" in fields: + cur = fields["real_kv_hash_mode"] + if hasattr(cur, "value"): + cls = cur.__class__ + if int(cur) == 2: + yield from emit("hash_mode_bit", real_kv_hash_mode=cls(1)) + elif int(cur) == 1: + yield from emit("hash_mode_off", real_kv_hash_mode=cls(0)) + + if "real_kv_sources" in fields: + srcs = fields["real_kv_sources"] + if isinstance(srcs, tuple) and len(srcs) > 1: + yield from emit("sources_to_one", real_kv_sources=srcs[:1]) + if "enable_write_verify_inputs" in fields: cur = fields["enable_write_verify_inputs"] if hasattr(cur, "value") and int(cur) != 0: @@ -398,13 +428,18 @@ def run_verify_diff( *, buf_pair: tuple[torch.Tensor, torch.Tensor], plan_pair: tuple[VerifyPlan, VerifyPlan], + real_kv_sources_pair: tuple[tuple[RealKvSource, ...], tuple[RealKvSource, ...]] = ( + (), + (), + ), + real_kv_hash_mode: consts.RealKvHashMode = consts.RealKvHashMode.NONE, kernel_kind: CanaryLaunchTag = CanaryLaunchTag.HEAD_K_FULL, device: torch.device = _DEVICE, assert_equal: bool = True, check_verify_expected_token: bool = True, ) -> tuple[FakeViolationLog, FakeViolationLog]: """Thin wrapper around ``_run_both_verify`` that creates a fresh log pair and packs (cuda, ref) - buf/plan arguments into 2-tuples to drop ~8 lines of boilerplate per call site. + buf/plan/source arguments into 2-tuples to drop ~8 lines of boilerplate per call site. """ cuda_log, ref_log = make_log_pair(device=device) _run_both_verify( @@ -414,6 +449,9 @@ def run_verify_diff( plan_ref=plan_pair[1], cuda_log=cuda_log, ref_log=ref_log, + real_kv_sources_cuda=real_kv_sources_pair[0], + real_kv_sources_ref=real_kv_sources_pair[1], + real_kv_hash_mode=real_kv_hash_mode, kernel_kind=kernel_kind, assert_equal=assert_equal, check_verify_expected_token=check_verify_expected_token, @@ -431,12 +469,17 @@ def run_write_diff( expected_input_tokens: torch.Tensor, expected_input_positions: torch.Tensor, enable_write_verify_inputs: bool = False, + real_kv_sources_pair: tuple[tuple[RealKvSource, ...], tuple[RealKvSource, ...]] = ( + (), + (), + ), + real_kv_hash_mode: consts.RealKvHashMode = consts.RealKvHashMode.NONE, kernel_kind: CanaryLaunchTag = CanaryLaunchTag.HEAD_K_FULL, device: torch.device = _DEVICE, assert_equal: bool = True, ) -> tuple[FakeViolationLog, FakeViolationLog]: """Thin wrapper around ``_run_both_write`` that creates a fresh log pair and packs (cuda, ref) - buf/plan arguments into 2-tuples to drop ~10 lines of boilerplate per call site. + buf/plan/source arguments into 2-tuples to drop ~10 lines of boilerplate per call site. """ cuda_log, ref_log = make_log_pair(device=device) _run_both_write( @@ -452,6 +495,9 @@ def run_write_diff( expected_input_positions=expected_input_positions, cuda_log=cuda_log, ref_log=ref_log, + real_kv_sources_cuda=real_kv_sources_pair[0], + real_kv_sources_ref=real_kv_sources_pair[1], + real_kv_hash_mode=real_kv_hash_mode, kernel_kind=kernel_kind, assert_equal=assert_equal, ) diff --git a/python/sglang/jit_kernel/tests/kv_canary/_fixtures.py b/python/sglang/jit_kernel/tests/kv_canary/_fixtures.py index 60d8ac46e..98d6d21e1 100644 --- a/python/sglang/jit_kernel/tests/kv_canary/_fixtures.py +++ b/python/sglang/jit_kernel/tests/kv_canary/_fixtures.py @@ -5,8 +5,12 @@ from typing import Literal, Optional import torch -from sglang.jit_kernel.kv_canary.verify import VerifyPlan +from sglang.jit_kernel.kv_canary.verify import ( + RealKvSource, + VerifyPlan, +) from sglang.jit_kernel.kv_canary.write import WritePlan +from sglang.jit_kernel.tests.kv_canary._constants import DEFAULT_NUM_SLOTS _DEVICE = torch.device("cuda") @@ -72,6 +76,83 @@ def make_req_to_token( return rtt.contiguous() +def make_real_kv_source( + *, + num_slots: int = DEFAULT_NUM_SLOTS, + num_bytes_per_token: int = 16, + page_size: int = 1, + read_bytes: Optional[int] = None, + pad_dim1: int = 0, + device: torch.device, + fill: int = 0, +) -> RealKvSource: + """Allocate one RealKvSource with the canonical [num_rows, dim1_bytes] uint8 shape. + + ``pad_dim1`` adds trailing per-row bytes the canary should skip — used by the "holey dim 1" case to + confirm the kernel never reads past ``page_size * num_bytes_per_token``. + """ + num_rows = (num_slots + page_size - 1) // page_size + cols = page_size * num_bytes_per_token + pad_dim1 + tensor = torch.full( + (num_rows, cols), fill_value=fill, dtype=torch.uint8, device=device + ) + effective_read = read_bytes if read_bytes is not None else num_bytes_per_token + return RealKvSource( + tensor=tensor, + page_size=page_size, + num_bytes_per_token=num_bytes_per_token, + read_bytes=effective_read, + ) + + +FillStrategy = Literal["constant_per_source", "random_bytes"] + + +def make_real_kv_sources( + *, + count: int, + num_bytes_per_token: int = 16, + page_size: int = 1, + num_slots: int = DEFAULT_NUM_SLOTS, + device: torch.device, + rng: Optional[random.Random] = None, + fill_strategy: FillStrategy = "constant_per_source", +) -> tuple[RealKvSource, ...]: + sources: list[RealKvSource] = [] + for i in range(count): + read_bytes_eff = num_bytes_per_token + src = make_real_kv_source( + num_slots=num_slots, + num_bytes_per_token=num_bytes_per_token, + page_size=page_size, + read_bytes=read_bytes_eff, + device=device, + fill=(i + 1) * 17, + ) + if fill_strategy == "random_bytes": + if rng is None: + rng = random.Random(0) + seed = rng.randint(0, 0xFFFFFFFF) + gen = torch.Generator(device=device).manual_seed(seed) + src.tensor.random_(generator=gen) + sources.append(src) + return tuple(sources) + + +def clone_real_kv_sources( + sources: tuple[RealKvSource, ...], +) -> tuple[RealKvSource, ...]: + return tuple( + RealKvSource( + tensor=src.tensor.clone(), + page_size=src.page_size, + num_bytes_per_token=src.num_bytes_per_token, + read_bytes=src.read_bytes, + ) + for src in sources + ) + + PaddingKind = Literal["none", "trailing", "interleaved"] diff --git a/python/sglang/jit_kernel/tests/kv_canary/_hand_oracle.py b/python/sglang/jit_kernel/tests/kv_canary/_hand_oracle.py new file mode 100644 index 000000000..2f9487845 --- /dev/null +++ b/python/sglang/jit_kernel/tests/kv_canary/_hand_oracle.py @@ -0,0 +1,31 @@ +"""Hand-computed Python re-implementation of the real-kv-source fold, kept independent from +``verify_ref._splitmix64_fold_bytes_scalar`` so a ref / kernel co-regression cannot silently fix the +diff comparison.""" + +from __future__ import annotations + +from sglang.jit_kernel.kv_canary.consts import splitmix64 + + +def _fold_words(padded: bytes) -> int: + """Pack padded bytes little-endian into 8-byte words, fold each via splitmix64 from acc=0.""" + num_words = len(padded) // 8 + acc = 0 + for w in range(num_words): + chunk = padded[w * 8 : (w + 1) * 8] + word = sum(b << (8 * k) for k, b in enumerate(chunk)) + acc = splitmix64(acc ^ word) + return splitmix64(0 ^ acc) + + +def _hand_fold_partial(raw_bytes: bytes) -> int: + """PARTIAL-mode fold: first min(16, len) bytes, little-endian word-pack + splitmix64, same as ALL.""" + truncated = raw_bytes[: min(16, len(raw_bytes))] + pad = (8 - len(truncated) % 8) % 8 + return _fold_words(bytes(truncated) + bytes(pad)) + + +def _hand_fold_all(raw_bytes: bytes) -> int: + """ALL-mode fold: pack bytes little-endian into 8-byte words, fold each via splitmix64, then mix into acc=0.""" + pad = (8 - len(raw_bytes) % 8) % 8 + return _fold_words(raw_bytes + bytes(pad)) 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 index ce8d6957a..c5fd4799a 100644 --- a/python/sglang/jit_kernel/tests/kv_canary/test_const_sync.py +++ b/python/sglang/jit_kernel/tests/kv_canary/test_const_sync.py @@ -52,7 +52,7 @@ def test_int_consts_sync() -> None: def test_enums_sync() -> None: cuh = _CONSTS_CUH.read_text(encoding="utf-8") - for enum_name in ("FailReason",): + for enum_name in ("RealKvHashMode", "FailReason"): cpp_members = _parse_enum_class(cuh, enum_name) py_enum = getattr(consts, enum_name) cpp_normalized = { diff --git a/python/sglang/jit_kernel/tests/kv_canary/test_kernel_config.py b/python/sglang/jit_kernel/tests/kv_canary/test_kernel_config.py index df2c1021a..3f357e0b9 100644 --- a/python/sglang/jit_kernel/tests/kv_canary/test_kernel_config.py +++ b/python/sglang/jit_kernel/tests/kv_canary/test_kernel_config.py @@ -3,6 +3,7 @@ from __future__ import annotations import pytest import torch +from sglang.jit_kernel.kv_canary import consts from sglang.jit_kernel.kv_canary.verify import ( CanaryLaunchTag, VerifyOrWriteContext, @@ -119,6 +120,9 @@ def test_verify_byte_equal_across_repeated_launches_10x() -> None: plan_ref=plan_ref, cuda_log=cuda_log, ref_log=ref_log, + real_kv_sources_cuda=(), + real_kv_sources_ref=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, kernel_kind=CanaryLaunchTag.HEAD_K_FULL, ) @@ -171,6 +175,9 @@ def test_write_byte_equal_across_repeated_launches_10x() -> None: expected_input_positions=pseudo_pos, cuda_log=cuda_log, ref_log=ref_log, + real_kv_sources_cuda=(), + real_kv_sources_ref=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, kernel_kind=CanaryLaunchTag.HEAD_K_FULL, ) @@ -284,6 +291,8 @@ def test_verify_multi_launch_100x_counter_linear() -> None: slot_run_counter=cuda_log.slot_run_counter, kernel_run_counter=cuda_log.kernel_run_counter, enable_chain_position_assert=cuda_log.enable_chain_position_assert, + real_kv_sources=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, ), plan=plan_cuda, check_verify_expected_token=True, @@ -339,6 +348,9 @@ def test_verify_check_disabled_byte_equal() -> None: plan_ref=plan_true_ref, cuda_log=cuda_log_true, ref_log=ref_log_true, + real_kv_sources_cuda=(), + real_kv_sources_ref=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, kernel_kind=CanaryLaunchTag.HEAD_K_FULL, check_verify_expected_token=True, ) @@ -349,6 +361,9 @@ def test_verify_check_disabled_byte_equal() -> None: plan_ref=plan_false_ref, cuda_log=cuda_log_false, ref_log=ref_log_false, + real_kv_sources_cuda=(), + real_kv_sources_ref=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, kernel_kind=CanaryLaunchTag.HEAD_K_FULL, check_verify_expected_token=False, ) diff --git a/python/sglang/jit_kernel/tests/kv_canary/test_pipeline_e2e.py b/python/sglang/jit_kernel/tests/kv_canary/test_pipeline_e2e.py index 56478806a..df5d2ff72 100644 --- a/python/sglang/jit_kernel/tests/kv_canary/test_pipeline_e2e.py +++ b/python/sglang/jit_kernel/tests/kv_canary/test_pipeline_e2e.py @@ -12,6 +12,7 @@ from sglang.jit_kernel.kv_canary.plan_ref import ( ) from sglang.jit_kernel.kv_canary.verify import ( CanaryLaunchTag, + RealKvSource, VerifyOrWriteContext, VerifyPlan, launch_canary_verify_kernel, @@ -28,10 +29,12 @@ from sglang.jit_kernel.tests.kv_canary._canary_helpers import ( assert_canary_buf_equal, assert_canary_state_equal, make_canary_buf, + make_real_kv_sources, stamp_clean_chain, write_slot_fields, ) from sglang.jit_kernel.tests.kv_canary._fixtures import ( + clone_real_kv_sources, empty_extras, make_req_to_token, ) @@ -62,6 +65,8 @@ def _run_pipeline( enable_write_verify_inputs: bool, expected_input_tokens: torch.Tensor, expected_input_positions: torch.Tensor, + real_kv_sources: tuple[RealKvSource, ...], + real_kv_hash_mode: consts.RealKvHashMode, verify_capacity: int, write_req_capacity: int, req_to_verify_expected_tokens: Optional[torch.Tensor] = None, @@ -115,6 +120,8 @@ def _run_pipeline( slot_run_counter=log.slot_run_counter, kernel_run_counter=log.kernel_run_counter, enable_chain_position_assert=log.enable_chain_position_assert, + real_kv_sources=real_kv_sources, + real_kv_hash_mode=real_kv_hash_mode, ) launch_canary_write_kernel( context=context, @@ -142,6 +149,8 @@ def _run_pipeline( slot_run_counter=log.slot_run_counter, kernel_run_counter=log.kernel_run_counter, enable_chain_position_assert=log.enable_chain_position_assert, + real_kv_sources=real_kv_sources, + real_kv_hash_mode=real_kv_hash_mode, ), plan=plan_w, input_ids=input_ids, @@ -160,6 +169,8 @@ def _run_pipeline( slot_run_counter=log.slot_run_counter, kernel_run_counter=log.kernel_run_counter, enable_chain_position_assert=log.enable_chain_position_assert, + real_kv_sources=real_kv_sources, + real_kv_hash_mode=real_kv_hash_mode, ), plan=plan_v, check_verify_expected_token=check_verify_expected_token, @@ -185,6 +196,9 @@ def _run_both_and_assert_pipeline_equal( enable_write_verify_inputs: bool = False, expected_input_tokens: Optional[torch.Tensor] = None, expected_input_positions: Optional[torch.Tensor] = None, + real_kv_sources_real: tuple[RealKvSource, ...] = (), + real_kv_sources_ref: tuple[RealKvSource, ...] = (), + real_kv_hash_mode: consts.RealKvHashMode = consts.RealKvHashMode.NONE, ring_capacity: int = 64, verify_capacity: int = 256, write_req_capacity: int = 16, @@ -240,6 +254,7 @@ def _run_both_and_assert_pipeline_equal( enable_write_verify_inputs=enable_write_verify_inputs, expected_input_tokens=expected_input_tokens, expected_input_positions=expected_input_positions, + real_kv_hash_mode=real_kv_hash_mode, verify_capacity=verify_capacity, write_req_capacity=write_req_capacity, req_to_verify_expected_tokens=req_to_verify_expected_tokens, @@ -251,12 +266,14 @@ def _run_both_and_assert_pipeline_equal( real=True, canary_buf=buf_real, log=log_real, + real_kv_sources=real_kv_sources_real, **shared, ) plan_v_ref, plan_w_ref = _run_pipeline( real=False, canary_buf=buf_ref, log=log_ref, + real_kv_sources=real_kv_sources_ref, **shared, ) @@ -436,6 +453,35 @@ def test_pipeline_sweep_no_write() -> None: assert int(log_ref.slot_run_counter[0].item()) == prefix_len +@pytest.mark.parametrize( + "real_kv_hash_mode", + [ + consts.RealKvHashMode.NONE, + consts.RealKvHashMode.PARTIAL, + consts.RealKvHashMode.ALL, + ], +) +def test_pipeline_real_kv_mode(real_kv_hash_mode: consts.RealKvHashMode) -> None: + """real_kv_hash_mode OFF/PARTIAL/ALL: real and ref use cloned sources to prevent ALL-mode hash aliasing.""" + sources_real = make_real_kv_sources(count=2, num_slots=64, device=_DEVICE) + sources_ref = clone_real_kv_sources(sources_real) + + _run_both_and_assert_pipeline_equal( + req_pool_indices=_t([1]), + prefix_lens=_t([0]), + extend_seq_lens=_t([3]), + input_ids=_t([5, 6, 7]), + positions=_t([0, 1, 2]), + out_cache_loc=_contiguous_out_cache_loc(req_pool_idx=1, start=0, count=3), + req_to_token=_linear_r2t(), + num_slots=64, + extras=empty_extras(), + real_kv_sources_real=sources_real, + real_kv_sources_ref=sources_ref, + real_kv_hash_mode=real_kv_hash_mode, + ) + + def test_pipeline_pseudo_mode_on_match() -> None: """enable_write_verify_inputs=ON, expected==actual: zero violations, buf byte-equal.""" input_ids = _t([1, 2, 3, 4]) @@ -560,6 +606,7 @@ def test_pipeline_ring_overflow_via_real_plan() -> None: token=slot_idx + 1, position=slot_idx, prev_hash=0x1234_DEAD_BEEF_0000 + slot_idx, + real_kv_hash=0, ) # Step 2: run real pipeline (plan + no write + verify); overflow ring capacity=4 with all n_slots violations. @@ -609,6 +656,8 @@ def test_pipeline_ring_overflow_via_real_plan() -> None: slot_run_counter=log_real.slot_run_counter, kernel_run_counter=log_real.kernel_run_counter, enable_chain_position_assert=log_real.enable_chain_position_assert, + real_kv_sources=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, ), plan=plan_v_real, check_verify_expected_token=True, @@ -624,6 +673,8 @@ def test_pipeline_ring_overflow_via_real_plan() -> None: slot_run_counter=log_ref.slot_run_counter, kernel_run_counter=log_ref.kernel_run_counter, enable_chain_position_assert=log_ref.enable_chain_position_assert, + real_kv_sources=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, ), plan=plan_v_ref, check_verify_expected_token=True, @@ -649,6 +700,7 @@ def test_pipeline_kernel_kind_propagates(kernel_kind: CanaryLaunchTag) -> None: token=7, position=99, prev_hash=0, + real_kv_hash=0, ) _, _, log_real, log_ref, _, _, _, _ = _run_both_and_assert_pipeline_equal( diff --git a/python/sglang/jit_kernel/tests/kv_canary/test_utils.py b/python/sglang/jit_kernel/tests/kv_canary/test_utils.py index 429cf2d26..83023f664 100644 --- a/python/sglang/jit_kernel/tests/kv_canary/test_utils.py +++ b/python/sglang/jit_kernel/tests/kv_canary/test_utils.py @@ -50,6 +50,8 @@ def test_cases_to_x_vals_includes_scenario_axis() -> None: case.mode, case.extend_len, case.pool_kind, + case.real_kv_kind, + case.hash_mode, ) ] diff --git a/python/sglang/jit_kernel/tests/kv_canary/test_verify_fuzz.py b/python/sglang/jit_kernel/tests/kv_canary/test_verify_fuzz.py index 22ecebd65..729f23391 100644 --- a/python/sglang/jit_kernel/tests/kv_canary/test_verify_fuzz.py +++ b/python/sglang/jit_kernel/tests/kv_canary/test_verify_fuzz.py @@ -9,6 +9,7 @@ import torch from sglang.jit_kernel.kv_canary import consts from sglang.jit_kernel.kv_canary.verify import ( CanaryLaunchTag, + RealKvSource, VerifyPlan, ) from sglang.jit_kernel.tests.kv_canary._canary_helpers import ( @@ -19,6 +20,10 @@ from sglang.jit_kernel.tests.kv_canary._canary_helpers import ( stamp_clean_chain, ) from sglang.jit_kernel.tests.kv_canary._differential import _run_both_verify +from sglang.jit_kernel.tests.kv_canary._fixtures import ( + clone_real_kv_sources, + make_real_kv_sources, +) from sglang.jit_kernel.tests.kv_canary._fuzz_driver import ( FUZZ_SEEDS_PR, run_fuzz_combo, @@ -41,16 +46,39 @@ class VerifyFuzzInputs: plan_cuda: VerifyPlan plan_ref: VerifyPlan kernel_kind: CanaryLaunchTag + real_kv_sources_cuda: tuple[RealKvSource, ...] + real_kv_sources_ref: tuple[RealKvSource, ...] + real_kv_hash_mode: consts.RealKvHashMode ring_capacity: int check_verify_expected_token: bool def _draw_random_verify_inputs(rng: random.Random) -> VerifyFuzzInputs: + hash_mode = rng.choice( + [ + consts.RealKvHashMode.NONE, + consts.RealKvHashMode.PARTIAL, + consts.RealKvHashMode.ALL, + ] + ) + src_count = rng.choice([1, 2, 4]) + page_size = rng.choice([1, 16]) + bytes_per = rng.choice([16, 64, 128]) kernel_kind = rng.choice(list(CanaryLaunchTag)) plan_size = rng.randint(0, 32) num_slots = max(plan_size + 8, 16) ring_capacity = rng.choice([16, 64, 256]) + sources_cuda = make_real_kv_sources( + count=src_count, + num_bytes_per_token=bytes_per, + page_size=page_size, + num_slots=num_slots, + device=_DEVICE, + rng=rng, + ) + sources_ref = clone_real_kv_sources(sources_cuda) + cuda_buf = make_canary_buf( num_slots=num_slots, slot_stride_bytes=32, device=_DEVICE ) @@ -68,7 +96,7 @@ def _draw_random_verify_inputs(rng: random.Random) -> VerifyFuzzInputs: else: prev_slot_indices.append(slot_indices[i - 1]) - if plan_size > 0: + if hash_mode == consts.RealKvHashMode.NONE and plan_size > 0: stamp_clean_chain( cuda_buf=cuda_buf, ref_buf=ref_buf, @@ -112,6 +140,9 @@ def _draw_random_verify_inputs(rng: random.Random) -> VerifyFuzzInputs: plan_cuda=plan_cuda, plan_ref=plan_ref, kernel_kind=kernel_kind, + real_kv_sources_cuda=sources_cuda, + real_kv_sources_ref=sources_ref, + real_kv_hash_mode=hash_mode, ring_capacity=ring_capacity, check_verify_expected_token=check_verify_expected_token, ) @@ -130,6 +161,9 @@ def _run_one(inputs: VerifyFuzzInputs) -> None: plan_ref=inputs.plan_ref, cuda_log=cuda_log, ref_log=ref_log, + real_kv_sources_cuda=inputs.real_kv_sources_cuda, + real_kv_sources_ref=inputs.real_kv_sources_ref, + real_kv_hash_mode=inputs.real_kv_hash_mode, kernel_kind=inputs.kernel_kind, assert_equal=False, check_verify_expected_token=inputs.check_verify_expected_token, @@ -155,6 +189,8 @@ def _summarize(inputs: VerifyFuzzInputs) -> str: n_active = int(inputs.plan_cuda.verify_num_valid[0].item()) return ( f"plan_size={n_active} kind={inputs.kernel_kind.name} " + f"hash_mode={inputs.real_kv_hash_mode.name} " + f"sources={len(inputs.real_kv_sources_cuda)} " f"ring={inputs.ring_capacity} " f"check_token={inputs.check_verify_expected_token}" ) @@ -162,7 +198,7 @@ def _summarize(inputs: VerifyFuzzInputs) -> str: @pytest.mark.parametrize("seed", FUZZ_SEEDS_PR) def test_verify_fuzz_full_combo(seed: int) -> None: - """Multi-dim verify fuzzer: random kernel kind × plan size × ring capacity × N iters, byte-equal.""" + """Multi-dim verify fuzzer: random hash mode × kernel kind × page × bytes × N iters, byte-equal.""" run_fuzz_combo( seed, draw_fn=_draw_random_verify_inputs, diff --git a/python/sglang/jit_kernel/tests/kv_canary/test_verify_hand.py b/python/sglang/jit_kernel/tests/kv_canary/test_verify_hand.py index 5fded210c..6e78f34a0 100644 --- a/python/sglang/jit_kernel/tests/kv_canary/test_verify_hand.py +++ b/python/sglang/jit_kernel/tests/kv_canary/test_verify_hand.py @@ -12,13 +12,18 @@ 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 ( CanaryLaunchTag, + RealKvSource, VerifyOrWriteContext, VerifyPlan, launch_canary_verify_kernel, ) from sglang.jit_kernel.kv_canary.verify_ref import ( + _compute_real_kv_hash_scalar, launch_canary_verify_kernel_torch_reference, ) +from sglang.jit_kernel.kv_canary.write_ref import ( + launch_canary_write_kernel_torch_reference, +) from sglang.jit_kernel.tests.kv_canary._canary_helpers import ( FakeViolationLog, assert_only_bits_set, @@ -26,8 +31,11 @@ from sglang.jit_kernel.tests.kv_canary._canary_helpers import ( make_canary_buf, make_canary_buf_pair, make_log_pair, + make_real_kv_source, + make_real_kv_sources, make_verify_plan, make_verify_plan_pair, + make_write_plan, read_slot_fields, stamp_clean_chain, stamp_pair, @@ -38,6 +46,11 @@ from sglang.jit_kernel.tests.kv_canary._differential import ( _run_both_verify, run_verify_diff, ) +from sglang.jit_kernel.tests.kv_canary._fixtures import clone_real_kv_sources +from sglang.jit_kernel.tests.kv_canary._hand_oracle import ( + _hand_fold_all, + _hand_fold_partial, +) from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large") @@ -66,6 +79,7 @@ def _stamp_head( token: int = 42, position: int = 0, prev_hash: int | None = None, + real_kv_hash: int = 0, ) -> None: """``stamp_pair`` with ``prev_hash`` defaulting to ``chain_anchor_signed()`` (the chain-head value).""" stamp_pair( @@ -74,6 +88,7 @@ def _stamp_head( token=token, position=position, prev_hash=chain_anchor_signed() if prev_hash is None else prev_hash, + real_kv_hash=real_kv_hash, ) @@ -114,7 +129,7 @@ def _run_both_verify_no_rkv( assert_equal: bool = True, kernel_kind: CanaryLaunchTag = CanaryLaunchTag.HEAD_K_FULL, ) -> None: - """``_run_both_verify`` — the most common in-place verify run.""" + """``_run_both_verify`` with empty real_kv sources / NONE mode — the most common in-place verify run.""" _run_both_verify( cuda_canary_buf=buf_pair[0], ref_canary_buf=buf_pair[1], @@ -122,6 +137,9 @@ def _run_both_verify_no_rkv( plan_ref=plan_pair[1], cuda_log=cuda_log, ref_log=ref_log, + real_kv_sources_cuda=(), + real_kv_sources_ref=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, kernel_kind=kernel_kind, assert_equal=assert_equal, ) @@ -132,6 +150,9 @@ class _VerifySingleSlotInput: token: int = 42 position: int = 0 stored_prev_hash_signed: int + stored_real_kv_hash_signed: int = 0 + real_kv_sources: tuple[RealKvSource, ...] = () + real_kv_hash_mode: consts.RealKvHashMode = consts.RealKvHashMode.NONE def _run_verify_single_slot_byte_equal(case: _VerifySingleSlotInput) -> None: @@ -142,14 +163,66 @@ def _run_verify_single_slot_byte_equal(case: _VerifySingleSlotInput) -> None: token=case.token, position=case.position, prev_hash=case.stored_prev_hash_signed, + real_kv_hash=case.stored_real_kv_hash_signed, ) + sources_cuda = case.real_kv_sources + sources_ref = clone_real_kv_sources(sources_cuda) plan_pair = _plan_pair_single(slot_idx=1, position=case.position) run_verify_diff( buf_pair=buf_pair, plan_pair=plan_pair, + real_kv_sources_pair=(sources_cuda, sources_ref), + real_kv_hash_mode=case.real_kv_hash_mode, ) +def _stamp_clean_kv_chain( + *, + buf_pair: tuple[torch.Tensor, torch.Tensor], + sources_cuda: tuple[RealKvSource, ...], + input_ids: torch.Tensor, + positions: torch.Tensor, + out_cache_loc: torch.Tensor, + real_kv_hash_mode: consts.RealKvHashMode, +) -> None: + """Use the Python write ref impl to populate the canary buf for a fresh chain. + + Lets verify tests start from a known-good chain without re-implementing splitmix64 by hand. + """ + n = int(input_ids.shape[0]) + cuda_buf, ref_buf = buf_pair + write_plan = make_write_plan( + write_offsets=[0, n], + seed_slot_indices=[-1], + num_valid_reqs=1, + device=_DEVICE, + ) + log = FakeViolationLog.allocate(device=_DEVICE) + # enable_write_input_assert=False is hard-wired here, so the kernel API requires the + # expected_* tensors be None (otherwise it raises ValueError). + launch_canary_write_kernel_torch_reference( + context=VerifyOrWriteContext( + canary_buf=cuda_buf, + kernel_kind=CanaryLaunchTag.HEAD_K_FULL, + violation_ring=log.ring, + violation_write_index=log.write_index, + slot_run_counter=log.slot_run_counter, + kernel_run_counter=log.kernel_run_counter, + enable_chain_position_assert=log.enable_chain_position_assert, + real_kv_sources=sources_cuda, + real_kv_hash_mode=real_kv_hash_mode, + ), + plan=write_plan, + input_ids=input_ids, + positions=positions, + out_cache_loc=out_cache_loc, + enable_write_input_assert=False, + expected_input_tokens=None, + expected_input_positions=None, + ) + ref_buf.copy_(cuda_buf) + + # --------------------------------------------------------------------------- # Kernel-contract invariants. # --------------------------------------------------------------------------- @@ -194,11 +267,12 @@ class TestChain: tokens = [101, 202, 303, 404, 505] positions = [0, 1, 2, 3, 4] slot_indices = [1, 2, 3, 4, 5] + real_kv_hashes = [0, 0, 0, 0, 0] # Step 1: compute the expected stored prev_hash sequence in pure Python via splitmix64. expected_prev_hashes_u64: list[int] = [] running = splitmix64(consts.CANARY_CHAIN_ANCHOR) - for token, position in zip(tokens, positions): + for token, position, real_kv_hash in zip(tokens, positions, real_kv_hashes): expected_prev_hashes_u64.append(running) running = splitmix64_mix3(running, token, position) expected_prev_hashes_signed = [ @@ -447,7 +521,40 @@ class TestViolationField: _fail_bits(cuda_log), consts.FailReason.VERIFY_CHAIN_HASH_MISMATCH ) - @pytest.mark.parametrize("bit_to_trigger", ["POSITION", "PREV_HASH"]) + def test_violation_real_kv_hash_mismatch(self) -> None: + """Mutate one byte of a RealKvSource tensor after writing the chain → REAL_KV_HASH bit on verify.""" + buf_pair = _buf_pair() + sources_cuda = make_real_kv_sources(count=1, device=_DEVICE) + + # Step: write a chain with real_kv_hash mixin, then mutate one byte in the source tensors so the next + # verify reconstructs a hash that differs from the stored one. + _stamp_clean_kv_chain( + buf_pair=buf_pair, + sources_cuda=sources_cuda, + input_ids=torch.tensor([7, 8, 9], dtype=torch.int64, device=_DEVICE), + positions=torch.tensor([0, 1, 2], dtype=torch.int64, device=_DEVICE), + out_cache_loc=torch.tensor([1, 2, 3], dtype=torch.int64, device=_DEVICE), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + + # Mutate one byte in BOTH copies so the verify recomputed hash diverges from stored. + sources_ref = clone_real_kv_sources(sources_cuda) + sources_cuda[0].tensor[1, 0] ^= 0xFF + sources_ref[0].tensor.copy_(sources_cuda[0].tensor) + + plan_pair = _plan_pair_single(slot_idx=1, position=0) + cuda_log, _ = run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair, + real_kv_sources_pair=(sources_cuda, sources_ref), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + + assert_only_bits_set( + _fail_bits(cuda_log), consts.FailReason.VERIFY_REAL_KV_HASH_MISMATCH + ) + + @pytest.mark.parametrize("bit_to_trigger", ["POSITION", "PREV_HASH", "REAL_KV"]) @pytest.mark.parametrize("injection_position", ["head", "mid", "last"]) @pytest.mark.parametrize("ring_state", ["open", "full"]) def test_violation_bit_injection_position_ring_state_matrix( @@ -468,41 +575,67 @@ class TestViolationField: expected_bit = { "POSITION": consts.FailReason.VERIFY_POSITION_MISMATCH, "PREV_HASH": consts.FailReason.VERIFY_CHAIN_HASH_MISMATCH, + "REAL_KV": consts.FailReason.VERIFY_REAL_KV_HASH_MISMATCH, }[bit_to_trigger] - cuda_buf, ref_buf = _buf_pair() - buf_pair = (cuda_buf, ref_buf) - stamp_clean_chain( - cuda_buf=cuda_buf, - ref_buf=ref_buf, - tokens=tokens, - positions=positions, - slot_indices=slot_indices, - ) - - if bit_to_trigger == "POSITION": - stored_token, stored_pos, stored_prev, _ = read_slot_fields( - canary_buf=cuda_buf, slot_idx=corrupt_slot - ) - stamp_pair( - buf_pair, - slot_idx=corrupt_slot, - token=stored_token, - position=stored_pos + 99, - prev_hash=stored_prev, + if bit_to_trigger == "REAL_KV": + buf_pair = _buf_pair() + sources_cuda = make_real_kv_sources(count=1, device=_DEVICE) + _stamp_clean_kv_chain( + buf_pair=buf_pair, + sources_cuda=sources_cuda, + input_ids=torch.tensor(tokens, dtype=torch.int64, device=_DEVICE), + positions=torch.tensor(positions, dtype=torch.int64, device=_DEVICE), + out_cache_loc=torch.tensor( + slot_indices, dtype=torch.int64, device=_DEVICE + ), + real_kv_hash_mode=consts.RealKvHashMode.ALL, ) + sources_ref = clone_real_kv_sources(sources_cuda) + sources_cuda[0].tensor[corrupt_slot, 0] ^= 0xFF + sources_ref[0].tensor.copy_(sources_cuda[0].tensor) + real_kv_hash_mode = consts.RealKvHashMode.ALL + real_kv_sources_cuda = sources_cuda + real_kv_sources_ref = sources_ref else: - stored_token, stored_pos, stored_prev, _ = read_slot_fields( - canary_buf=cuda_buf, slot_idx=corrupt_slot - ) - flipped_prev = stored_prev ^ 1 - stamp_pair( - buf_pair, - slot_idx=corrupt_slot, - token=stored_token, - position=stored_pos, - prev_hash=flipped_prev, + cuda_buf, ref_buf = _buf_pair() + buf_pair = (cuda_buf, ref_buf) + stamp_clean_chain( + cuda_buf=cuda_buf, + ref_buf=ref_buf, + tokens=tokens, + positions=positions, + slot_indices=slot_indices, ) + real_kv_hash_mode = consts.RealKvHashMode.NONE + real_kv_sources_cuda = () + real_kv_sources_ref = () + + if bit_to_trigger == "POSITION": + stored_token, stored_pos, stored_prev, stored_rkv = read_slot_fields( + canary_buf=cuda_buf, slot_idx=corrupt_slot + ) + stamp_pair( + buf_pair, + slot_idx=corrupt_slot, + token=stored_token, + position=stored_pos + 99, + prev_hash=stored_prev, + real_kv_hash=stored_rkv, + ) + else: + stored_token, stored_pos, stored_prev, stored_rkv = read_slot_fields( + canary_buf=cuda_buf, slot_idx=corrupt_slot + ) + flipped_prev = stored_prev ^ 1 + stamp_pair( + buf_pair, + slot_idx=corrupt_slot, + token=stored_token, + position=stored_pos, + prev_hash=flipped_prev, + real_kv_hash=stored_rkv, + ) ring_capacity = _RING_CAPACITY cuda_log, ref_log = make_log_pair(capacity=ring_capacity, device=_DEVICE) @@ -540,6 +673,9 @@ class TestViolationField: plan_ref=plan_ref, cuda_log=cuda_log, ref_log=ref_log, + real_kv_sources_cuda=real_kv_sources_cuda, + real_kv_sources_ref=real_kv_sources_ref, + real_kv_hash_mode=real_kv_hash_mode, assert_equal=False, ) @@ -577,7 +713,451 @@ class TestViolationField: ) == 0, f"chain hash bit unexpectedly set: {bits:#b}" +class TestRealKvHash: + def test_real_kv_mode_off_yields_zero(self) -> None: + """OFF mode → stored real_kv_hash field stays zero post-write; verify with OFF agrees byte-equal.""" + buf_pair = _buf_pair() + sources = make_real_kv_sources(count=2, device=_DEVICE) + + plan_pair = _plan_pair_single(slot_idx=1, position=0) + _stamp_head(buf_pair, slot_idx=1, token=1) + + cuda_log, _ = run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair, + real_kv_sources_pair=(sources, sources), + ) + + assert _n_violations(cuda_log) == 0 + + @pytest.mark.parametrize( + "mode", + [ + pytest.param(consts.RealKvHashMode.PARTIAL, id="partial"), + pytest.param(consts.RealKvHashMode.ALL, id="all"), + ], + ) + def test_real_kv_mode_byte_equal(self, mode: consts.RealKvHashMode) -> None: + """PARTIAL / ALL modes both produce CUDA-vs-ref byte-equal state on a clean 3-step chain.""" + buf_pair = _buf_pair() + sources_cuda = make_real_kv_sources(count=2, device=_DEVICE) + sources_ref = clone_real_kv_sources(sources_cuda) + + # Write a chain through the ref so both buffers are byte-equal post-write. + _stamp_clean_kv_chain( + buf_pair=buf_pair, + sources_cuda=sources_cuda, + input_ids=torch.tensor([10, 20, 30], dtype=torch.int64, device=_DEVICE), + positions=torch.tensor([0, 1, 2], dtype=torch.int64, device=_DEVICE), + out_cache_loc=torch.tensor([1, 2, 3], dtype=torch.int64, device=_DEVICE), + real_kv_hash_mode=mode, + ) + + plan_pair = make_verify_plan_pair( + slot_indices=[1, 2, 3], + positions=[0, 1, 2], + prev_slot_indices=[-1, 1, 2], + device=_DEVICE, + ) + cuda_log, _ = run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair, + real_kv_sources_pair=(sources_cuda, sources_ref), + real_kv_hash_mode=mode, + ) + + assert _n_violations(cuda_log) == 0 + + @pytest.mark.parametrize("count", [1, 2, 3, 4]) + def test_real_kv_sources_fold_1_to_4(self, count: int) -> None: + """Fold ``count`` sources sequentially → CUDA matches ref for every count in {1..4}.""" + buf_pair = _buf_pair() + sources_cuda = make_real_kv_sources(count=count, device=_DEVICE) + sources_ref = clone_real_kv_sources(sources_cuda) + + _stamp_clean_kv_chain( + buf_pair=buf_pair, + sources_cuda=sources_cuda, + input_ids=torch.tensor([1, 2], dtype=torch.int64, device=_DEVICE), + positions=torch.tensor([0, 1], dtype=torch.int64, device=_DEVICE), + out_cache_loc=torch.tensor([1, 2], dtype=torch.int64, device=_DEVICE), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + + plan_pair = make_verify_plan_pair( + slot_indices=[1, 2], + positions=[0, 1], + prev_slot_indices=[-1, 1], + device=_DEVICE, + ) + cuda_log, _ = run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair, + real_kv_sources_pair=(sources_cuda, sources_ref), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + + assert _n_violations(cuda_log) == 0 + + @pytest.mark.parametrize( + "mode,fold_fn,expected_hash", + [ + pytest.param( + consts.RealKvHashMode.PARTIAL, + _hand_fold_partial, + 0x6041580849E6407D, + id="partial", + ), + pytest.param( + consts.RealKvHashMode.ALL, + _hand_fold_all, + 0x6041580849E6407D, + id="all", + ), + ], + ) + def test_real_kv_hash_fold_mode_hardcoded( + self, + mode: consts.RealKvHashMode, + fold_fn: Callable[[bytes], int], + expected_hash: int, + ) -> None: + # Step 1: build one RealKvSource with read_bytes=16 and a fixed byte pattern at slot 1. + _PATTERN = bytes( + [ + 0x01, + 0x02, + 0x04, + 0x08, + 0x10, + 0x20, + 0x40, + 0x80, + 0x81, + 0x82, + 0x84, + 0x88, + 0x90, + 0xA0, + 0xC0, + 0xFF, + ] + ) + + buf_pair = _buf_pair() + source_cuda = make_real_kv_source( + num_slots=16, + num_bytes_per_token=16, + page_size=1, + read_bytes=16, + device=_DEVICE, + ) + source_cuda.tensor[1, :16] = torch.tensor(list(_PATTERN), dtype=torch.uint8) + source_ref = RealKvSource( + tensor=source_cuda.tensor.clone(), + page_size=source_cuda.page_size, + num_bytes_per_token=source_cuda.num_bytes_per_token, + read_bytes=source_cuda.read_bytes, + ) + + # Step 2: verify hand-computed fold matches the hex literal. + assert fold_fn(_PATTERN) == expected_hash + + # Step 3: stamp slot 1 with a chain-head entry whose real_kv_hash equals the expected value. + _stamp_head( + buf_pair, + slot_idx=1, + token=7, + real_kv_hash=to_signed_int64(expected_hash), + ) + + # Step 4: 1-entry verify plan; no violation because stored matches recomputed. + plan_pair = _plan_pair_single(slot_idx=1, position=0) + cuda_log, _ = run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair, + real_kv_sources_pair=((source_cuda,), (source_ref,)), + real_kv_hash_mode=mode, + assert_equal=False, + ) + + assert _n_violations(cuda_log) == 0 + + # Step 5: mutate one byte in the source so the recomputed hash diverges from stored. + source_cuda.tensor[1, 0] ^= 0xFF + source_ref.tensor.copy_(source_cuda.tensor) + + plan_pair2 = _plan_pair_single(slot_idx=1, position=0) + cuda_log2, _ = run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair2, + real_kv_sources_pair=((source_cuda,), (source_ref,)), + real_kv_hash_mode=mode, + ) + + assert_only_bits_set( + _fail_bits(cuda_log2), consts.FailReason.VERIFY_REAL_KV_HASH_MISMATCH + ) + + def test_real_kv_hash_all_mode_with_multiple_sources(self) -> None: + """ALL mode with count=2 page=16 bytes=128 sources: chain still verifies clean.""" + buf_pair = _buf_pair(num_slots=32) + sources_cuda = make_real_kv_sources( + count=2, + num_bytes_per_token=128, + page_size=16, + num_slots=32, + device=_DEVICE, + ) + sources_ref = clone_real_kv_sources(sources_cuda) + + slot_indices = [1, 2, 3] + tokens = [100, 200, 300] + positions = [0, 1, 2] + + running = splitmix64(consts.CANARY_CHAIN_ANCHOR) + real_kv_hashes: list[int] = [] + for slot_idx in slot_indices: + real_kv_hashes.append( + _compute_real_kv_hash_scalar( + real_kv_sources=sources_cuda, + real_kv_hash_mode=consts.RealKvHashMode.ALL, + slot_idx=slot_idx, + work_device=torch.device("cpu"), + ) + ) + + for slot_idx, token, position, rkv in zip( + slot_indices, tokens, positions, real_kv_hashes + ): + signed_prev = to_signed_int64(running) + stamp_pair( + buf_pair, + slot_idx=slot_idx, + token=token, + position=position, + prev_hash=signed_prev, + real_kv_hash=to_signed_int64(rkv), + ) + running = splitmix64_mix3(running, token, position) + + plan_pair = make_verify_plan_pair( + slot_indices=slot_indices, + positions=positions, + prev_slot_indices=[-1, 1, 2], + device=_DEVICE, + ) + cuda_log, _ = run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair, + real_kv_sources_pair=(sources_cuda, sources_ref), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + assert _n_violations(cuda_log) == 0 + + def test_real_kv_hash_partial_mode_detects_single_bit_flip(self) -> None: + """PARTIAL mode + 1-bit flip in source tensor → REAL_KV_HASH bit set in violation row.""" + buf_pair = _buf_pair() + sources_cuda = make_real_kv_sources( + count=1, num_bytes_per_token=16, device=_DEVICE + ) + slot_idx = 3 + row_bytes = ( + sources_cuda[0] + .tensor[slot_idx, : sources_cuda[0].read_bytes] + .detach() + .cpu() + .tolist() + ) + rkv_clean = _hand_fold_partial(bytes(row_bytes)) + _stamp_head( + buf_pair, + slot_idx=slot_idx, + real_kv_hash=to_signed_int64(rkv_clean), + ) + + sources_cuda[0].tensor[slot_idx, 0] ^= 1 + sources_ref = clone_real_kv_sources(sources_cuda) + + plan_pair = _plan_pair_single(slot_idx=slot_idx, position=0) + cuda_log, _ = run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair, + real_kv_sources_pair=(sources_cuda, sources_ref), + real_kv_hash_mode=consts.RealKvHashMode.PARTIAL, + ) + + assert _n_violations(cuda_log) >= 1 + bits = _fail_bits(cuda_log) + assert ( + bits & consts.FailReason.VERIFY_REAL_KV_HASH_MISMATCH + ), f"expected REAL_KV_HASH bit, got {bits:#b}" + + def test_real_kv_off_does_not_deref_real_kv_sources(self) -> None: + buf_pair = _buf_pair(num_slots=8) + _stamp_head(buf_pair, slot_idx=1, token=1) + + garbage_source = make_real_kv_source( + num_slots=8, + num_bytes_per_token=16, + page_size=1, + read_bytes=16, + device=_DEVICE, + fill=0xDE, + ) + + plan_pair = _plan_pair_single(slot_idx=1, position=0) + cuda_log, ref_log = run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair, + real_kv_sources_pair=((garbage_source,), (garbage_source,)), + assert_equal=False, + ) + + assert _n_violations(cuda_log) == 0 + assert _n_violations(ref_log) == 0 + + +class TestRealKvSource: + def test_real_kv_source_rejects_zero_read_bytes(self) -> None: + """RealKvSource has no \"skip me\" sentinel — read_bytes=0 must raise rather than silently pass.""" + with pytest.raises(ValueError, match="read_bytes"): + RealKvSource( + tensor=torch.zeros((1, 16), dtype=torch.uint8, device=_DEVICE), + page_size=1, + num_bytes_per_token=16, + read_bytes=0, + ) + + def test_real_kv_source_padding_below_4(self) -> None: + """Host wrapper pads to 4 slots when fewer sources are supplied; dummy slots are never dereferenced.""" + buf_pair = _buf_pair() + sources = make_real_kv_sources(count=2, device=_DEVICE) + + plan_pair = _plan_pair_single(slot_idx=1, position=0) + _stamp_head(buf_pair, slot_idx=1, token=1) + run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair, + real_kv_sources_pair=(sources, sources), + ) + + def test_real_kv_source_above_4_raises(self) -> None: + """``len(real_kv_sources) > 4`` → host wrapper raises ValueError before launching.""" + canary_buf = make_canary_buf(device=_DEVICE) + plan = make_verify_plan( + slot_indices=[1], positions=[0], prev_slot_indices=[-1], device=_DEVICE + ) + log = FakeViolationLog.allocate(device=_DEVICE) + sources = make_real_kv_sources(count=4, device=_DEVICE) + extra = make_real_kv_source(device=_DEVICE) + too_many = sources + (extra,) + + with pytest.raises(ValueError, match="at most 4 RealKvSource"): + launch_canary_verify_kernel( + context=VerifyOrWriteContext( + canary_buf=canary_buf, + kernel_kind=CanaryLaunchTag.HEAD_K_FULL, + violation_ring=log.ring, + violation_write_index=log.write_index, + slot_run_counter=log.slot_run_counter, + kernel_run_counter=log.kernel_run_counter, + enable_chain_position_assert=log.enable_chain_position_assert, + real_kv_sources=too_many, + real_kv_hash_mode=consts.RealKvHashMode.NONE, + ), + plan=plan, + check_verify_expected_token=True, + ) + + def test_real_kv_source_holey_dim1(self) -> None: + """``tensor.shape[1] > page_size * num_bytes_per_token`` → trailing bytes are skipped.""" + buf_pair = _buf_pair() + holey_source = make_real_kv_source( + num_slots=16, + num_bytes_per_token=16, + page_size=1, + read_bytes=16, + pad_dim1=16, # 16 trailing pad bytes per row; must be skipped. + device=_DEVICE, + ) + trailing_start = holey_source.page_size * holey_source.num_bytes_per_token + # Fill those skipped trailing bytes with garbage; CUDA must not read them. + holey_source.tensor[:, trailing_start:].fill_(0xAA) + sources = (holey_source,) + sources_ref = clone_real_kv_sources(sources) + + _stamp_clean_kv_chain( + buf_pair=buf_pair, + sources_cuda=sources, + input_ids=torch.tensor([1, 2], dtype=torch.int64, device=_DEVICE), + positions=torch.tensor([0, 1], dtype=torch.int64, device=_DEVICE), + out_cache_loc=torch.tensor([1, 2], dtype=torch.int64, device=_DEVICE), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + + plan_pair = make_verify_plan_pair( + slot_indices=[1, 2], + positions=[0, 1], + prev_slot_indices=[-1, 1], + device=_DEVICE, + ) + cuda_log, _ = run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair, + real_kv_sources_pair=(sources, sources_ref), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + + assert _n_violations(cuda_log) == 0 + + class TestLayoutAndScheduling: + def test_page_size_gt_1_access_pattern(self) -> None: + """``page_size > 1`` → byte access follows ``(row=slot//page, col=(slot%page)*bpt:)``.""" + buf_pair = _buf_pair(num_slots=8) + src = make_real_kv_source( + num_slots=8, + num_bytes_per_token=16, + page_size=4, # 2 rows × 4 slots/page × 16 bytes/slot. + read_bytes=16, + device=_DEVICE, + ) + # Each slot's 16 bytes get a slot-specific signature so kernel mis-indexing would shift the hash. + flat = src.tensor.view(-1) + for slot_idx in range(8): + row = slot_idx // src.page_size + col = (slot_idx % src.page_size) * src.num_bytes_per_token + for k in range(src.num_bytes_per_token): + flat_index = row * (src.page_size * src.num_bytes_per_token) + col + k + flat[flat_index] = (slot_idx * 13 + k) & 0xFF + sources = (src,) + sources_ref = clone_real_kv_sources(sources) + + _stamp_clean_kv_chain( + buf_pair=buf_pair, + sources_cuda=sources, + input_ids=torch.tensor([1, 2], dtype=torch.int64, device=_DEVICE), + positions=torch.tensor([0, 1], dtype=torch.int64, device=_DEVICE), + out_cache_loc=torch.tensor([1, 5], dtype=torch.int64, device=_DEVICE), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + + plan_pair = make_verify_plan_pair( + slot_indices=[1, 5], + positions=[0, 1], + prev_slot_indices=[-1, 1], + device=_DEVICE, + ) + cuda_log, _ = run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair, + real_kv_sources_pair=(sources, sources_ref), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + + assert _n_violations(cuda_log) == 0 + def test_swa_translated_slot_indices(self) -> None: """SWA-translated slots already passed in plan; verify kernel does no further translation.""" # SWA verify plans carry pre-translated slot indices — the verify kernel never sees the FULL slot @@ -615,6 +1195,7 @@ class TestLayoutAndScheduling: token=999, position=123, prev_hash=to_signed_int64(0xDEADBEEF), + real_kv_hash=0, ) plan = make_verify_plan( slot_indices=[0], @@ -633,6 +1214,8 @@ class TestLayoutAndScheduling: slot_run_counter=log.slot_run_counter, kernel_run_counter=log.kernel_run_counter, enable_chain_position_assert=log.enable_chain_position_assert, + real_kv_sources=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, ), plan=plan, check_verify_expected_token=True, @@ -659,6 +1242,7 @@ class TestLayoutAndScheduling: token=42, position=1, prev_hash=to_signed_int64(0x1234), + real_kv_hash=0, ) plan = make_verify_plan( slot_indices=[slot_idx], @@ -687,6 +1271,8 @@ class TestLayoutAndScheduling: slot_run_counter=log.slot_run_counter, kernel_run_counter=log.kernel_run_counter, enable_chain_position_assert=log.enable_chain_position_assert, + real_kv_sources=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, ), plan=plan, check_verify_expected_token=True, @@ -702,6 +1288,65 @@ class TestLayoutAndScheduling: == int(kernel_run_before[0].item()) + 1 ) + def test_paged_layout_page_size_16(self) -> None: + """page_size=16: slot→page mapping doesn't change verify chain semantics on a clean chain.""" + buf_pair = _buf_pair(num_slots=64) + sources_cuda = make_real_kv_sources( + count=1, + num_bytes_per_token=16, + page_size=16, + num_slots=64, + device=_DEVICE, + ) + sources_ref = clone_real_kv_sources(sources_cuda) + + # Step: cross a page boundary by writing slots [15, 16] which straddle pages 0 and 1. + slot_indices = [15, 16] + tokens = [77, 88] + positions = [0, 1] + running = splitmix64(consts.CANARY_CHAIN_ANCHOR) + # Use the reference fold (8-byte little-endian word pack + splitmix64), not a + # byte-by-byte loop, so the stamped real_kv_hash matches what the kernel / + # verify reference will recompute. A byte-by-byte fold was the previous bug + # here and triggered REAL_KV_HASH violations on otherwise clean chains. + rkv_values = [ + _compute_real_kv_hash_scalar( + slot_idx=slot_idx, + real_kv_sources=sources_cuda, + real_kv_hash_mode=consts.RealKvHashMode.ALL, + work_device=_DEVICE, + ) + for slot_idx in slot_indices + ] + + for slot_idx, token, position, rkv in zip( + slot_indices, tokens, positions, rkv_values + ): + signed_prev = to_signed_int64(running) + stamp_pair( + buf_pair, + slot_idx=slot_idx, + token=token, + position=position, + prev_hash=signed_prev, + real_kv_hash=to_signed_int64(rkv), + ) + running = splitmix64_mix3(running, token, position) + + plan_pair = make_verify_plan_pair( + slot_indices=slot_indices, + positions=positions, + prev_slot_indices=[-1, 15], + device=_DEVICE, + ) + cuda_log, _ = run_verify_diff( + buf_pair=buf_pair, + plan_pair=plan_pair, + real_kv_sources_pair=(sources_cuda, sources_ref), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + assert _n_violations(cuda_log) == 0 + def test_empty_plan_keeps_slot_counter_unchanged(self) -> None: buf_pair = _buf_pair(num_slots=8) _stamp_head(buf_pair, slot_idx=1, token=7) @@ -1166,6 +1811,22 @@ class TestBoundarySweep: ) ) + @pytest.mark.parametrize( + "stored_rkv_val", + [0, 1, 0xFFFFFFFFFFFFFFFF, 0x8000000000000000], + ) + def test_real_kv_hash_boundary_byte_equal_sweep(self, stored_rkv_val: int) -> None: + """Sweep real_kv_hash boundary values; assert CUDA vs ref state byte-equal.""" + sources_cuda = make_real_kv_sources(count=1, device=_DEVICE) + _run_verify_single_slot_byte_equal( + _VerifySingleSlotInput( + stored_prev_hash_signed=chain_anchor_signed(), + stored_real_kv_hash_signed=to_signed_int64(stored_rkv_val), + real_kv_sources=sources_cuda, + real_kv_hash_mode=consts.RealKvHashMode.PARTIAL, + ) + ) + class TestVerifyExpectedInputIds: """Cover the new verify-time token-id check via VerifyPlan.verify_expected_tokens.""" diff --git a/python/sglang/jit_kernel/tests/kv_canary/test_write_fuzz.py b/python/sglang/jit_kernel/tests/kv_canary/test_write_fuzz.py index 6010796a5..3d56dc889 100644 --- a/python/sglang/jit_kernel/tests/kv_canary/test_write_fuzz.py +++ b/python/sglang/jit_kernel/tests/kv_canary/test_write_fuzz.py @@ -6,7 +6,11 @@ from dataclasses import dataclass import pytest import torch -from sglang.jit_kernel.kv_canary.verify import CanaryLaunchTag +from sglang.jit_kernel.kv_canary import consts +from sglang.jit_kernel.kv_canary.verify import ( + CanaryLaunchTag, + RealKvSource, +) from sglang.jit_kernel.kv_canary.write import WritePlan from sglang.jit_kernel.tests.kv_canary._canary_helpers import ( FakeViolationLog, @@ -16,6 +20,10 @@ from sglang.jit_kernel.tests.kv_canary._canary_helpers import ( stamp_pair, ) from sglang.jit_kernel.tests.kv_canary._differential import _run_both_write +from sglang.jit_kernel.tests.kv_canary._fixtures import ( + clone_real_kv_sources, + make_real_kv_sources, +) from sglang.jit_kernel.tests.kv_canary._fuzz_driver import ( FUZZ_SEEDS_PR, run_fuzz_combo, @@ -44,11 +52,24 @@ class WriteFuzzInputs: enable_write_verify_inputs: bool expected_input_tokens: torch.Tensor expected_input_positions: torch.Tensor + real_kv_sources_cuda: tuple[RealKvSource, ...] + real_kv_sources_ref: tuple[RealKvSource, ...] + real_kv_hash_mode: consts.RealKvHashMode ring_capacity: int def _draw_random_write_inputs(rng: random.Random) -> WriteFuzzInputs: enable_write_verify_inputs = rng.choice([False, True]) + hash_mode = rng.choice( + [ + consts.RealKvHashMode.NONE, + consts.RealKvHashMode.PARTIAL, + consts.RealKvHashMode.ALL, + ] + ) + src_count = rng.choice([1, 2, 4]) + page_size = rng.choice([1, 16]) + bytes_per = rng.choice([16, 64, 128]) kernel_kind = rng.choice(list(CanaryLaunchTag)) ring_capacity = rng.choice([16, 64, 256]) @@ -57,6 +78,16 @@ def _draw_random_write_inputs(rng: random.Random) -> WriteFuzzInputs: total_tokens = sum(per_req_tokens) num_slots = max(total_tokens + 8, 16) + sources_cuda = make_real_kv_sources( + count=src_count, + num_bytes_per_token=bytes_per, + page_size=page_size, + num_slots=num_slots, + device=_DEVICE, + rng=rng, + ) + sources_ref = clone_real_kv_sources(sources_cuda) + cuda_buf = make_canary_buf( num_slots=num_slots, slot_stride_bytes=32, device=_DEVICE ) @@ -140,6 +171,9 @@ def _draw_random_write_inputs(rng: random.Random) -> WriteFuzzInputs: enable_write_verify_inputs=enable_write_verify_inputs, expected_input_tokens=expected_input_tokens, expected_input_positions=expected_input_positions, + real_kv_sources_cuda=sources_cuda, + real_kv_sources_ref=sources_ref, + real_kv_hash_mode=hash_mode, ring_capacity=ring_capacity, ) @@ -163,6 +197,9 @@ def _run_one(inputs: WriteFuzzInputs) -> None: expected_input_positions=inputs.expected_input_positions, cuda_log=cuda_log, ref_log=ref_log, + real_kv_sources_cuda=inputs.real_kv_sources_cuda, + real_kv_sources_ref=inputs.real_kv_sources_ref, + real_kv_hash_mode=inputs.real_kv_hash_mode, kernel_kind=inputs.kernel_kind, assert_equal=False, ) @@ -196,13 +233,14 @@ def _summarize(inputs: WriteFuzzInputs) -> str: total = int(inputs.plan_cuda.write_offsets[n_active].item()) return ( f"n_reqs={n_active} total_tokens={total} kind={inputs.kernel_kind.name} " - f"pseudo={inputs.enable_write_verify_inputs}" + f"pseudo={inputs.enable_write_verify_inputs} hash_mode={inputs.real_kv_hash_mode.name} " + f"sources={len(inputs.real_kv_sources_cuda)}" ) @pytest.mark.parametrize("seed", FUZZ_SEEDS_PR) def test_write_fuzz_full_combo(seed: int) -> None: - """Multi-dim write fuzzer: random pseudo/kernel × N iters, byte-equal.""" + """Multi-dim write fuzzer: random pseudo/hash/kernel/page/source × N iters, byte-equal.""" run_fuzz_combo( seed, draw_fn=_draw_random_write_inputs, diff --git a/python/sglang/jit_kernel/tests/kv_canary/test_write_hand.py b/python/sglang/jit_kernel/tests/kv_canary/test_write_hand.py index a7c5689f2..b3fee829c 100644 --- a/python/sglang/jit_kernel/tests/kv_canary/test_write_hand.py +++ b/python/sglang/jit_kernel/tests/kv_canary/test_write_hand.py @@ -1,6 +1,7 @@ from __future__ import annotations from dataclasses import dataclass +from typing import Callable from unittest.mock import patch import pytest @@ -12,6 +13,7 @@ from sglang.jit_kernel.kv_canary.consts import splitmix64, splitmix64_mix3 from sglang.jit_kernel.kv_canary.verify import ( CANARY_SLOT_BYTES, CanaryLaunchTag, + RealKvSource, VerifyOrWriteContext, launch_canary_verify_kernel, ) @@ -26,6 +28,8 @@ from sglang.jit_kernel.tests.kv_canary._canary_helpers import ( make_canary_buf, make_canary_buf_pair, make_log_pair, + make_real_kv_source, + make_real_kv_sources, make_verify_plan, make_write_plan, make_write_plan_pair, @@ -38,8 +42,13 @@ from sglang.jit_kernel.tests.kv_canary._differential import ( run_write_diff, ) from sglang.jit_kernel.tests.kv_canary._fixtures import ( + clone_real_kv_sources, dummy_pseudo_tensors, ) +from sglang.jit_kernel.tests.kv_canary._hand_oracle import ( + _hand_fold_all, + _hand_fold_partial, +) from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large") @@ -73,6 +82,10 @@ def _run_write( enable_write_verify_inputs: bool = False, expected_input_tokens: torch.Tensor | None = None, expected_input_positions: torch.Tensor | None = None, + real_kv_sources_pair: ( + tuple[tuple[RealKvSource, ...], tuple[RealKvSource, ...]] | None + ) = None, + real_kv_hash_mode: consts.RealKvHashMode = consts.RealKvHashMode.NONE, assert_equal: bool = True, ) -> tuple[FakeViolationLog, FakeViolationLog]: """Shared scaffold: build write plan + pseudo tensors and call ``run_write_diff``. @@ -117,6 +130,9 @@ def _run_write( if expected_input_positions is None: expected_input_positions = pseudo_positions + extra_kwargs: dict = {} + if real_kv_sources_pair is not None: + extra_kwargs["real_kv_sources_pair"] = real_kv_sources_pair return run_write_diff( buf_pair=buf_pair, plan_pair=plan_pair, @@ -126,7 +142,9 @@ def _run_write( enable_write_verify_inputs=enable_write_verify_inputs, expected_input_tokens=expected_input_tokens, expected_input_positions=expected_input_positions, + real_kv_hash_mode=real_kv_hash_mode, assert_equal=assert_equal, + **extra_kwargs, ) @@ -135,6 +153,8 @@ class _WriteSingleSlotInput: token: int = 42 position: int = 0 enable_write_verify_inputs: bool = False + real_kv_sources: tuple[RealKvSource, ...] = () + real_kv_hash_mode: consts.RealKvHashMode = consts.RealKvHashMode.NONE class _RecordingWriteModule: @@ -146,12 +166,16 @@ class _RecordingWriteModule: def _run_write_single_slot_byte_equal(case: _WriteSingleSlotInput) -> None: + sources_cuda = case.real_kv_sources + sources_ref = clone_real_kv_sources(sources_cuda) _run_write( buf_pair=_make_default_buf_pair(), input_ids=[case.token], positions=[case.position], out_cache_loc=[0], enable_write_verify_inputs=case.enable_write_verify_inputs, + real_kv_sources_pair=(sources_cuda, sources_ref), + real_kv_hash_mode=case.real_kv_hash_mode, ) @@ -242,6 +266,8 @@ class TestSeedSlot: slot_run_counter=verify_log.slot_run_counter, kernel_run_counter=verify_log.kernel_run_counter, enable_chain_position_assert=verify_log.enable_chain_position_assert, + real_kv_sources=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, ), plan=verify_plan, check_verify_expected_token=True, @@ -270,11 +296,12 @@ class TestSeedSlot: tokens = [101, 202, 303, 404, 505] positions = [11, 12, 13, 14, 15] + real_kv = [0, 0, 0, 0, 0] out_cache_loc = [0, 1, 2, 3, 4] expected_prev_hashes: list[int] = [] running = predecessor_advance - for t, p in zip(tokens, positions): + for t, p, r in zip(tokens, positions, real_kv): expected_prev_hashes.append(running) running = splitmix64_mix3(running, t, p) @@ -305,6 +332,7 @@ class TestSeedSlot: seed_slot = 3 seed_token = 7 seed_position = 1 + seed_real_kv = 0 expected_seed_prev_hash = splitmix64(consts.CANARY_CHAIN_ANCHOR) stamp_pair( self.buf_pair, @@ -312,6 +340,7 @@ class TestSeedSlot: token=seed_token, position=seed_position, prev_hash=to_signed_int64(expected_seed_prev_hash), + real_kv_hash=to_signed_int64(seed_real_kv), ) new_slot = 4 @@ -355,11 +384,12 @@ class TestChain: tokens = [101, 202, 303, 404, 505] positions = [0, 1, 2, 3, 4] out_cache_loc = [0, 1, 2, 3, 4] + real_kv_hashes = [0, 0, 0, 0, 0] # Step 1: compute the expected stored prev_hash sequence in pure Python via splitmix64. expected_prev_hashes_u64: list[int] = [] running = splitmix64(consts.CANARY_CHAIN_ANCHOR) - for token, position in zip(tokens, positions): + for token, position, real_kv_hash in zip(tokens, positions, real_kv_hashes): expected_prev_hashes_u64.append(running) running = splitmix64_mix3(running, token, position) expected_prev_hashes_signed = [ @@ -385,6 +415,41 @@ class TestChain: assert stored_prev_hash == expected_prev_signed assert stored_real_kv_hash == 0 + def test_chain_advances_with_real_kv_hash_all(self) -> None: + """ALL mode + 2 sources + 5-step chain: stored prev_hash recoverable from seed.""" + cuda_buf = self.buf_pair[0] + sources_cuda = make_real_kv_sources( + count=2, + num_bytes_per_token=16, + page_size=1, + num_slots=16, + device=_DEVICE, + ) + sources_ref = clone_real_kv_sources(sources_cuda) + + slot_indices = [1, 2, 3, 4, 5] + tokens = [11, 22, 33, 44, 55] + positions = [0, 1, 2, 3, 4] + + _run_write( + buf_pair=self.buf_pair, + input_ids=tokens, + positions=positions, + out_cache_loc=slot_indices, + real_kv_sources_pair=(sources_cuda, sources_ref), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + + running = splitmix64(consts.CANARY_CHAIN_ANCHOR) + for slot_idx, token, position in zip(slot_indices, tokens, positions): + stored_prev_signed, stored_real_kv_hash = read_slot_fields( + canary_buf=cuda_buf, slot_idx=slot_idx + )[2:] + assert stored_prev_signed == to_signed_int64( + running + ), f"slot {slot_idx}: stored prev_hash != recomputed chain step" + running = splitmix64_mix3(running, token, position) + class TestMockMode: def setup_method(self) -> None: @@ -496,6 +561,8 @@ class TestMockMode: slot_run_counter=verify_log.slot_run_counter, kernel_run_counter=verify_log.kernel_run_counter, enable_chain_position_assert=verify_log.enable_chain_position_assert, + real_kv_sources=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, ), plan=verify_plan, check_verify_expected_token=True, @@ -686,6 +753,265 @@ class TestSlotHandling: ), f"slot {slot} from earlier bs=8 run was overwritten by bs=3 run" +class TestRealKvHash: + def setup_method(self) -> None: + self.buf_pair = _make_default_buf_pair() + + def test_real_kv_mode_off_writes_zero(self) -> None: + """``consts.RealKvHashMode.NONE`` → ``real_kv_hash`` field is written as 0 regardless of source presence.""" + sources = make_real_kv_sources(count=2, device=_DEVICE) + + _run_write( + buf_pair=self.buf_pair, + input_ids=[1, 2], + positions=[0, 1], + out_cache_loc=[0, 1], + real_kv_sources_pair=(sources, sources), + ) + + _, _, _, real_kv_0 = read_slot_fields(canary_buf=self.buf_pair[0], slot_idx=0) + _, _, _, real_kv_1 = read_slot_fields(canary_buf=self.buf_pair[0], slot_idx=1) + assert real_kv_0 == 0 + assert real_kv_1 == 0 + + @pytest.mark.parametrize( + "mode", + [ + pytest.param(consts.RealKvHashMode.PARTIAL, id="partial"), + pytest.param(consts.RealKvHashMode.ALL, id="all"), + ], + ) + def test_real_kv_mode_byte_equal(self, mode: consts.RealKvHashMode) -> None: + """PARTIAL / ALL modes both produce CUDA-vs-ref byte-equal write state on a 3-entry chain.""" + sources_cuda = make_real_kv_sources(count=2, device=_DEVICE) + sources_ref = clone_real_kv_sources(sources_cuda) + + _run_write( + buf_pair=self.buf_pair, + input_ids=[10, 20, 30], + positions=[0, 1, 2], + out_cache_loc=[0, 1, 2], + real_kv_sources_pair=(sources_cuda, sources_ref), + real_kv_hash_mode=mode, + ) + + @pytest.mark.parametrize("count", [1, 2, 3, 4]) + def test_real_kv_sources_fold_1_to_4(self, count: int) -> None: + """Folding ``count`` sources sequentially → CUDA matches ref for every count in {1..4}.""" + sources_cuda = make_real_kv_sources(count=count, device=_DEVICE) + sources_ref = clone_real_kv_sources(sources_cuda) + + _run_write( + buf_pair=self.buf_pair, + input_ids=[1, 2], + positions=[0, 1], + out_cache_loc=[0, 1], + real_kv_sources_pair=(sources_cuda, sources_ref), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + + def test_real_kv_source_above_4_raises(self) -> None: + """``len(real_kv_sources) > 4`` → host wrapper raises ValueError before launching.""" + cuda_buf = make_canary_buf(device=_DEVICE) + plan = make_write_plan( + write_offsets=[0, 1], + seed_slot_indices=[-1], + num_valid_reqs=1, + device=_DEVICE, + ) + input_ids = _int32_tensor([1]) + positions = _int32_tensor([0]) + out_cache_loc = _int32_tensor([0]) + log = FakeViolationLog.allocate(device=_DEVICE) + sources = make_real_kv_sources(count=4, device=_DEVICE) + extra = make_real_kv_source(device=_DEVICE) + too_many = sources + (extra,) + + with pytest.raises(ValueError, match="at most 4 RealKvSource"): + launch_canary_write_kernel( + context=VerifyOrWriteContext( + canary_buf=cuda_buf, + kernel_kind=CanaryLaunchTag.HEAD_K_FULL, + violation_ring=log.ring, + violation_write_index=log.write_index, + slot_run_counter=log.slot_run_counter, + kernel_run_counter=log.kernel_run_counter, + enable_chain_position_assert=log.enable_chain_position_assert, + real_kv_sources=too_many, + real_kv_hash_mode=consts.RealKvHashMode.NONE, + ), + plan=plan, + input_ids=input_ids, + positions=positions, + out_cache_loc=out_cache_loc, + enable_write_input_assert=False, + expected_input_tokens=None, + expected_input_positions=None, + ) + + @pytest.mark.parametrize( + "mode,fold_fn,expected_hash", + [ + pytest.param( + consts.RealKvHashMode.PARTIAL, + _hand_fold_partial, + 0x6041580849E6407D, + id="partial", + ), + pytest.param( + consts.RealKvHashMode.ALL, + _hand_fold_all, + 0x6041580849E6407D, + id="all", + ), + ], + ) + def test_real_kv_hash_fold_mode_writes_expected_hash_hardcoded( + self, + mode: consts.RealKvHashMode, + fold_fn: Callable[[bytes], int], + expected_hash: int, + ) -> None: + # Step 1: build one RealKvSource with read_bytes=16 and a fixed byte pattern at slot 0. + _PATTERN = bytes( + [ + 0x01, + 0x02, + 0x04, + 0x08, + 0x10, + 0x20, + 0x40, + 0x80, + 0x81, + 0x82, + 0x84, + 0x88, + 0x90, + 0xA0, + 0xC0, + 0xFF, + ] + ) + + # Step 2: verify hand-computed fold matches the hex literal. + assert fold_fn(_PATTERN) == expected_hash + + source_cuda = make_real_kv_source( + num_slots=16, + num_bytes_per_token=16, + page_size=1, + read_bytes=16, + device=_DEVICE, + ) + source_cuda.tensor[0, :16] = torch.tensor(list(_PATTERN), dtype=torch.uint8) + source_ref = RealKvSource( + tensor=source_cuda.tensor.clone(), + page_size=source_cuda.page_size, + num_bytes_per_token=source_cuda.num_bytes_per_token, + read_bytes=source_cuda.read_bytes, + ) + + # Step 3: run write kernel on slot 0 with the given mode. + _run_write( + buf_pair=self.buf_pair, + input_ids=[7], + positions=[0], + out_cache_loc=[0], + real_kv_sources_pair=((source_cuda,), (source_ref,)), + real_kv_hash_mode=mode, + ) + + # Step 4: assert stored real_kv_hash equals the hand-computed hex literal. + _, _, _, stored_real_kv_hash = read_slot_fields( + canary_buf=self.buf_pair[0], slot_idx=0 + ) + assert stored_real_kv_hash == to_signed_int64( + expected_hash + ), f"stored_real_kv_hash={stored_real_kv_hash:#x} expected={to_signed_int64(expected_hash):#x}" + + def test_paged_real_kv_hash_consistent_across_slots(self) -> None: + """page=16: writing two slots inside same page yields independent real_kv_hash per slot.""" + sources_cuda = make_real_kv_sources( + count=1, + num_bytes_per_token=16, + page_size=16, + num_slots=16, + device=_DEVICE, + ) + pattern_slot3 = bytes(range(1, 17)) + pattern_slot7 = bytes(range(101, 117)) + sources_cuda[0].tensor[0, 3 * 16 : 4 * 16] = torch.tensor( + list(pattern_slot3), dtype=torch.uint8, device=_DEVICE + ) + sources_cuda[0].tensor[0, 7 * 16 : 8 * 16] = torch.tensor( + list(pattern_slot7), dtype=torch.uint8, device=_DEVICE + ) + sources_ref = clone_real_kv_sources(sources_cuda) + + _run_write( + buf_pair=self.buf_pair, + input_ids=[42, 84], + positions=[0, 1], + out_cache_loc=[3, 7], + real_kv_sources_pair=(sources_cuda, sources_ref), + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ) + + slot3 = read_slot_fields(canary_buf=self.buf_pair[0], slot_idx=3) + slot7 = read_slot_fields(canary_buf=self.buf_pair[0], slot_idx=7) + assert slot3[3] == to_signed_int64(_hand_fold_all(pattern_slot3)) + assert slot7[3] == to_signed_int64(_hand_fold_all(pattern_slot7)) + assert slot3[3] != slot7[3] + + def test_multi_source_real_kv_fold_order_matters(self) -> None: + """Two sources folded in reverse order yields a different real_kv_hash (fold is ordered).""" + sources_a = make_real_kv_sources( + count=2, num_bytes_per_token=16, num_slots=8, device=_DEVICE + ) + sources_b = tuple(reversed(sources_a)) + + def _run_with(srcs: tuple[RealKvSource, ...]) -> tuple[int, int, int, int]: + buf = make_canary_buf(num_slots=16, slot_stride_bytes=32, device=_DEVICE) + plan = make_write_plan( + write_offsets=[0, 1], + seed_slot_indices=[-1], + num_valid_reqs=1, + device=_DEVICE, + ) + log = FakeViolationLog.allocate(device=_DEVICE) + launch_canary_write_kernel( + context=VerifyOrWriteContext( + canary_buf=buf, + kernel_kind=CanaryLaunchTag.HEAD_K_FULL, + violation_ring=log.ring, + violation_write_index=log.write_index, + slot_run_counter=log.slot_run_counter, + kernel_run_counter=log.kernel_run_counter, + enable_chain_position_assert=log.enable_chain_position_assert, + real_kv_sources=srcs, + real_kv_hash_mode=consts.RealKvHashMode.ALL, + ), + plan=plan, + input_ids=_int32_tensor([1]), + positions=_int32_tensor([0]), + out_cache_loc=_int32_tensor([2]), + enable_write_input_assert=False, + expected_input_tokens=None, + expected_input_positions=None, + ) + torch.cuda.synchronize() + return read_slot_fields(canary_buf=buf, slot_idx=2) + + fields_a = _run_with(sources_a) + fields_b = _run_with(sources_b) + assert fields_a[3] != 0 + assert fields_b[3] != 0 + assert ( + fields_a[3] != fields_b[3] + ), "reversing source order must change real_kv_hash (fold is ordered)" + + class TestRunCounter: def setup_method(self) -> None: self.buf_pair = _make_default_buf_pair() @@ -718,6 +1044,9 @@ class TestRunCounter: expected_input_positions=pseudo_positions, cuda_log=cuda_log, ref_log=ref_log, + real_kv_sources_cuda=(), + real_kv_sources_ref=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, assert_equal=False, ) @@ -806,6 +1135,8 @@ class TestMisc: slot_run_counter=log.slot_run_counter, kernel_run_counter=log.kernel_run_counter, enable_chain_position_assert=log.enable_chain_position_assert, + real_kv_sources=(), + real_kv_hash_mode=consts.RealKvHashMode.NONE, ) module = _RecordingWriteModule() diff --git a/python/sglang/srt/kv_canary/buffer_group.py b/python/sglang/srt/kv_canary/buffer_group.py index 527b47bd3..8f6206f1f 100644 --- a/python/sglang/srt/kv_canary/buffer_group.py +++ b/python/sglang/srt/kv_canary/buffer_group.py @@ -6,6 +6,8 @@ from typing import Optional import torch +from sglang.jit_kernel.kv_canary.verify import RealKvSource + class PoolKind(IntEnum): """Which attention regime a canary group belongs to. @@ -22,12 +24,13 @@ class PoolKind(IntEnum): @dataclass(frozen=True, slots=True, kw_only=True) class CanaryBufferGroup: - """Canary buffers for one (PoolKind × K-half | V-half) on a pool. + """Canary buffers + real-KV sources for one (PoolKind × K-half | V-half) on a pool. - Each (head | tail) launch sees a single 2-D uint8 buf for the canary. Head and tail use separate canary - buffers so they can be staged at different points in the forward pass without overwriting each other. + Each (head | tail) launch sees a single 2-D uint8 buf for the canary, plus a list of RealKvSource for the + real-KV mixin. Head and tail use separate canary buffers so they can be staged at different points in the + forward pass without overwriting each other. - MLA-style pools have no V half (v_head / v_tail = None). SWA pools have two + MLA-style pools have no V half (v_head / v_tail = None; real_kv_sources_v is empty). SWA pools have two CanaryBufferGroup instances (FULL sized to the full sub-pool, SWA sized to the swa sub-pool). Fields: @@ -36,6 +39,10 @@ class CanaryBufferGroup: k_tail: Tail canary buffer for K-half launches, same shape, uint8. v_head: Same for V-half, or None for MLA-style pools. v_tail: Same for V-half, or None. + real_kv_sources_k: Real KV pieces folded into the K-half canary's real_kv_hash. Tuple length is + pool-specific (1 for simple MHA, more for multi-layer / weird-layout pools). Empty tuple = + real-KV mixin disabled for this half. + real_kv_sources_v: Same for V-half. Empty tuple iff v_head is None or the mixin is disabled. swa_index_lut: SWA full-to-swa index mapping LUT, shape [full_pool_size + 1], int64, or None for FULL groups. Used by launch_canary_plan_kernels to translate verify/seed slot indices at plan time, and by launch_canary_write_kernel to translate write slots inline. None iff kind == PoolKind.FULL. @@ -50,6 +57,8 @@ class CanaryBufferGroup: k_tail: torch.Tensor v_head: Optional[torch.Tensor] v_tail: Optional[torch.Tensor] + real_kv_sources_k: tuple[RealKvSource, ...] + real_kv_sources_v: tuple[RealKvSource, ...] swa_index_lut: Optional[torch.Tensor] kv_token_id_vs_position_offset: int diff --git a/python/sglang/srt/kv_canary/config.py b/python/sglang/srt/kv_canary/config.py index 0b7f6cd73..7da1c2a34 100644 --- a/python/sglang/srt/kv_canary/config.py +++ b/python/sglang/srt/kv_canary/config.py @@ -4,6 +4,9 @@ from dataclasses import dataclass from enum import Enum from typing import TYPE_CHECKING +from sglang.jit_kernel.kv_canary.consts import ( + RealKvHashMode, +) from sglang.srt.environ import envs if TYPE_CHECKING: @@ -33,6 +36,8 @@ class CanaryConfig: sweep_interval: 0 disables sweep entirely; positive N means every N-th forward step the runner additionally walks all radix-tree-held slots (overlap with per-forward HEAD/TAIL is harmless redundancy) and verifies them. + real_kv_hash_mode: RealKvHashMode (NONE / PARTIAL / ALL). Uniform across head/tail/sweep launches; + PARTIAL (first 16B, hard cap) is cheap enough for production defaults. enable_write_input_assert: bool. True = launch_canary_write_kernel additionally compares forward_batch.input_ids[i] / positions[i] against caller-supplied expected_input_tokens[i] / expected_input_positions[i]; mismatch records a violation. Only useful when something else @@ -43,6 +48,7 @@ class CanaryConfig: mode: CanaryMode ring_capacity: int sweep_interval: int + real_kv_hash_mode: RealKvHashMode enable_write_input_assert: bool @classmethod @@ -53,9 +59,12 @@ class CanaryConfig: f"kv-canary: kv_canary must be one of none/log/raise, got {mode_raw!r}" ) + real_kv_raw = server_args.kv_canary_real_data.strip().upper() + return cls( mode=CanaryMode(mode_raw), ring_capacity=envs.SGLANG_KV_CANARY_RING_CAPACITY.get(), sweep_interval=server_args.kv_canary_sweep_interval, + real_kv_hash_mode=RealKvHashMode[real_kv_raw], enable_write_input_assert=envs.SGLANG_KV_CANARY_ENABLE_WRITE_INPUT_ASSERT.get(), ) diff --git a/python/sglang/srt/kv_canary/endpoint.py b/python/sglang/srt/kv_canary/endpoint.py index cd0ce0403..c143d1280 100644 --- a/python/sglang/srt/kv_canary/endpoint.py +++ b/python/sglang/srt/kv_canary/endpoint.py @@ -5,8 +5,12 @@ from typing import Optional import torch +from sglang.jit_kernel.kv_canary.consts import ( + RealKvHashMode, +) from sglang.jit_kernel.kv_canary.verify import ( CanaryLaunchTag, + RealKvSource, VerifyOrWriteContext, VerifyPlan, launch_canary_verify_kernel, @@ -28,6 +32,7 @@ class CanaryEndpoint: kernel_kind: CanaryLaunchTag canary_buf: torch.Tensor full_to_swa_index_mapping: Optional[torch.Tensor] + real_kv_sources: tuple[RealKvSource, ...] slot_run_counter_view: torch.Tensor kernel_run_counter_view: torch.Tensor enable_chain_position_assert: torch.Tensor @@ -44,6 +49,7 @@ class CanaryEndpoint: enable_verify_token_assert: bool, expected_inputs: ExpectedInputs, violation_log: ViolationLog, + real_kv_hash_mode: RealKvHashMode, ) -> None: if _is_sweep_tag(self.kernel_kind): raise NotImplementedError( @@ -52,6 +58,7 @@ class CanaryEndpoint: context = self._make_verify_or_write_context( violation_log=violation_log, + real_kv_hash_mode=real_kv_hash_mode, ) launch_canary_verify_kernel( context=context, @@ -87,6 +94,7 @@ class CanaryEndpoint: *, verify_plan: VerifyPlan, violation_log: ViolationLog, + real_kv_hash_mode: RealKvHashMode, ) -> None: if not _is_sweep_tag(self.kernel_kind): raise NotImplementedError( @@ -96,6 +104,7 @@ class CanaryEndpoint: launch_canary_verify_kernel( context=self._make_verify_or_write_context( violation_log=violation_log, + real_kv_hash_mode=real_kv_hash_mode, ), plan=verify_plan, check_verify_expected_token=False, @@ -105,6 +114,7 @@ class CanaryEndpoint: self, *, violation_log: ViolationLog, + real_kv_hash_mode: RealKvHashMode, ) -> VerifyOrWriteContext: return VerifyOrWriteContext( canary_buf=self.canary_buf, @@ -113,6 +123,8 @@ class CanaryEndpoint: violation_write_index=violation_log.violation_write_index, slot_run_counter=self.slot_run_counter_view, kernel_run_counter=self.kernel_run_counter_view, + real_kv_sources=self.real_kv_sources, + real_kv_hash_mode=real_kv_hash_mode, enable_chain_position_assert=self.enable_chain_position_assert, ) @@ -141,6 +153,16 @@ def _resolve_canary_buf( return group.v_tail +def _resolve_real_kv_sources( + *, + half: str, + group: CanaryBufferGroup, +) -> tuple[RealKvSource, ...]: + if half == "K": + return group.real_kv_sources_k + return group.real_kv_sources_v + + _FULL_LAYOUT: tuple[tuple[CanaryLaunchTag, str, str], ...] = ( (CanaryLaunchTag.HEAD_K_FULL, "HEAD", "K"), (CanaryLaunchTag.HEAD_V_FULL, "HEAD", "V"), @@ -177,6 +199,9 @@ def build_endpoints_from_group( buf_slot = "TAIL" if slot == "SWEEP" else slot canary_buf = _resolve_canary_buf(slot=buf_slot, half=half, group=group) + real_kv_sources = ( + () if slot == "HEAD" else _resolve_real_kv_sources(half=half, group=group) + ) lut = group.swa_index_lut if pool_kind is PoolKind.SWA else None slot_view = device_state.slot_run_counters[tag.value : tag.value + 1] kernel_view = device_state.kernel_run_counters[tag.value : tag.value + 1] @@ -185,6 +210,7 @@ def build_endpoints_from_group( kernel_kind=tag, canary_buf=canary_buf, full_to_swa_index_mapping=lut, + real_kv_sources=real_kv_sources, slot_run_counter_view=slot_view, kernel_run_counter_view=kernel_view, enable_chain_position_assert=device_state.enable_chain_position_assert, diff --git a/python/sglang/srt/kv_canary/perturb/utils.py b/python/sglang/srt/kv_canary/perturb/utils.py index db9dc17d7..47c4a41c6 100644 --- a/python/sglang/srt/kv_canary/perturb/utils.py +++ b/python/sglang/srt/kv_canary/perturb/utils.py @@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Optional import torch +from sglang.jit_kernel.kv_canary.verify import RealKvSource from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup, PoolKind from sglang.srt.kv_canary.perturb.config import PerturbConfig, TargetGroupKind @@ -90,18 +91,111 @@ def pick_target_group( buffer_groups: tuple[CanaryBufferGroup, ...], target_kind: TargetGroupKind, ) -> Optional[CanaryBufferGroup]: - """Filter buffer_groups by target_kind. + """Filter buffer_groups by target_kind restricted to groups with non-empty real_kv_sources_k. Returns None if no group matches. """ + eligible = [group for group in buffer_groups if group.real_kv_sources_k] + if not eligible: + return None if target_kind == TargetGroupKind.FULL: want = PoolKind.FULL elif target_kind == TargetGroupKind.SWA: want = PoolKind.SWA else: raise ValueError(f"Unsupported target_group_kind: {target_kind!r}") - filtered = [group for group in buffer_groups if group.kind == want] + filtered = [group for group in eligible if group.kind == want] if not filtered: return None pick = random.randrange(len(filtered)) return filtered[pick] + + +def flip_random_source_byte_and_log( + *, + perturb_name: str, + group: CanaryBufferGroup, + slot_idx: int, +) -> None: + """Pick a random K-half real_kv source on group, flip byte 0 of slot_idx's tile + in it, and log the result. Logs and returns silently when the group has no + real_kv_sources_k or the slot cannot be mapped into the chosen source.""" + if not group.real_kv_sources_k: + logger.info( + "kv_canary perturb %s: skipped because group=%s has no real_kv_sources_k", + perturb_name, + group.kind.name, + ) + return + source_pick = random.randrange(len(group.real_kv_sources_k)) + source = group.real_kv_sources_k[source_pick] + flip_result = flip_first_byte_in_source( + group=group, source=source, slot_idx=slot_idx + ) + if flip_result is None: + logger.info( + "kv_canary perturb %s: skipped because slot=%d could not be mapped " + "into group=%s source_idx=%d", + perturb_name, + slot_idx, + group.kind.name, + source_pick, + ) + return + row, col, original_byte = flip_result + logger.info( + "kv_canary perturb %s: group=%s source_idx=%d slot=%d row=%d col=%d " + "original_byte=0x%02X new_byte=0x%02X", + perturb_name, + group.kind.name, + source_pick, + slot_idx, + row, + col, + original_byte, + original_byte ^ 0xFF, + ) + + +def flip_first_byte_in_source( + *, + group: CanaryBufferGroup, + source: RealKvSource, + slot_idx: int, + slot_is_physical: bool = False, +) -> Optional[tuple[int, int, int]]: + """XOR 0xFF on byte 0 of slot_idx's tile in source.tensor (column + `(physical_slot % page_size) * num_bytes_per_token`, not row offset 0). + + For SWA groups, slot_idx is translated through group.swa_index_lut before computing + (row, col). Returns (row, col, original_byte) for logging, or None if the slot is + out-of-range / source is degenerate. + """ + if source.num_bytes_per_token <= 0 or source.read_bytes <= 0: + return None + + physical_slot = slot_idx + if ( + not slot_is_physical + and group.kind == PoolKind.SWA + and group.swa_index_lut is not None + ): + lut = group.swa_index_lut + if slot_idx < 0 or slot_idx >= int(lut.shape[0]): + return None + physical_slot = int(lut[slot_idx].detach().to("cpu").item()) + if physical_slot < 0: + return None + + page_size = max(1, source.page_size) + row = physical_slot // page_size + col = (physical_slot % page_size) * source.num_bytes_per_token + if row < 0 or row >= int(source.tensor.shape[0]): + return None + if col < 0 or col >= int(source.tensor.shape[1]): + return None + + flat = source.tensor + original_byte = int(flat[row, col].item()) + flat[row, col] = original_byte ^ 0xFF + return row, col, original_byte diff --git a/python/sglang/srt/kv_canary/pool_patcher/adapters/dsv4.py b/python/sglang/srt/kv_canary/pool_patcher/adapters/dsv4.py index 4204e2dc6..aeb473f88 100644 --- a/python/sglang/srt/kv_canary/pool_patcher/adapters/dsv4.py +++ b/python/sglang/srt/kv_canary/pool_patcher/adapters/dsv4.py @@ -11,13 +11,19 @@ def attach_dsv4( *, pool: object, device: torch.device, + read_bytes: int, kv_token_id_vs_position_offset: int, ) -> tuple[CanaryBufferGroup, ...]: """Attach canary buffers to a DeepSeekV4TokenToKVPool. TODO: only the swa_kv_pool sub-pool is wired; c4_kv_pool / c128_kv_pool / c4_indexer_kv_pool / compress state pools are left uncovered. + TODO: even on swa_kv_pool, real-KV fingerprint is disabled (read_bytes is + ignored). DSV4 stores 584 B/token which is not 16-aligned (584 % 16 == 8), + so num_bytes_per_token cannot satisfy the 128-bit load alignment precondition. """ + del read_bytes + sub_pool = pool.swa_kv_pool num_slots = int(sub_pool.size) @@ -30,6 +36,8 @@ def attach_dsv4( k_tail=k_tail, v_head=None, v_tail=None, + real_kv_sources_k=(), + real_kv_sources_v=(), swa_index_lut=pool.full_to_swa_index_mapping, kv_token_id_vs_position_offset=kv_token_id_vs_position_offset, ) diff --git a/python/sglang/srt/kv_canary/pool_patcher/adapters/mha.py b/python/sglang/srt/kv_canary/pool_patcher/adapters/mha.py index fc8dd5d0c..c4cb4868d 100644 --- a/python/sglang/srt/kv_canary/pool_patcher/adapters/mha.py +++ b/python/sglang/srt/kv_canary/pool_patcher/adapters/mha.py @@ -4,13 +4,17 @@ import torch from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup, PoolKind from sglang.srt.kv_canary.pool_patcher.buf_info_splice import patch_buf_info_method -from sglang.srt.kv_canary.pool_patcher.buffer_alloc import alloc_canary_buf +from sglang.srt.kv_canary.pool_patcher.buffer_alloc import ( + alloc_canary_buf, + make_row_source, +) def attach_mha( *, pool: object, device: torch.device, + read_bytes: int, kv_token_id_vs_position_offset: int, ) -> tuple[CanaryBufferGroup, ...]: num_slots = int(pool.k_buffer[0].shape[0]) @@ -25,6 +29,12 @@ def attach_mha( k_tail=k_tail, v_head=v_head, v_tail=v_tail, + real_kv_sources_k=make_row_source( + layer_buffer=pool.k_buffer[0], read_bytes=read_bytes + ), + real_kv_sources_v=make_row_source( + layer_buffer=pool.v_buffer[0], read_bytes=read_bytes + ), swa_index_lut=None, kv_token_id_vs_position_offset=kv_token_id_vs_position_offset, ) diff --git a/python/sglang/srt/kv_canary/pool_patcher/adapters/swa.py b/python/sglang/srt/kv_canary/pool_patcher/adapters/swa.py index d9012b0c3..a66b73e09 100644 --- a/python/sglang/srt/kv_canary/pool_patcher/adapters/swa.py +++ b/python/sglang/srt/kv_canary/pool_patcher/adapters/swa.py @@ -6,19 +6,24 @@ import torch from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup, PoolKind from sglang.srt.kv_canary.pool_patcher.buf_info_splice import patch_buf_info_method -from sglang.srt.kv_canary.pool_patcher.buffer_alloc import alloc_canary_buf +from sglang.srt.kv_canary.pool_patcher.buffer_alloc import ( + alloc_canary_buf, + make_row_source, +) def attach_swa( *, pool: object, device: torch.device, + read_bytes: int, kv_token_id_vs_position_offset: int, ) -> tuple[CanaryBufferGroup, ...]: full_group = _build_subpool_group( sub_pool=pool.full_kv_pool, kind=PoolKind.FULL, device=device, + read_bytes=read_bytes, swa_lut=None, kv_token_id_vs_position_offset=kv_token_id_vs_position_offset, ) @@ -26,6 +31,7 @@ def attach_swa( sub_pool=pool.swa_kv_pool, kind=PoolKind.SWA, device=device, + read_bytes=read_bytes, swa_lut=pool.full_to_swa_index_mapping, kv_token_id_vs_position_offset=kv_token_id_vs_position_offset, ) @@ -52,6 +58,7 @@ def _build_subpool_group( sub_pool: object, kind: PoolKind, device: torch.device, + read_bytes: int, swa_lut: Optional[torch.Tensor], kv_token_id_vs_position_offset: int, ) -> CanaryBufferGroup: @@ -66,6 +73,12 @@ def _build_subpool_group( k_tail=k_tail, v_head=v_head, v_tail=v_tail, + real_kv_sources_k=make_row_source( + layer_buffer=sub_pool.k_buffer[0], read_bytes=read_bytes + ), + real_kv_sources_v=make_row_source( + layer_buffer=sub_pool.v_buffer[0], read_bytes=read_bytes + ), swa_index_lut=swa_lut, kv_token_id_vs_position_offset=kv_token_id_vs_position_offset, ) diff --git a/python/sglang/srt/kv_canary/pool_patcher/api.py b/python/sglang/srt/kv_canary/pool_patcher/api.py index 6c1438412..b73599ff8 100644 --- a/python/sglang/srt/kv_canary/pool_patcher/api.py +++ b/python/sglang/srt/kv_canary/pool_patcher/api.py @@ -10,6 +10,7 @@ from sglang.srt.kv_canary.config import CanaryConfig from sglang.srt.kv_canary.pool_patcher.adapters.dsv4 import attach_dsv4 from sglang.srt.kv_canary.pool_patcher.adapters.mha import attach_mha from sglang.srt.kv_canary.pool_patcher.adapters.swa import attach_swa +from sglang.srt.kv_canary.pool_patcher.buffer_alloc import resolve_real_kv_read_bytes from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.memory_pool import ( KVCache, @@ -53,16 +54,19 @@ def attach_canary_buffers( f"supported: {sorted(cls.__name__ for cls in _POOL_ATTACHERS)}" ) + read_bytes = resolve_real_kv_read_bytes(config) groups = attacher( pool=pool, device=device, + read_bytes=read_bytes, kv_token_id_vs_position_offset=kv_token_id_vs_position_offset, ) logger.info( - "attach_canary_buffers: pool=%s attacher=%s n_groups=%d kinds=%s " + "attach_canary_buffers: pool=%s attacher=%s read_bytes=%d n_groups=%d kinds=%s " "kv_token_id_vs_position_offset=%d", type(pool).__name__, attacher.__name__, + read_bytes, len(groups), [g.kind.name for g in groups], kv_token_id_vs_position_offset, diff --git a/python/sglang/srt/kv_canary/pool_patcher/buffer_alloc.py b/python/sglang/srt/kv_canary/pool_patcher/buffer_alloc.py index 8e2f265ee..e40494d1d 100644 --- a/python/sglang/srt/kv_canary/pool_patcher/buffer_alloc.py +++ b/python/sglang/srt/kv_canary/pool_patcher/buffer_alloc.py @@ -1,8 +1,27 @@ from __future__ import annotations +import sys +from typing import Tuple + import torch -from sglang.jit_kernel.kv_canary.verify import CANARY_SLOT_BYTES +from sglang.jit_kernel.kv_canary.consts import RealKvHashMode +from sglang.jit_kernel.kv_canary.verify import ( + CANARY_SLOT_BYTES, + RealKvSource, +) +from sglang.srt.kv_canary.config import CanaryConfig + +_PARTIAL_REAL_KV_READ_BYTES = 16 +_REAL_KV_READ_ALIGN = 16 + + +def resolve_real_kv_read_bytes(config: CanaryConfig) -> int: + if config.real_kv_hash_mode is RealKvHashMode.NONE: + return 0 + if config.real_kv_hash_mode is RealKvHashMode.ALL: + return sys.maxsize + return _PARTIAL_REAL_KV_READ_BYTES def alloc_canary_buf( @@ -11,3 +30,86 @@ def alloc_canary_buf( device: torch.device, ) -> torch.Tensor: return torch.zeros(num_slots, CANARY_SLOT_BYTES, dtype=torch.uint8, device=device) + + +def _clip_read_bytes_aligned(*, requested: int, num_bytes_per_token: int) -> int: + """Validate and clip read_bytes for the CUDA fold kernel's 128-bit aligned loads. + + Normalizes sentinels (``sys.maxsize`` -> ``num_bytes_per_token``, ``0`` -> ``0``) and + rejects negative / unaligned / oversized requests. + """ + if num_bytes_per_token <= 0 or num_bytes_per_token % _REAL_KV_READ_ALIGN != 0: + raise ValueError( + "kv-canary: num_bytes_per_token must be a positive multiple of " + f"{_REAL_KV_READ_ALIGN}, got {num_bytes_per_token}" + ) + if requested == 0: + return 0 + if requested == sys.maxsize: + return num_bytes_per_token + if requested < 0: + raise ValueError(f"kv-canary: read_bytes must be non-negative, got {requested}") + if requested > num_bytes_per_token: + raise ValueError( + "kv-canary: read_bytes must be <= num_bytes_per_token " + f"({num_bytes_per_token}), got {requested}" + ) + if requested % _REAL_KV_READ_ALIGN != 0: + raise ValueError( + "kv-canary: read_bytes must be a multiple of " + f"{_REAL_KV_READ_ALIGN}, got {requested}" + ) + return requested + + +def make_row_source( + *, + layer_buffer: torch.Tensor, + read_bytes: int, +) -> Tuple[RealKvSource, ...]: + contiguous = layer_buffer.contiguous() + num_slots = int(contiguous.shape[0]) + if num_slots == 0 or read_bytes == 0: + return () + flat = contiguous.view(torch.uint8).reshape(num_slots, -1) + num_bytes_per_token = int(flat.shape[1]) + clipped = _clip_read_bytes_aligned( + requested=read_bytes, num_bytes_per_token=num_bytes_per_token + ) + if clipped == 0: + return () + return ( + RealKvSource( + tensor=flat, + page_size=1, + num_bytes_per_token=num_bytes_per_token, + read_bytes=clipped, + ), + ) + + +def make_packed_source( + *, + page_buffer: torch.Tensor, + page_size: int, + bytes_per_token: int, + read_bytes: int, +) -> Tuple[RealKvSource, ...]: + if read_bytes == 0 or page_buffer.numel() == 0: + return () + flat = page_buffer.contiguous().view(torch.uint8) + if flat.ndim == 1: + flat = flat.reshape(1, -1) + clipped = _clip_read_bytes_aligned( + requested=read_bytes, num_bytes_per_token=bytes_per_token + ) + if clipped == 0: + return () + return ( + RealKvSource( + tensor=flat, + page_size=page_size, + num_bytes_per_token=bytes_per_token, + read_bytes=clipped, + ), + ) diff --git a/python/sglang/srt/kv_canary/runner/kernel_launcher.py b/python/sglang/srt/kv_canary/runner/kernel_launcher.py index 63f9f97c4..8bf459772 100644 --- a/python/sglang/srt/kv_canary/runner/kernel_launcher.py +++ b/python/sglang/srt/kv_canary/runner/kernel_launcher.py @@ -4,6 +4,9 @@ from typing import TYPE_CHECKING, Callable, Optional import torch +from sglang.jit_kernel.kv_canary.consts import ( + RealKvHashMode, +) from sglang.jit_kernel.kv_canary.plan import launch_canary_plan_kernels from sglang.jit_kernel.kv_canary.verify import ( CanaryLaunchTag, @@ -64,6 +67,7 @@ def launch_endpoints_per_forward( forward_batch: "ForwardBatch", expected_inputs: ExpectedInputs, violation_log: ViolationLog, + real_kv_hash_mode: RealKvHashMode, enable_write_input_assert: bool, enable_verify_token_assert: bool, ) -> None: @@ -104,6 +108,7 @@ def launch_endpoints_per_forward( enable_verify_token_assert=enable_verify_token_assert, expected_inputs=expected_inputs, violation_log=violation_log, + real_kv_hash_mode=real_kv_hash_mode, ) @@ -113,6 +118,7 @@ def launch_endpoints_sweep( group: CanaryBufferGroup, verify_plan: VerifyPlan, violation_log: ViolationLog, + real_kv_hash_mode: RealKvHashMode, ) -> None: active_endpoints = [ endpoint @@ -127,6 +133,7 @@ def launch_endpoints_sweep( endpoint.launch_sweep( verify_plan=verify_plan, violation_log=violation_log, + real_kv_hash_mode=real_kv_hash_mode, ) diff --git a/python/sglang/srt/kv_canary/runner/sweep.py b/python/sglang/srt/kv_canary/runner/sweep.py index 32d343f28..496991e1f 100644 --- a/python/sglang/srt/kv_canary/runner/sweep.py +++ b/python/sglang/srt/kv_canary/runner/sweep.py @@ -73,6 +73,7 @@ class SweepOrchestrator: group=group, verify_plan=verify_plan, violation_log=violation_log, + real_kv_hash_mode=self._config.real_kv_hash_mode, ) self._sweep_passes += 1 diff --git a/python/sglang/srt/kv_canary/single_forward_manager/manager.py b/python/sglang/srt/kv_canary/single_forward_manager/manager.py index c85f9395e..5f4293176 100644 --- a/python/sglang/srt/kv_canary/single_forward_manager/manager.py +++ b/python/sglang/srt/kv_canary/single_forward_manager/manager.py @@ -201,6 +201,7 @@ class SingleForwardManager: forward_batch=forward_batch, expected_inputs=expected_inputs_slice, violation_log=violation_log, + real_kv_hash_mode=self._config.real_kv_hash_mode, enable_write_input_assert=enable_write_input_assert, enable_verify_token_assert=False, ) @@ -238,6 +239,7 @@ class SingleForwardManager: forward_batch=forward_batch, expected_inputs=expected_inputs_slice, violation_log=violation_log, + real_kv_hash_mode=self._config.real_kv_hash_mode, enable_write_input_assert=enable_write_input_assert, enable_verify_token_assert=False, ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 920a78d33..2723d9749 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -29,6 +29,7 @@ import uuid from functools import cached_property from typing import Any, Callable, Dict, List, Literal, Optional, Union +from sglang.jit_kernel.kv_canary.consts import RealKvHashMode from sglang.srt.arg_groups.argparse_actions import ( DeprecatedAction, DeprecatedAliasStoreAction, @@ -780,6 +781,7 @@ class ServerArgs: disable_attn_tp_gather: bool = False gc_threshold: Optional[List[int]] = None kv_canary: str = "none" + kv_canary_real_data: str = "none" kv_canary_sweep_interval: int = 0 # Context parallelism used in the long sequence prefill phase of DeepSeek v3.2 enable_dsa_prefill_context_parallel: bool = False @@ -6321,6 +6323,18 @@ class ServerArgs: "'raise' fails the server on the first detected mismatch (CI lane)." ), ) + parser.add_argument( + "--kv-canary-real-data", + type=str, + default=ServerArgs.kv_canary_real_data, + choices=[m.name.lower() for m in RealKvHashMode], + help=( + "Check the real KV-cache in the canary. " + "'none' (default) disables the feature. " + "'partial' checks the first 16 bytes of each real-KV slot. " + "'all' checks the full real-KV slot." + ), + ) parser.add_argument( "--kv-canary-sweep-interval", type=int, diff --git a/python/sglang/test/kv_canary/fixtures.py b/python/sglang/test/kv_canary/fixtures.py index c6d934d36..bf2b361c8 100644 --- a/python/sglang/test/kv_canary/fixtures.py +++ b/python/sglang/test/kv_canary/fixtures.py @@ -6,7 +6,8 @@ from typing import List, Optional import torch -from sglang.jit_kernel.kv_canary.verify import CANARY_SLOT_BYTES +from sglang.jit_kernel.kv_canary import consts +from sglang.jit_kernel.kv_canary.verify import CANARY_SLOT_BYTES, RealKvSource from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup, PoolKind from sglang.srt.kv_canary.config import CanaryConfig, CanaryMode from sglang.srt.kv_canary.pool_patcher.adapters.mha import attach_mha @@ -136,6 +137,7 @@ def make_base_config() -> CanaryConfig: mode=CanaryMode.RAISE, ring_capacity=1024, sweep_interval=0, + real_kv_hash_mode=consts.RealKvHashMode.NONE, enable_write_input_assert=False, ) @@ -223,6 +225,8 @@ def make_buffer_group( device: torch.device = DEFAULT_DEVICE, kind: PoolKind = PoolKind.FULL, has_v: bool = True, + has_real_kv: bool = False, + real_kv_source: Optional[RealKvSource] = None, swa_index_lut: Optional[torch.Tensor] = None, num_slots: int = 4, kv_token_id_vs_position_offset: int = 0, @@ -232,12 +236,25 @@ def make_buffer_group( num_slots, CANARY_SLOT_BYTES, dtype=torch.uint8, device=device ) + if has_real_kv: + source = real_kv_source or RealKvSource( + tensor=torch.zeros(num_slots, 16, dtype=torch.uint8, device=device), + page_size=1, + num_bytes_per_token=16, + read_bytes=16, + ) + real_kv_sources = (source,) + else: + real_kv_sources = () + return CanaryBufferGroup( kind=kind, k_head=_zero(), k_tail=_zero(), v_head=_zero() if has_v else None, v_tail=_zero() if has_v else None, + real_kv_sources_k=real_kv_sources, + real_kv_sources_v=real_kv_sources if has_v else (), swa_index_lut=swa_index_lut, kv_token_id_vs_position_offset=kv_token_id_vs_position_offset, ) diff --git a/python/sglang/test/kv_canary/runner_test_base.py b/python/sglang/test/kv_canary/runner_test_base.py index 0e0ac942c..e84a94904 100644 --- a/python/sglang/test/kv_canary/runner_test_base.py +++ b/python/sglang/test/kv_canary/runner_test_base.py @@ -5,6 +5,7 @@ from unittest.mock import patch import torch +from sglang.jit_kernel.kv_canary.consts import RealKvHashMode from sglang.jit_kernel.kv_canary.verify import CanaryLaunchTag from sglang.srt.kv_canary import endpoint as endpoint_module from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup @@ -26,12 +27,14 @@ def make_config( mode: CanaryMode = CanaryMode.RAISE, ring_capacity: int = 1024, sweep_interval: int = 0, + real_kv_hash_mode: RealKvHashMode = RealKvHashMode.NONE, enable_write_input_assert: bool = False, ) -> CanaryConfig: return CanaryConfig( mode=mode, ring_capacity=ring_capacity, sweep_interval=sweep_interval, + real_kv_hash_mode=real_kv_hash_mode, enable_write_input_assert=enable_write_input_assert, ) diff --git a/test/registered/kv_canary/test_self_unit_buffer_alloc.py b/test/registered/kv_canary/test_self_unit_buffer_alloc.py index db8930186..3e0ced43c 100644 --- a/test/registered/kv_canary/test_self_unit_buffer_alloc.py +++ b/test/registered/kv_canary/test_self_unit_buffer_alloc.py @@ -1,26 +1,140 @@ from __future__ import annotations +import sys import unittest import torch -from sglang.jit_kernel.kv_canary.verify import CANARY_SLOT_BYTES -from sglang.srt.kv_canary.pool_patcher.buffer_alloc import alloc_canary_buf +from sglang.jit_kernel.kv_canary.consts import RealKvHashMode +from sglang.srt.kv_canary.config import CanaryConfig, CanaryMode +from sglang.srt.kv_canary.pool_patcher.buffer_alloc import ( + make_packed_source, + make_row_source, + resolve_real_kv_read_bytes, +) from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.kv_canary.fixtures import DEFAULT_DEVICE from sglang.test.test_utils import CustomTestCase register_cuda_ci(est_time=10, stage="extra-a", runner_config="1-gpu-small") -class TestAllocCanaryBuf(CustomTestCase): - def test_alloc_canary_buf_shape_and_dtype(self) -> None: - """Verify alloc_canary_buf returns a zeroed uint8 buffer of [num_slots, CANARY_SLOT_BYTES].""" - buf = alloc_canary_buf(num_slots=8, device=DEFAULT_DEVICE) - self.assertEqual(buf.shape, (8, CANARY_SLOT_BYTES)) - self.assertEqual(buf.dtype, torch.uint8) - self.assertEqual(buf.device.type, DEFAULT_DEVICE.type) - self.assertTrue(torch.equal(buf, torch.zeros_like(buf))) +def _config(mode: RealKvHashMode) -> CanaryConfig: + return CanaryConfig( + mode=CanaryMode.RAISE, + ring_capacity=1024, + sweep_interval=0, + real_kv_hash_mode=mode, + enable_write_input_assert=False, + ) + + +class TestResolveRealKvReadBytes(CustomTestCase): + def test_resolve_real_kv_read_bytes_off_returns_zero(self) -> None: + """Verify NONE mode disables real KV byte reads.""" + self.assertEqual(resolve_real_kv_read_bytes(_config(RealKvHashMode.NONE)), 0) + + def test_resolve_real_kv_read_bytes_partial_returns_16(self) -> None: + """Verify PARTIAL mode reads the fixed byte prefix.""" + self.assertEqual( + resolve_real_kv_read_bytes(_config(RealKvHashMode.PARTIAL)), 16 + ) + + def test_resolve_real_kv_read_bytes_all_returns_sentinel_so_full_stride_used( + self, + ) -> None: + """Verify ALL mode requests the full token stride.""" + self.assertEqual( + resolve_real_kv_read_bytes(_config(RealKvHashMode.ALL)), sys.maxsize + ) + + +class TestMakeRowSource(CustomTestCase): + def test_make_row_source_large_stride(self) -> None: + """Verify row sources with 128-byte stride return the requested clip / full stride.""" + bytes_per_token = 128 + layer_buf = torch.zeros(4, bytes_per_token, dtype=torch.uint8) + cases = [ + ("partial", 32, 32), + ("all", sys.maxsize, bytes_per_token), + ] + for label, read_bytes, expected_read in cases: + with self.subTest(label=label): + sources = make_row_source(layer_buffer=layer_buf, read_bytes=read_bytes) + self.assertEqual(len(sources), 1) + self.assertEqual(sources[0].read_bytes, expected_read) + self.assertEqual(sources[0].num_bytes_per_token, bytes_per_token) + + def test_make_row_source_small_stride_raises(self) -> None: + """Verify row sources reject 8-byte strides (cannot satisfy 16-byte aligned loads).""" + layer_buf = torch.zeros(4, 8, dtype=torch.uint8) + for label, read_bytes in [("partial", 32), ("all", sys.maxsize)]: + with self.subTest(label=label): + with self.assertRaisesRegex(ValueError, "num_bytes_per_token"): + make_row_source(layer_buffer=layer_buf, read_bytes=read_bytes) + + +class TestMakePackedSource(CustomTestCase): + def test_make_packed_source_large_stride(self) -> None: + """Verify packed sources with 128-byte stride return the requested clip / full stride.""" + bytes_per_token = 128 + page_size = 2 + page_buffer = torch.zeros(4, bytes_per_token * page_size, dtype=torch.uint8) + cases = [ + ("partial", 32, 32), + ("all", sys.maxsize, bytes_per_token), + ] + for label, read_bytes, expected_read in cases: + with self.subTest(label=label): + sources = make_packed_source( + page_buffer=page_buffer, + page_size=page_size, + bytes_per_token=bytes_per_token, + read_bytes=read_bytes, + ) + self.assertEqual(len(sources), 1) + self.assertEqual(sources[0].read_bytes, expected_read) + self.assertEqual(sources[0].num_bytes_per_token, bytes_per_token) + + def test_make_packed_source_small_stride_raises(self) -> None: + """Verify packed sources reject 8-byte strides (cannot satisfy 16-byte aligned loads).""" + bytes_per_token = 8 + page_size = 1 + page_buffer = torch.zeros(4, bytes_per_token, dtype=torch.uint8) + for label, read_bytes in [("partial", 32), ("all", sys.maxsize)]: + with self.subTest(label=label): + with self.assertRaisesRegex(ValueError, "num_bytes_per_token"): + make_packed_source( + page_buffer=page_buffer, + page_size=page_size, + bytes_per_token=bytes_per_token, + read_bytes=read_bytes, + ) + + def test_make_packed_source_unaligned_read_bytes_raises(self) -> None: + """Verify packed sources reject unaligned explicit reads.""" + bytes_per_token = 128 + page_size = 1 + page_buffer = torch.zeros(4, bytes_per_token, dtype=torch.uint8) + with self.assertRaisesRegex(ValueError, "multiple of 16"): + make_packed_source( + page_buffer=page_buffer, + page_size=page_size, + bytes_per_token=bytes_per_token, + read_bytes=24, + ) + + def test_make_packed_source_oversized_read_bytes_raises(self) -> None: + """Verify packed sources reject oversized explicit reads.""" + bytes_per_token = 128 + page_size = 1 + page_buffer = torch.zeros(4, bytes_per_token, dtype=torch.uint8) + with self.assertRaisesRegex(ValueError, "<= num_bytes_per_token"): + make_packed_source( + page_buffer=page_buffer, + page_size=page_size, + bytes_per_token=bytes_per_token, + read_bytes=256, + ) if __name__ == "__main__": diff --git a/test/registered/kv_canary/test_self_unit_endpoint.py b/test/registered/kv_canary/test_self_unit_endpoint.py index 8c406d2dd..b301b27b5 100644 --- a/test/registered/kv_canary/test_self_unit_endpoint.py +++ b/test/registered/kv_canary/test_self_unit_endpoint.py @@ -6,6 +6,7 @@ from unittest.mock import patch import torch +from sglang.jit_kernel.kv_canary.consts import RealKvHashMode from sglang.jit_kernel.kv_canary.verify import ( CANARY_SLOT_BYTES, CanaryLaunchTag, @@ -36,6 +37,7 @@ def _make_endpoint(*, device, kernel_kind=CanaryLaunchTag.HEAD_K_FULL, swa_lut=N kernel_kind=kernel_kind, canary_buf=canary_buf, full_to_swa_index_mapping=swa_lut, + real_kv_sources=(), slot_run_counter_view=slot_view, kernel_run_counter_view=kernel_view, enable_chain_position_assert=enable_chain_position_assert, @@ -56,6 +58,7 @@ def _make_kernel_args(device): enable_write_input_assert=False, enable_verify_token_assert=False, expected_inputs=ExpectedInputs.allocate(capacity=1, device=device), + real_kv_hash_mode=RealKvHashMode.NONE, ) @@ -82,6 +85,7 @@ class TestSelfUnitEndpoint(CustomTestCase): ep.launch_sweep( verify_plan=args.verify_plan, violation_log=args.violation_log, + real_kv_hash_mode=args.real_kv_hash_mode, ) self.assertEqual(calls, ["verify"]) @@ -111,6 +115,7 @@ class TestSelfUnitEndpoint(CustomTestCase): enable_verify_token_assert=args.enable_verify_token_assert, expected_inputs=args.expected_inputs, violation_log=args.violation_log, + real_kv_hash_mode=args.real_kv_hash_mode, ) self.assertIn(("verify", CanaryLaunchTag.TAIL_V_SWA), captured) self.assertIn(("write", CanaryLaunchTag.TAIL_V_SWA), captured) @@ -141,10 +146,12 @@ class TestSelfUnitEndpoint(CustomTestCase): ep_a.launch_sweep( verify_plan=plan, violation_log=shared_log, + real_kv_hash_mode=RealKvHashMode.NONE, ) ep_b.launch_sweep( verify_plan=plan, violation_log=shared_log, + real_kv_hash_mode=RealKvHashMode.NONE, ) self.assertEqual(captured_rings[0], captured_rings[1]) self.assertEqual(captured_rings[0], shared_log.violation_ring.data_ptr()) @@ -183,6 +190,7 @@ class TestSelfUnitEndpoint(CustomTestCase): enable_verify_token_assert=args.enable_verify_token_assert, expected_inputs=args.expected_inputs, violation_log=args.violation_log, + real_kv_hash_mode=args.real_kv_hash_mode, ) full_ep.launch_per_forward( verify_plan=args.verify_plan, @@ -194,6 +202,7 @@ class TestSelfUnitEndpoint(CustomTestCase): enable_verify_token_assert=args.enable_verify_token_assert, expected_inputs=args.expected_inputs, violation_log=args.violation_log, + real_kv_hash_mode=args.real_kv_hash_mode, ) # SWA call: out_cache_loc was rewritten via lut gather (so identity-shifted by +100 here). expected_swa = lut[args.out_cache_loc] @@ -235,6 +244,7 @@ class TestSelfUnitEndpoint(CustomTestCase): enable_verify_token_assert=args.enable_verify_token_assert, expected_inputs=args.expected_inputs, violation_log=args.violation_log, + real_kv_hash_mode=args.real_kv_hash_mode, ) self.assertTrue( diff --git a/test/registered/kv_canary/test_self_unit_perturb.py b/test/registered/kv_canary/test_self_unit_perturb.py index d808c981e..9ba493fb7 100644 --- a/test/registered/kv_canary/test_self_unit_perturb.py +++ b/test/registered/kv_canary/test_self_unit_perturb.py @@ -2,14 +2,26 @@ from __future__ import annotations import os import unittest +from typing import cast from unittest.mock import patch +import torch + +from sglang.jit_kernel.kv_canary.verify import RealKvSource +from sglang.srt.kv_canary.buffer_group import PoolKind from sglang.srt.kv_canary.perturb.config import ( PerturbConfig, TargetGroupKind, _parse_target_group_kind, ) +from sglang.srt.kv_canary.perturb.utils import ( + flip_first_byte_in_source, + pick_target_group, +) from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kv_canary.fixtures import ( + make_buffer_group, +) from sglang.test.test_utils import CustomTestCase register_cuda_ci(est_time=10, stage="extra-a", runner_config="1-gpu-small") @@ -55,5 +67,98 @@ class TestParseTargetGroupKind(CustomTestCase): self.assertIsNone(config.target_group_kind) +class TestPickTargetGroup(CustomTestCase): + def test_pick_target_group_filters_exact_kind(self) -> None: + """Verify target group selection returns only the requested pool kind.""" + cases = [ + (TargetGroupKind.FULL, PoolKind.FULL), + (TargetGroupKind.SWA, PoolKind.SWA), + ] + + for target_kind, expected_kind in cases: + with self.subTest(target_kind=target_kind): + full_group = make_buffer_group(kind=PoolKind.FULL, has_real_kv=True) + swa_group = make_buffer_group(kind=PoolKind.SWA, has_real_kv=True) + + group = pick_target_group( + buffer_groups=(full_group, swa_group), + target_kind=target_kind, + ) + + self.assertIsNotNone(group) + self.assertEqual(group.kind, expected_kind) + + def test_pick_target_group_rejects_unsupported_kind(self) -> None: + """Verify target group selection rejects unsupported enum values.""" + full_group = make_buffer_group(kind=PoolKind.FULL, has_real_kv=True) + + with self.assertRaisesRegex(ValueError, "Unsupported target_group_kind"): + pick_target_group( + buffer_groups=(full_group,), + target_kind=cast(TargetGroupKind, 2), + ) + + def test_pick_target_group_ignores_groups_without_real_kv_sources(self) -> None: + """Verify target group selection skips groups without real KV sources.""" + full_group = make_buffer_group(kind=PoolKind.FULL, has_real_kv=False) + swa_group = make_buffer_group(kind=PoolKind.SWA, has_real_kv=True) + + group = pick_target_group( + buffer_groups=(full_group, swa_group), + target_kind=TargetGroupKind.FULL, + ) + + self.assertIsNone(group) + + +class TestPerturbWarmupAndUtils(CustomTestCase): + def test_flip_first_byte_in_source_maps_swa_logical_slot_through_lut( + self, + ) -> None: + """Verify SWA groups map logical slots through swa_index_lut before flipping bytes.""" + source = RealKvSource( + tensor=torch.arange(64, dtype=torch.uint8).view(2, 32), + page_size=2, + num_bytes_per_token=16, + read_bytes=16, + ) + group = make_buffer_group( + kind=PoolKind.SWA, + has_real_kv=True, + real_kv_source=source, + swa_index_lut=torch.tensor([0, 3], dtype=torch.int32), + ) + + snapshot = source.tensor.clone() + result = flip_first_byte_in_source(group=group, source=source, slot_idx=1) + + self.assertEqual(result, (1, 16, int(snapshot[1, 16].item()))) + expected = snapshot.clone() + expected[1, 16] = int(snapshot[1, 16].item()) ^ 0xFF + self.assertTrue(torch.equal(source.tensor, expected)) + + +class TestPerturbUtils(CustomTestCase): + def test_flip_first_byte_in_physical_swa_slot_does_not_translate_twice( + self, + ) -> None: + """Verify a physical SWA slot selected from sweep is not LUT-translated again.""" + group = make_buffer_group(kind=PoolKind.SWA, has_real_kv=True) + source = group.real_kv_sources_k[0] + source.tensor[2, 0] = 0x12 + source.tensor[3, 0] = 0x34 + + result = flip_first_byte_in_source( + group=group, + source=source, + slot_idx=2, + slot_is_physical=True, + ) + + self.assertEqual(result, (2, 0, 0x12)) + self.assertEqual(int(source.tensor[2, 0].item()), 0xED) + self.assertEqual(int(source.tensor[3, 0].item()), 0x34) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/kv_canary/test_self_unit_pool_patcher.py b/test/registered/kv_canary/test_self_unit_pool_patcher.py index 7760c845e..859914721 100644 --- a/test/registered/kv_canary/test_self_unit_pool_patcher.py +++ b/test/registered/kv_canary/test_self_unit_pool_patcher.py @@ -2,7 +2,17 @@ from __future__ import annotations import unittest -from sglang.jit_kernel.kv_canary.verify import CANARY_SLOT_BYTES +import torch + +from sglang.jit_kernel.kv_canary.consts import MAX_REAL_KV_SOURCES, RealKvHashMode +from sglang.jit_kernel.kv_canary.verify import ( + CANARY_SLOT_BYTES, + CanaryLaunchTag, + RealKvSource, + VerifyOrWriteContext, + VerifyPlan, + launch_canary_verify_kernel, +) from sglang.srt.kv_canary.buffer_group import PoolKind from sglang.srt.kv_canary.pool_patcher.api import attach_canary_buffers from sglang.test.ci.ci_register import register_cuda_ci @@ -59,6 +69,47 @@ class TestAttachCanaryBuffers(PoolPatcherHelper, CustomTestCase): self.assertIsNone(groups[PoolKind.FULL].swa_index_lut) +class TestRealKvSources(PoolPatcherHelper, CustomTestCase): + def test_real_kv_sources_above_4_raises(self): + """Verify too many real KV sources are rejected.""" + tensor = torch.zeros(4, 16, dtype=torch.uint8, device=self.device) + sources = tuple( + RealKvSource( + tensor=tensor, page_size=1, num_bytes_per_token=16, read_bytes=16 + ) + for _ in range(MAX_REAL_KV_SOURCES + 1) + ) + canary_buf = torch.zeros( + 4, CANARY_SLOT_BYTES, dtype=torch.uint8, device=self.device + ) + plan = VerifyPlan.allocate(verify_capacity=1, device=self.device) + violation_ring = torch.zeros(2, 8, dtype=torch.int64, device=self.device) + violation_write_index = torch.zeros(1, dtype=torch.int32, device=self.device) + slot_run_counter = torch.zeros(1, dtype=torch.int64, device=self.device) + kernel_run_counter = torch.zeros(1, dtype=torch.int64, device=self.device) + + enable_chain_position_assert = torch.ones( + 1, dtype=torch.int32, device=self.device + ) + + with self.assertRaises(ValueError): + launch_canary_verify_kernel( + context=VerifyOrWriteContext( + canary_buf=canary_buf, + kernel_kind=CanaryLaunchTag.HEAD_K_FULL, + violation_ring=violation_ring, + violation_write_index=violation_write_index, + slot_run_counter=slot_run_counter, + kernel_run_counter=kernel_run_counter, + real_kv_sources=sources, + real_kv_hash_mode=RealKvHashMode.NONE, + enable_chain_position_assert=enable_chain_position_assert, + ), + plan=plan, + check_verify_expected_token=True, + ) + + class TestPoolPatcherBufferInfos(PoolPatcherHelper, CustomTestCase): def test_get_contiguous_buf_infos_inserts_canary_entries(self): """Verify contiguous buffer metadata includes canary entries after patching.""" diff --git a/test/registered/kv_canary/test_self_unit_runner_per_forward.py b/test/registered/kv_canary/test_self_unit_runner_per_forward.py index d981f686b..de5ebcf32 100644 --- a/test/registered/kv_canary/test_self_unit_runner_per_forward.py +++ b/test/registered/kv_canary/test_self_unit_runner_per_forward.py @@ -5,6 +5,7 @@ from unittest.mock import patch import torch +from sglang.jit_kernel.kv_canary.consts import RealKvHashMode from sglang.jit_kernel.kv_canary.verify import CanaryLaunchTag, VerifyPlan from sglang.jit_kernel.kv_canary.write import WritePlan from sglang.srt.kv_canary import endpoint as endpoint_module @@ -96,6 +97,7 @@ class TestLaunchEndpointsPerForward(CanaryManagerTestCase): forward_batch=forward_batch, expected_inputs=ExpectedInputs.allocate(capacity=3, device=self.device), violation_log=ViolationLog.allocate(ring_capacity=2, device=self.device), + real_kv_hash_mode=RealKvHashMode.NONE, enable_write_input_assert=False, enable_verify_token_assert=False, ) @@ -148,6 +150,7 @@ class TestLaunchEndpointsPerForward(CanaryManagerTestCase): forward_batch=forward_batch, expected_inputs=ExpectedInputs.allocate(capacity=1, device=self.device), violation_log=ViolationLog.allocate(ring_capacity=2, device=self.device), + real_kv_hash_mode=RealKvHashMode.NONE, enable_write_input_assert=False, enable_verify_token_assert=False, ) @@ -185,6 +188,7 @@ class TestLaunchEndpointsPerForward(CanaryManagerTestCase): forward_batch=forward_batch, expected_inputs=ExpectedInputs.allocate(capacity=3, device=self.device), violation_log=ViolationLog.allocate(ring_capacity=2, device=self.device), + real_kv_hash_mode=RealKvHashMode.NONE, enable_write_input_assert=False, enable_verify_token_assert=True, ) @@ -220,6 +224,7 @@ class TestLaunchEndpointsPerForward(CanaryManagerTestCase): forward_batch=forward_batch, expected_inputs=ExpectedInputs.allocate(capacity=1, device=self.device), violation_log=ViolationLog.allocate(ring_capacity=2, device=self.device), + real_kv_hash_mode=RealKvHashMode.NONE, enable_write_input_assert=False, enable_verify_token_assert=False, ) diff --git a/test/registered/kv_canary/test_self_unit_violation.py b/test/registered/kv_canary/test_self_unit_violation.py index ac272cd19..6a9ff778a 100644 --- a/test/registered/kv_canary/test_self_unit_violation.py +++ b/test/registered/kv_canary/test_self_unit_violation.py @@ -53,6 +53,7 @@ class TestViolationReporter(CustomTestCase): fail_reason_bits=int( FailReason.VERIFY_CHAIN_HASH_MISMATCH | FailReason.VERIFY_POSITION_MISMATCH + | FailReason.VERIFY_REAL_KV_HASH_MISMATCH ), ) out = _format_violation( @@ -60,12 +61,12 @@ class TestViolationReporter(CustomTestCase): ) self.assertEqual( out, - "kv_canary violation: launch_tag=HEAD_K_FULL fail_reason=verify_chain_hash+verify_position " + "kv_canary violation: launch_tag=HEAD_K_FULL fail_reason=verify_chain_hash+verify_position+verify_real_kv_hash " "slot_idx=17 position=42 stored_token=111 expected_token=0 stored_chain_hash=0x1111111111111111 " "expected_aux=0x2222222222222222\n" "KV cache canary violation detected (kernel_kind=HEAD_K_FULL, slot_idx=17, position=42)\n" "canary_kind: per_forward_head_k_full\n" - " fail_reasons: verify_chain_hash verify_position\n" + " fail_reasons: verify_chain_hash verify_position verify_real_kv_hash\n" " stored: token_id=111 position=42 prev_hash=0x1111111111111111\n" " expected: prev_hash=0x2222222222222222\n" " total_violations=1 ring_overflow=False step_when_pumped=7", @@ -199,7 +200,7 @@ class TestLogOrRaiseViolation(CustomTestCase): self.assertFalse(reporter.is_raised) def test_log_mode_emits_one_warning_per_violation(self) -> None: - """Log mode with 2 valid rows emits 2 warnings, each a full _format_violation snapshot for that row.""" + """Log mode with 3 valid rows emits 3 warnings, each a full _format_violation snapshot for that row.""" rows = [ _make_row( slot_idx=11, @@ -211,14 +212,19 @@ class TestLogOrRaiseViolation(CustomTestCase): position=202, fail_reason_bits=int(FailReason.VERIFY_POSITION_MISMATCH), ), + _make_row( + slot_idx=33, + position=303, + fail_reason_bits=int(FailReason.VERIFY_REAL_KV_HASH_MISMATCH), + ), ] reporter = _make_reporter( - rows=rows, write_index=2, ring_capacity=4, mode=CanaryMode.LOG + rows=rows, write_index=3, ring_capacity=4, mode=CanaryMode.LOG ) with patch.object(violation_reporter_module.logger, "warning") as mock_warning: reporter.log_or_raise_violation(outer_step_counter=7) - self.assertEqual(mock_warning.call_count, 2) + self.assertEqual(mock_warning.call_count, 3) messages: list[str] = [call.args[0] for call in mock_warning.call_args_list] self.assertEqual( messages[0], @@ -230,7 +236,7 @@ class TestLogOrRaiseViolation(CustomTestCase): " fail_reasons: verify_chain_hash\n" " stored: token_id=111 position=101 prev_hash=0x0000000000000000\n" " expected: prev_hash=0x0000000000000000\n" - " total_violations=2 ring_overflow=False step_when_pumped=7", + " total_violations=3 ring_overflow=False step_when_pumped=7", ) self.assertEqual( messages[1], @@ -242,12 +248,24 @@ class TestLogOrRaiseViolation(CustomTestCase): " fail_reasons: verify_position\n" " stored: token_id=111 position=202 prev_hash=0x0000000000000000\n" " expected: prev_hash=0x0000000000000000\n" - " total_violations=2 ring_overflow=False step_when_pumped=7", + " total_violations=3 ring_overflow=False step_when_pumped=7", + ) + self.assertEqual( + messages[2], + "kv_canary violation: launch_tag=HEAD_K_FULL fail_reason=verify_real_kv_hash slot_idx=33 position=303 " + "stored_token=111 expected_token=0 stored_chain_hash=0x0000000000000000 " + "expected_aux=0x0000000000000000\n" + "KV cache canary violation detected (kernel_kind=HEAD_K_FULL, slot_idx=33, position=303)\n" + "canary_kind: per_forward_head_k_full\n" + " fail_reasons: verify_real_kv_hash\n" + " stored: token_id=111 position=303 prev_hash=0x0000000000000000\n" + " expected: prev_hash=0x0000000000000000\n" + " total_violations=3 ring_overflow=False step_when_pumped=7", ) self.assertFalse(reporter.is_raised) def test_raise_mode_raises_one_error_containing_all_violations(self) -> None: - """Raise mode raises a single RuntimeError whose text is the 2 formatted rows joined with single newlines.""" + """Raise mode raises a single RuntimeError whose text is the 3 formatted rows joined with single newlines.""" rows = [ _make_row( slot_idx=11, @@ -259,9 +277,14 @@ class TestLogOrRaiseViolation(CustomTestCase): position=202, fail_reason_bits=int(FailReason.VERIFY_POSITION_MISMATCH), ), + _make_row( + slot_idx=33, + position=303, + fail_reason_bits=int(FailReason.VERIFY_REAL_KV_HASH_MISMATCH), + ), ] reporter = _make_reporter( - rows=rows, write_index=2, ring_capacity=4, mode=CanaryMode.RAISE + rows=rows, write_index=3, ring_capacity=4, mode=CanaryMode.RAISE ) with self.assertRaises(RuntimeError) as ctx: reporter.log_or_raise_violation(outer_step_counter=5) @@ -276,7 +299,7 @@ class TestLogOrRaiseViolation(CustomTestCase): " fail_reasons: verify_chain_hash\n" " stored: token_id=111 position=101 prev_hash=0x0000000000000000\n" " expected: prev_hash=0x0000000000000000\n" - " total_violations=2 ring_overflow=False step_when_pumped=5\n" + " total_violations=3 ring_overflow=False step_when_pumped=5\n" "kv_canary violation: launch_tag=HEAD_K_FULL fail_reason=verify_position slot_idx=22 position=202 " "stored_token=111 expected_token=0 stored_chain_hash=0x0000000000000000 " "expected_aux=0x0000000000000000\n" @@ -285,7 +308,16 @@ class TestLogOrRaiseViolation(CustomTestCase): " fail_reasons: verify_position\n" " stored: token_id=111 position=202 prev_hash=0x0000000000000000\n" " expected: prev_hash=0x0000000000000000\n" - " total_violations=2 ring_overflow=False step_when_pumped=5", + " total_violations=3 ring_overflow=False step_when_pumped=5\n" + "kv_canary violation: launch_tag=HEAD_K_FULL fail_reason=verify_real_kv_hash slot_idx=33 position=303 " + "stored_token=111 expected_token=0 stored_chain_hash=0x0000000000000000 " + "expected_aux=0x0000000000000000\n" + "KV cache canary violation detected (kernel_kind=HEAD_K_FULL, slot_idx=33, position=303)\n" + "canary_kind: per_forward_head_k_full\n" + " fail_reasons: verify_real_kv_hash\n" + " stored: token_id=111 position=303 prev_hash=0x0000000000000000\n" + " expected: prev_hash=0x0000000000000000\n" + " total_violations=3 ring_overflow=False step_when_pumped=5", ) self.assertTrue(reporter.is_raised)