From 70e983a2da60c560f4193e1e80927b68474be72a Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Sun, 31 May 2026 09:53:34 +0800 Subject: [PATCH] Add the KV-canary write JIT kernel and reference implementation (#26806) --- .../benchmark/kv_canary/bench_write.py | 289 ++++++++++ .../csrc/kv_canary/canary_write.cuh | 287 ++++++++++ python/sglang/jit_kernel/kv_canary/write.py | 241 ++++++++ .../sglang/jit_kernel/kv_canary/write_ref.py | 201 +++++++ .../tests/kv_canary/_canary_helpers.py | 330 +++++++++++ .../jit_kernel/tests/kv_canary/_fixtures.py | 150 +++++ .../jit_kernel/tests/kv_canary/_invariants.py | 520 ++++++++++++++++++ 7 files changed, 2018 insertions(+) create mode 100644 python/sglang/jit_kernel/benchmark/kv_canary/bench_write.py create mode 100644 python/sglang/jit_kernel/csrc/kv_canary/canary_write.cuh create mode 100644 python/sglang/jit_kernel/kv_canary/write.py create mode 100644 python/sglang/jit_kernel/kv_canary/write_ref.py create mode 100644 python/sglang/jit_kernel/tests/kv_canary/_canary_helpers.py create mode 100644 python/sglang/jit_kernel/tests/kv_canary/_fixtures.py create mode 100644 python/sglang/jit_kernel/tests/kv_canary/_invariants.py diff --git a/python/sglang/jit_kernel/benchmark/kv_canary/bench_write.py b/python/sglang/jit_kernel/benchmark/kv_canary/bench_write.py new file mode 100644 index 000000000..1721f4a3b --- /dev/null +++ b/python/sglang/jit_kernel/benchmark/kv_canary/bench_write.py @@ -0,0 +1,289 @@ +from __future__ import annotations + +from typing import Tuple + +import torch +import triton +import triton.testing + +from sglang.jit_kernel.benchmark.kv_canary.utils import ( + RING_CAPACITY, + SWA_WINDOW, + BenchCase, + build_fast_matrix_cases, + build_full_matrix_cases, + cases_to_x_vals, + naive_slot_copy_fn, +) +from sglang.jit_kernel.benchmark.utils import ( + DEFAULT_DEVICE, + get_benchmark_range, + run_benchmark, +) +from sglang.jit_kernel.kv_canary import consts +from sglang.jit_kernel.kv_canary.verify import ( + CANARY_SLOT_BYTES, + CanaryLaunchTag, + VerifyOrWriteContext, +) +from sglang.jit_kernel.kv_canary.write import WritePlan, launch_canary_write_kernel +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=900, suite="nightly-kernel-1-gpu", nightly=True) + + +_X_NAMES = [ + "scenario", + "bs", + "prefix_len", + "mode", + "extend_len", + "pool_kind", +] +_X_VALS = cases_to_x_vals( + get_benchmark_range( + full_range=build_full_matrix_cases(), + ci_range=build_fast_matrix_cases(), + ) +) + +_KERNEL_KIND_X_NAMES = ["kernel_kind_name", "enable_write_verify_inputs_name"] +_KERNEL_KIND_X_VALS = [ + (tag.name, str(enable)) for tag in CanaryLaunchTag for enable in (False, True) +] + + +def _write_entry_count(case: BenchCase) -> int: + return case.bs * case.extend_len + + +def _write_num_slots(case: BenchCase) -> int: + per_req_slots = max( + SWA_WINDOW if case.pool_kind == "swa_window_128" else 1, + case.prefix_len + case.extend_len, + ) + return max(2, case.bs * per_req_slots + 1) + + +def _build_write_inputs( + case: BenchCase, *, device: torch.device, mirror_expected_inputs: bool = False +) -> dict: + total_entries = _write_entry_count(case) + num_tokens_padded = max(1, total_entries) + + per_req_slots = max( + SWA_WINDOW if case.pool_kind == "swa_window_128" else 1, + case.prefix_len + case.extend_len, + ) + num_slots = _write_num_slots(case) + + canary_buf = torch.zeros( + num_slots, CANARY_SLOT_BYTES, dtype=torch.uint8, device=device + ) + + write_offsets = torch.zeros(case.bs + 1, dtype=torch.int64, device=device) + if case.bs > 0: + offsets_host = torch.arange(0, case.bs + 1, dtype=torch.int64) * case.extend_len + write_offsets.copy_(offsets_host.to(device)) + + write_seed_slots = torch.empty(case.bs, dtype=torch.int64, device=device) + if case.bs > 0: + if case.prefix_len == 0: + write_seed_slots.fill_(-1) + else: + per_req_stride = per_req_slots + seeds = ( + torch.arange(case.bs, dtype=torch.int32, device=device) * per_req_stride + + case.prefix_len + - 1 + ) + write_seed_slots.copy_(seeds.to(torch.int64)) + + write_num_valid_reqs = torch.tensor([case.bs], dtype=torch.int32, device=device) + + plan = WritePlan( + write_offsets=write_offsets, + write_seed_slot_indices=write_seed_slots, + write_num_valid_reqs=write_num_valid_reqs, + ) + + input_ids = torch.zeros(num_tokens_padded, dtype=torch.int64, device=device) + positions = torch.zeros(num_tokens_padded, dtype=torch.int64, device=device) + out_cache_loc = torch.zeros(num_tokens_padded, dtype=torch.int64, device=device) + if total_entries > 0: + flat_idx = torch.arange(total_entries, device=device, dtype=torch.int64) + per_req_idx = flat_idx % max(case.extend_len, 1) + req_idx = flat_idx // max(case.extend_len, 1) + per_req_stride = per_req_slots + slots = (req_idx * per_req_stride + case.prefix_len + per_req_idx) % max( + num_slots, 1 + ) + input_ids[:total_entries] = (flat_idx % 32768).to(torch.int64) + positions[:total_entries] = (case.prefix_len + per_req_idx).to(torch.int64) + out_cache_loc[:total_entries] = slots.to(torch.int64) + + if case.pool_kind == "swa_window_128": + full_to_swa = torch.arange(num_slots + 1, dtype=torch.int64, device=device) + full_to_swa[-1] = -1 + out_cache_loc = full_to_swa[out_cache_loc] + + if mirror_expected_inputs: + expected_input_tokens = input_ids.clone() + expected_input_positions = positions.clone() + else: + expected_input_tokens = None + expected_input_positions = None + + violation_ring = torch.zeros( + RING_CAPACITY, consts.VIOLATION_FIELDS, dtype=torch.int64, device=device + ) + violation_write_index = torch.zeros(1, dtype=torch.int32, device=device) + slot_run_counter = torch.zeros(1, dtype=torch.int64, device=device) + kernel_run_counter = torch.zeros(1, dtype=torch.int64, device=device) + enable_chain_position_assert = torch.ones(1, dtype=torch.int32, device=device) + + return dict( + canary_buf=canary_buf, + plan=plan, + input_ids=input_ids, + positions=positions, + out_cache_loc=out_cache_loc, + expected_input_tokens=expected_input_tokens, + expected_input_positions=expected_input_positions, + violation_ring=violation_ring, + violation_write_index=violation_write_index, + slot_run_counter=slot_run_counter, + kernel_run_counter=kernel_run_counter, + enable_chain_position_assert=enable_chain_position_assert, + ) + + +def _build_context( + *, + inputs: dict, + kernel_kind: CanaryLaunchTag, +) -> VerifyOrWriteContext: + return VerifyOrWriteContext( + canary_buf=inputs["canary_buf"], + kernel_kind=kernel_kind, + violation_ring=inputs["violation_ring"], + violation_write_index=inputs["violation_write_index"], + slot_run_counter=inputs["slot_run_counter"], + kernel_run_counter=inputs["kernel_run_counter"], + enable_chain_position_assert=inputs["enable_chain_position_assert"], + ) + + +@triton.testing.perf_report( + triton.testing.Benchmark( + x_names=_X_NAMES, + x_vals=_X_VALS, + line_arg="provider", + line_vals=["canary", "naive"], + line_names=["canary_write_step", "naive index_copy_"], + styles=[("blue", "-"), ("red", "--")], + ylabel="us", + plot_name="kv-canary-write-perf", + args={}, + ) +) +def benchmark( + scenario: str, + bs: int, + prefix_len: int, + mode: str, + extend_len: int, + pool_kind: str, + provider: str, +) -> Tuple[float, float, float]: + case = BenchCase( + scenario=scenario, + bs=bs, + prefix_len=prefix_len, + mode=mode, + extend_len=extend_len, + pool_kind=pool_kind, + ) + device = torch.device(DEFAULT_DEVICE) + + if provider == "canary": + inputs = _build_write_inputs(case, device=device) + context = _build_context( + inputs=inputs, + kernel_kind=CanaryLaunchTag.HEAD_K_FULL, + ) + + def fn() -> None: + launch_canary_write_kernel( + context=context, + plan=inputs["plan"], + input_ids=inputs["input_ids"], + positions=inputs["positions"], + out_cache_loc=inputs["out_cache_loc"], + enable_write_input_assert=False, + expected_input_tokens=inputs["expected_input_tokens"], + expected_input_positions=inputs["expected_input_positions"], + ) + + else: + fn = naive_slot_copy_fn(total=_write_entry_count(case), device=device) + + return run_benchmark(fn) + + +@triton.testing.perf_report( + triton.testing.Benchmark( + x_names=_KERNEL_KIND_X_NAMES, + x_vals=_KERNEL_KIND_X_VALS, + line_arg="provider", + line_vals=["canary"], + line_names=["canary_write_step"], + styles=[("blue", "-")], + ylabel="us", + plot_name="kv-canary-write-kernel-kind-perf", + args={}, + ) +) +def benchmark_kernel_kind( + kernel_kind_name: str, + enable_write_verify_inputs_name: str, + provider: str, +) -> Tuple[float, float, float]: + case = BenchCase( + scenario="kernel_kind", + bs=32, + prefix_len=4096, + mode="extend", + extend_len=128, + pool_kind="full", + ) + device = torch.device(DEFAULT_DEVICE) + + enable_write_verify_inputs = enable_write_verify_inputs_name == "True" + inputs = _build_write_inputs( + case, device=device, mirror_expected_inputs=enable_write_verify_inputs + ) + kernel_kind = CanaryLaunchTag[kernel_kind_name] + context = _build_context( + inputs=inputs, + kernel_kind=kernel_kind, + ) + + def fn() -> None: + launch_canary_write_kernel( + context=context, + plan=inputs["plan"], + input_ids=inputs["input_ids"], + positions=inputs["positions"], + out_cache_loc=inputs["out_cache_loc"], + enable_write_input_assert=enable_write_verify_inputs, + expected_input_tokens=inputs["expected_input_tokens"], + expected_input_positions=inputs["expected_input_positions"], + ) + + return run_benchmark(fn) + + +if __name__ == "__main__": + benchmark.run(print_data=True) + benchmark_kernel_kind.run(print_data=True) diff --git a/python/sglang/jit_kernel/csrc/kv_canary/canary_write.cuh b/python/sglang/jit_kernel/csrc/kv_canary/canary_write.cuh new file mode 100644 index 000000000..1401e11ff --- /dev/null +++ b/python/sglang/jit_kernel/csrc/kv_canary/canary_write.cuh @@ -0,0 +1,287 @@ +#pragma once + +#include // For TensorMatcher, SymbolicSize, SymbolicDevice +#include // For RuntimeCheck + +#include // For LaunchKernel, SGL_DEVICE + +#include +#include + +#include "canary_common.cuh" +#include + +namespace canary { + +namespace { + +// Single thread per block — chain advance is inherently serial. +constexpr uint32_t kWriteBlockSize = 1; + +struct WriteKernelParams { + uint8_t* canary_buf; + int64_t slot_stride_bytes; + + // Plan tensors. + const int64_t* write_offsets; + const int64_t* write_seed_slot_indices; + const int32_t* write_num_valid_reqs; + int32_t write_req_capacity; + + // ForwardBatch passthroughs. out_cache_loc is caller-pre-translated for SWA groups; the kernel + // treats it opaquely and skips entries with slot < 0. + const int64_t* input_ids; + const int64_t* positions; + const int64_t* out_cache_loc; + + // Pseudo-mode oracle inputs. + bool enable_write_input_assert; + const int64_t* expected_input_tokens; + const int64_t* expected_input_positions; + + // Violation sink (ring + write_index + capacity + kernel_kind bundled in canary_common.cuh). + ViolationSink violation_sink; + + // Health counters. + int64_t* slot_run_counter; + int64_t* kernel_run_counter; + + // 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; +}; + +__global__ void canary_write_kernel(const WriteKernelParams __grid_constant__ p) { + const uint32_t r = blockIdx.x; + + // Unconditional kernel_run_counter bump (block 0 is always present). + if (r == 0) { + atomicAdd(reinterpret_cast(p.kernel_run_counter), 1ULL); + } + + const int32_t active = *p.write_num_valid_reqs; + if (r >= static_cast(active)) { + return; + } + + const int64_t entry_start = p.write_offsets[r]; + const int64_t entry_end = p.write_offsets[r + 1]; + const int64_t entry_count = entry_end - entry_start; + if (entry_count <= 0) { + return; + } + + const int64_t seed_slot_idx = p.write_seed_slot_indices[r]; + + // Initialize running_prev_hash by advancing the chain from the seed slot + uint64_t running_prev_hash = + compute_slot_hash(p.canary_buf, p.slot_stride_bytes, static_cast(seed_slot_idx)); + + // Assumes eagle topk=1 (linear chain). Under topk>1 target_verify would be a tree and + // sibling positions share parent.pos+1, breaking this invariant. + const bool do_chain_position_assert = (seed_slot_idx >= 0) && (*p.enable_chain_position_assert != 0); + int64_t running_prev_position = 0; + if (do_chain_position_assert) { + running_prev_position = canary_load_field(p.canary_buf, seed_slot_idx, p.slot_stride_bytes, kCanaryFieldPosition); + } + + int64_t entries_written = 0; + for (int64_t entry_offset = 0; entry_offset < entry_count; ++entry_offset) { + const int64_t entry_idx = entry_start + entry_offset; + const int64_t slot = p.out_cache_loc[entry_idx]; + + if (slot < 0) { + continue; + } + ++entries_written; + + 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; + + if (p.enable_write_input_assert) { + const int64_t expected_token = p.expected_input_tokens[entry_idx]; + const int64_t expected_position = p.expected_input_positions[entry_idx]; + FailReason mismatch_bits{}; + if (token != expected_token) { + mismatch_bits |= FailReason::kWriteTokenMismatch; + } + if (position != expected_position) { + mismatch_bits |= FailReason::kWritePositionMismatch; + } + if (mismatch_bits != FailReason{}) { + record_violation( + p.violation_sink, + ViolationRow{ + /* slot_idx = */ slot, + /* position = */ position, + /* stored_token = */ token, + /* expected_token = */ expected_token, + /* stored_chain_hash = */ static_cast(running_prev_hash), + /* expected_aux = expected_position */ expected_position, + /* fail_reason_bits = */ static_cast(mismatch_bits), + }); + } + } + + if (do_chain_position_assert) { + const int64_t expected_position_chain = running_prev_position + 1; + if (position != expected_position_chain) { + record_violation( + p.violation_sink, + ViolationRow{ + /* slot_idx = */ slot, + /* position = */ position, + /* stored_token = */ token, + /* expected_token = */ token, + /* stored_chain_hash = */ static_cast(running_prev_hash), + /* expected_aux = expected_position */ expected_position_chain, + /* fail_reason_bits = */ static_cast(FailReason::kWritePositionMismatch), + }); + } + running_prev_position = position; + } + + canary_store_field(p.canary_buf, slot, p.slot_stride_bytes, kCanaryFieldToken, token); + canary_store_field(p.canary_buf, slot, p.slot_stride_bytes, kCanaryFieldPosition, position); + canary_store_field( + p.canary_buf, slot, p.slot_stride_bytes, kCanaryFieldPrevHash, static_cast(running_prev_hash)); + canary_store_field(p.canary_buf, slot, p.slot_stride_bytes, kCanaryFieldRealKvHash, real_kv_hash); + + running_prev_hash = + splitmix64_mix3(running_prev_hash, static_cast(token), static_cast(position)); + } + + // Each block contributes its non-skipped entry count to slot_run_counter once at exit. + atomicAdd( + reinterpret_cast(p.slot_run_counter), static_cast(entries_written)); +} + +} // namespace + +// API source of truth: docstring of canary_write_step in python/sglang/jit_kernel/kv_canary/write.py. +// +// ABI notes: +// - 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( + tvm::ffi::TensorView canary_buf, + tvm::ffi::TensorView write_offsets, + tvm::ffi::TensorView write_seed_slot_indices, + tvm::ffi::TensorView write_num_valid_reqs, + tvm::ffi::TensorView input_ids, + tvm::ffi::TensorView positions, + tvm::ffi::TensorView out_cache_loc, + int64_t kernel_kind, + int64_t enable_write_input_assert, + const tvm::ffi::Optional expected_input_tokens, + const tvm::ffi::Optional expected_input_positions, + 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 enable_chain_position_assert) { + using namespace host; + + SymbolicSize N_slots = {"num_canary_slots"}; + SymbolicSize N_stride = {"slot_stride_bytes"}; + SymbolicSize N_write_reqs = {"write_req_capacity"}; + SymbolicSize N_tokens = {"num_tokens_padded"}; + SymbolicDevice device_; + device_.set_options(); + + TensorMatcher({N_slots, N_stride}).with_dtype().with_device(device_).verify(canary_buf); + + // write_offsets has shape [write_req_capacity + 1]; the length relationship is pinned by the + // RuntimeCheck below, this matcher pins dtype + device. + SymbolicSize N_write_offsets = {"write_offsets_len"}; + TensorMatcher({N_write_offsets}).with_dtype().with_device(device_).verify(write_offsets); + TensorMatcher({N_write_reqs}).with_dtype().with_device(device_).verify(write_seed_slot_indices); + TensorMatcher({1}).with_dtype().with_device(device_).verify(write_num_valid_reqs); + + TensorMatcher({N_tokens}) + .with_dtype() + .with_device(device_) + .verify(input_ids) + .verify(positions) + .verify(out_cache_loc); + const bool enable_write_input_assert_bool = (enable_write_input_assert != 0); + RuntimeCheck( + enable_write_input_assert_bool == expected_input_tokens.has_value(), + "canary_write: expected_input_tokens presence must match enable_write_input_assert"); + RuntimeCheck( + enable_write_input_assert_bool == expected_input_positions.has_value(), + "canary_write: expected_input_positions presence must match enable_write_input_assert"); + if (enable_write_input_assert_bool) { + TensorMatcher({N_tokens}) + .with_dtype() + .with_device(device_) + .verify(expected_input_tokens.value()) + .verify(expected_input_positions.value()); + } + + TensorMatcher({1}).with_dtype().with_device(device_).verify(violation_write_index); + SymbolicSize N_ring = {"ring_capacity"}; + TensorMatcher({N_ring, static_cast(kViolationFields)}) + .with_dtype() + .with_device(device_) + .verify(violation_ring); + TensorMatcher({1}) + .with_dtype() + .with_device(device_) + .verify(slot_run_counter) + .verify(kernel_run_counter); + TensorMatcher({1}).with_dtype().with_device(device_).verify(enable_chain_position_assert); + + 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()); + const DLDevice device = device_.unwrap(); + + RuntimeCheck( + write_offsets.size(0) == static_cast(write_req_capacity) + 1, + "canary_write: write_offsets.size(0) must equal write_req_capacity + 1 (", + static_cast(write_req_capacity) + 1, + "), got ", + write_offsets.size(0)); + RuntimeCheck( + slot_stride_bytes >= static_cast(kCanaryFieldsPerSlot * sizeof(int64_t)), + "canary_write: slot_stride_bytes must hold at least ", + static_cast(kCanaryFieldsPerSlot * sizeof(int64_t)), + " bytes per slot, got ", + slot_stride_bytes); + + WriteKernelParams p{}; + p.canary_buf = static_cast(canary_buf.data_ptr()); + p.slot_stride_bytes = slot_stride_bytes; + p.write_offsets = static_cast(write_offsets.data_ptr()); + p.write_seed_slot_indices = static_cast(write_seed_slot_indices.data_ptr()); + p.write_num_valid_reqs = static_cast(write_num_valid_reqs.data_ptr()); + p.write_req_capacity = write_req_capacity; + p.input_ids = static_cast(input_ids.data_ptr()); + p.positions = static_cast(positions.data_ptr()); + p.out_cache_loc = static_cast(out_cache_loc.data_ptr()); + p.enable_write_input_assert = enable_write_input_assert_bool; + p.expected_input_tokens = + enable_write_input_assert_bool ? static_cast(expected_input_tokens.value().data_ptr()) : nullptr; + p.expected_input_positions = enable_write_input_assert_bool + ? static_cast(expected_input_positions.value().data_ptr()) + : nullptr; + p.violation_sink.ring = static_cast(violation_ring.data_ptr()); + p.violation_sink.write_index = static_cast(violation_write_index.data_ptr()); + p.violation_sink.ring_capacity = ring_capacity; + p.violation_sink.kernel_kind = static_cast(kernel_kind); + p.slot_run_counter = static_cast(slot_run_counter.data_ptr()); + p.kernel_run_counter = static_cast(kernel_run_counter.data_ptr()); + p.enable_chain_position_assert = static_cast(enable_chain_position_assert.data_ptr()); + + // 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. + const uint32_t grid = write_req_capacity == 0 ? 1u : static_cast(write_req_capacity); + LaunchKernel(grid, kWriteBlockSize, device)(canary_write_kernel, p); +} + +} // namespace canary diff --git a/python/sglang/jit_kernel/kv_canary/write.py b/python/sglang/jit_kernel/kv_canary/write.py new file mode 100644 index 000000000..80e6c89eb --- /dev/null +++ b/python/sglang/jit_kernel/kv_canary/write.py @@ -0,0 +1,241 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING + +import torch + +from sglang.jit_kernel.kv_canary.verify import ( + VerifyOrWriteContext, + _assert_contiguous, +) +from sglang.jit_kernel.utils import cache_once, load_jit + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@dataclass(frozen=True, slots=True, kw_only=True) +class WritePlan: + """Write plan consumed by launch_canary_write_kernel: per-token slot indices + per-req metadata. + + Fully per-req — no per-token tile. launch_canary_write_kernel uses write_offsets to map each thread's + (req, j) into a flat index i, then reads token-level data from input_ids / positions / + out_cache_loc[i] directly. + SWA translation of per-token slots is done **host-side by the caller** (typically the endpoint) before + invoking launch_canary_write_kernel — the kernel is SWA-agnostic and only understands "slot ≥ 0 ⇒ write; + slot < 0 ⇒ skip this entry". Only the chain-seed slot (a per-req gather from req_to_token at plan time) + is SWA-translated by the plan kernel and lives in write_seed_slot_indices. + + Req r's write entries occupy flat indices [write_offsets[r], write_offsets[r+1]). seed_slot_idx == -1 means + K_req_old == 0 (anchor on CANARY_CHAIN_ANCHOR). + + Fields: + write_offsets: Exclusive prefix-sum offsets indexing into ForwardBatch's input_ids / positions / + out_cache_loc, shape [write_req_capacity + 1], int64. write_offsets[0] == 0; + write_offsets[write_num_valid_reqs[0]] == total_write_entries. + write_seed_slot_indices: Chain-seed slot per write req, shape [write_req_capacity], int64. Already + SWA-translated. -1 = no prefix (chain anchors on CANARY_CHAIN_ANCHOR). + write_num_valid_reqs: Active write-req count, shape [1], int32. launch_canary_write_kernel skips blocks + with block_id >= write_num_valid_reqs[0]. + """ + + write_offsets: torch.Tensor + write_seed_slot_indices: torch.Tensor + write_num_valid_reqs: torch.Tensor + + @classmethod + def allocate( + cls, + *, + write_req_capacity: int, + device: torch.device, + ) -> "WritePlan": + if write_req_capacity <= 0: + raise ValueError( + f"kv-canary: WritePlan write_req_capacity must be positive, got {write_req_capacity}" + ) + return cls( + write_offsets=torch.empty( + write_req_capacity + 1, dtype=torch.int64, device=device + ), + write_seed_slot_indices=torch.empty( + write_req_capacity, dtype=torch.int64, device=device + ), + write_num_valid_reqs=torch.empty(1, dtype=torch.int32, device=device), + ) + + def zero_for_testing_(self) -> "WritePlan": + """WARN: ONLY use it when testing plan kernel. Do not use it when testing verify or + write kernel to avoid hiding bugs.""" + self.write_offsets.zero_() + self.write_seed_slot_indices.zero_() + self.write_num_valid_reqs.zero_() + return self + + +def launch_canary_write_kernel( + *, + context: VerifyOrWriteContext, + plan: WritePlan, + input_ids: torch.Tensor, + positions: torch.Tensor, + out_cache_loc: torch.Tensor, + enable_write_input_assert: bool, + expected_input_tokens: torch.Tensor | None, + expected_input_positions: torch.Tensor | None, +) -> None: + """Write canary fingerprints into one canary buffer per a WritePlan. + + Grid: one CUDA block per active write req, single thread per block (chain is intrinsically serial). + Block r walks entries ``[plan.write_offsets[r], plan.write_offsets[r+1])``. Per chain step ``i``: + + - ``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]). + - 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)``. + ``real_kv_hash`` is intentionally not folded into the chain — see ``compute_slot_hash`` in + ``csrc/kv_canary/canary_common.cuh`` for the radix-folding rationale. + + Initial ``running_prev_hash`` when ``seed_slot_idx >= 0``: load (token, position, prev_hash) from + ``canary_buf[plan.write_seed_slot_indices[r]]`` and set + ``running_prev_hash = splitmix64_mix3(seed.prev_hash, seed.token, seed.position)`` + (i.e. apply the same advance step that produced ``seed``'s successor — this keeps slot[0]'s stored + ``prev_hash`` consistent with the chain link). Else + ``running_prev_hash = splitmix64(CANARY_CHAIN_ANCHOR)``. ``write_seed_slot_indices`` is already + SWA-translated by the plan kernel; ``CANARY_CHAIN_ANCHOR`` is hardcoded module-level (no runtime seed). + + Write-time input verification (caller-driven, kernel is oracle-agnostic): when + ``enable_write_input_assert`` is True the kernel additionally compares ``input_ids[i]`` against + ``expected_input_tokens[i]`` and ``positions[i]`` against ``expected_input_positions[i]``; mismatch + on either field records a violation. The chain still advances on the actual values (not the expected + ones) so a downstream verify won't cascade. Whoever produced the expected tensors is responsible for + filling them; the kernel runs no oracle internally. + + Write only writes canary_buf (reads only at seed slots). Block uses no shared memory. + + The ForwardBatch-derived arguments are passed through unchanged from the source ForwardBatch — canary does not transform + them. + + Args: + context: Shared verify/write launch context, including canary buffer, launch tag, violation sink, + and health counters. + 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 + plan.write_offsets[plan.write_num_valid_reqs[0]] is cuda-graph padding. + positions: ForwardBatch.positions; sequence positions of input_ids, shape [num_tokens_padded], int64. + out_cache_loc: Per-token canary slot index, shape [num_tokens_padded], int64. The caller is + responsible for translating ForwardBatch.out_cache_loc into the canary's index space for SWA + groups (typically a host-side LUT gather in the endpoint); FULL groups pass it through + unchanged. A -1 entry signals skip-this-token (used for SWA out-of-window slots or padding). + The kernel does not consult any LUT. + enable_write_input_assert: bool toggle. False = expected_input_* tensors must be None. True = compare + each chain step's actual (token, position) against the caller-supplied expected tensors below. + expected_input_tokens: Expected token id per write entry, shape [num_tokens_padded], int64. Only read + when enable_write_input_assert is True; must be None when enable_write_input_assert is False. + Layout mirrors input_ids (flattened across reqs in plan.write_offsets order); padding tail + is ignored. Filled by the caller from whichever oracle produces expected inputs — the kernel + knows no oracle. + expected_input_positions: Expected position per write entry, shape [num_tokens_padded], int64, or None. + Same shape/layout/lifetime rules as expected_input_tokens. + + Implementation: + - CUDA __global__ `canary_write_kernel`: 1-D grid `(write_req_capacity, 1, 1)` blocks × `(1, 1, 1)` thread + per block. block_id r = blockIdx.x = one write req; chains are intrinsically serial so a single thread + per block is optimal (warp-level parallelism would idle 31 lanes). + - Per block, early-exit on r >= plan.write_num_valid_reqs[0]. Else load entry_start = plan.write_offsets[r], + entry_count = plan.write_offsets[r+1] - entry_start, seed_slot_idx = plan.write_seed_slot_indices[r] into + registers. + - Initialize running_prev_hash: if seed_slot_idx >= 0, load (token, position, prev_hash) from + canary_buf[seed_slot_idx] and set running_prev_hash = splitmix64_mix3(prev_hash, token, position); + else running_prev_hash = splitmix64(kCanaryChainAnchor). + - Serial chain loop `for j in range(entry_count)`: + i = entry_start + j; + 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] + 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 + store (token, position, running_prev_hash, real_kv_hash) to canary_buf[slot] as 4 int64 fields; + running_prev_hash = splitmix64_mix3(running_prev_hash, token, position); + - All chain state lives in the block's single thread's registers. No shared memory, no cross-block + coordination. + - record_violation() identical to verify (atomicAdd + atomic-write). + - Counters: thread of block 0 does atomicAdd(kernel_run_counter, 1); each block accumulates its + entry_count and atomicAdds to slot_run_counter once at exit. + + Calling contract: + - Pure side-effect; never raises. + - Input-verification mismatch records violations but does NOT abort the chain. + - kernel_run_counter is bumped every call. + - Safe in cuda-graph capture; caller refills input_ids / positions / out_cache_loc / plan + in-place before replay. + + Pinned by torch reference + :func:`sglang.jit_kernel.kv_canary.write_ref.launch_canary_write_kernel_torch_reference`; CUDA must match + byte-for-byte. + """ + canary_buf = context.canary_buf + + _assert_contiguous(canary_buf, "canary_buf") + _assert_contiguous(plan.write_offsets, "plan.write_offsets") + _assert_contiguous(plan.write_seed_slot_indices, "plan.write_seed_slot_indices") + _assert_contiguous(plan.write_num_valid_reqs, "plan.write_num_valid_reqs") + _assert_contiguous(input_ids, "input_ids") + _assert_contiguous(positions, "positions") + _assert_contiguous(out_cache_loc, "out_cache_loc") + if enable_write_input_assert: + if expected_input_tokens is None or expected_input_positions is None: + raise ValueError( + "kv-canary: expected input tensors are required when enable_write_input_assert=True" + ) + _assert_contiguous(expected_input_tokens, "expected_input_tokens") + _assert_contiguous(expected_input_positions, "expected_input_positions") + else: + if expected_input_tokens is not None or expected_input_positions is not None: + raise ValueError( + "kv-canary: expected input tensors must be None when enable_write_input_assert=False" + ) + _assert_contiguous(context.violation_ring, "violation_ring") + _assert_contiguous(context.violation_write_index, "violation_write_index") + _assert_contiguous(context.slot_run_counter, "slot_run_counter") + _assert_contiguous(context.kernel_run_counter, "kernel_run_counter") + _assert_contiguous( + context.enable_chain_position_assert, "enable_chain_position_assert" + ) + + module = _jit_canary_write_module() + module.canary_write_step_cuda( + canary_buf, + plan.write_offsets, + plan.write_seed_slot_indices, + plan.write_num_valid_reqs, + input_ids, + positions, + out_cache_loc, + int(context.kernel_kind), + int(enable_write_input_assert), + expected_input_tokens, + expected_input_positions, + context.violation_ring, + context.violation_write_index, + context.slot_run_counter, + context.kernel_run_counter, + context.enable_chain_position_assert, + ) + + +@cache_once +def _jit_canary_write_module() -> "Module": + return load_jit( + "kv_canary_write", + cuda_files=["kv_canary/canary_write.cuh"], + cuda_wrappers=[ + ("canary_write_step_cuda", "canary::canary_write_step_cuda"), + ], + ) diff --git a/python/sglang/jit_kernel/kv_canary/write_ref.py b/python/sglang/jit_kernel/kv_canary/write_ref.py new file mode 100644 index 000000000..21f9a1442 --- /dev/null +++ b/python/sglang/jit_kernel/kv_canary/write_ref.py @@ -0,0 +1,201 @@ +from __future__ import annotations + +import torch + +from sglang.jit_kernel.kv_canary import consts +from sglang.jit_kernel.kv_canary.verify import ( + VerifyOrWriteContext, +) +from sglang.jit_kernel.kv_canary.verify_ref import ( + _to_signed_int64, + compute_slot_hash, + splitmix64_mix3, +) +from sglang.jit_kernel.kv_canary.write import WritePlan + + +def launch_canary_write_kernel_torch_reference( + *, + context: VerifyOrWriteContext, + plan: WritePlan, + input_ids: torch.Tensor, + positions: torch.Tensor, + out_cache_loc: torch.Tensor, + enable_write_input_assert: bool, + expected_input_tokens: torch.Tensor | None, + expected_input_positions: torch.Tensor | None, +) -> None: + canary_buf = context.canary_buf + kernel_kind = context.kernel_kind + violation_ring = context.violation_ring + violation_write_index = context.violation_write_index + slot_run_counter = context.slot_run_counter + kernel_run_counter = context.kernel_run_counter + enable_chain_position_assert_value = int( + context.enable_chain_position_assert.detach().to("cpu").item() + ) + + work_device = torch.device("cpu") + + kernel_run_counter.add_(1) + + num_valid_reqs = int(plan.write_num_valid_reqs.detach().to("cpu").item()) + req_capacity = int(plan.write_seed_slot_indices.shape[0]) + active_reqs = max(0, min(num_valid_reqs, req_capacity)) + if active_reqs <= 0: + return + + write_offsets_host = plan.write_offsets.detach().to( + device=work_device, dtype=torch.int64 + ) + seed_slot_indices_host = plan.write_seed_slot_indices[:active_reqs].to( + device=work_device, dtype=torch.int64 + ) + input_ids_host = input_ids.detach().to(device=work_device, dtype=torch.int64) + positions_host = positions.detach().to(device=work_device, dtype=torch.int64) + out_cache_loc_host = out_cache_loc.detach().to( + device=work_device, dtype=torch.int64 + ) + + total_entries = int(write_offsets_host[active_reqs].item()) + if total_entries <= 0: + return + + buf_i64 = ( + canary_buf.detach() + .to(device=work_device) + .contiguous() + .view(torch.int64) + .clone() + ) + slot_stride_i64 = int(buf_i64.shape[1]) + if slot_stride_i64 < 4: + raise ValueError( + f"kv-canary: canary_buf slot stride must hold at least 4 int64 fields, got {slot_stride_i64}" + ) + + if enable_write_input_assert: + if expected_input_tokens is None or expected_input_positions is None: + raise ValueError( + "kv-canary: expected input tensors are required when enable_write_input_assert=True" + ) + expected_input_tokens_host = expected_input_tokens.detach().to( + device=work_device, dtype=torch.int64 + ) + expected_input_positions_host = expected_input_positions.detach().to( + device=work_device, dtype=torch.int64 + ) + else: + if expected_input_tokens is not None or expected_input_positions is not None: + raise ValueError( + "kv-canary: expected input tensors must be None when enable_write_input_assert=False" + ) + expected_input_tokens_host = None + expected_input_positions_host = None + + violation_rows: list[list[int]] = [] + total_slots_written = 0 + + for r in range(active_reqs): + entry_start = int(write_offsets_host[r].item()) + entry_end = int(write_offsets_host[r + 1].item()) + entry_count = entry_end - entry_start + if entry_count <= 0: + continue + + seed_slot = int(seed_slot_indices_host[r].item()) + running_prev_hash = compute_slot_hash(buf_i64, seed_slot) + + do_chain_position_assert = (seed_slot >= 0) and ( + enable_chain_position_assert_value != 0 + ) + if do_chain_position_assert: + running_prev_position = int( + buf_i64[seed_slot, consts.CANARY_FIELD_POSITION].item() + ) + else: + running_prev_position = 0 + + for entry_offset in range(entry_count): + entry_idx = entry_start + entry_offset + slot = int(out_cache_loc_host[entry_idx].item()) + if slot < 0: + continue + token = int(input_ids_host[entry_idx].item()) + position = int(positions_host[entry_idx].item()) + + if enable_write_input_assert: + assert expected_input_tokens_host is not None + assert expected_input_positions_host is not None + mismatch_bits = consts.FailReason(0) + expected_token = int(expected_input_tokens_host[entry_idx].item()) + expected_position = int(expected_input_positions_host[entry_idx].item()) + if token != expected_token: + mismatch_bits |= consts.FailReason.WRITE_TOKEN_MISMATCH + if position != expected_position: + mismatch_bits |= consts.FailReason.WRITE_POSITION_MISMATCH + if mismatch_bits != consts.FailReason(0): + row = [0] * consts.VIOLATION_FIELDS + row[consts.VIOLATION_FIELD_KERNEL_KIND] = int(kernel_kind) + row[consts.VIOLATION_FIELD_SLOT_IDX] = slot + row[consts.VIOLATION_FIELD_POSITION] = position + row[consts.VIOLATION_FIELD_STORED_TOKEN] = token + row[consts.VIOLATION_FIELD_EXPECTED_TOKEN] = expected_token + row[consts.VIOLATION_FIELD_STORED_CHAIN_HASH] = _to_signed_int64( + running_prev_hash + ) + row[consts.VIOLATION_FIELD_EXPECTED_AUX] = expected_position + row[consts.VIOLATION_FIELD_FAIL_REASON_BITS] = int(mismatch_bits) + violation_rows.append(row) + + if do_chain_position_assert: + expected_position_chain = running_prev_position + 1 + if position != expected_position_chain: + row = [0] * consts.VIOLATION_FIELDS + row[consts.VIOLATION_FIELD_KERNEL_KIND] = int(kernel_kind) + row[consts.VIOLATION_FIELD_SLOT_IDX] = slot + row[consts.VIOLATION_FIELD_POSITION] = position + row[consts.VIOLATION_FIELD_STORED_TOKEN] = token + row[consts.VIOLATION_FIELD_EXPECTED_TOKEN] = token + row[consts.VIOLATION_FIELD_STORED_CHAIN_HASH] = _to_signed_int64( + running_prev_hash + ) + row[consts.VIOLATION_FIELD_EXPECTED_AUX] = expected_position_chain + row[consts.VIOLATION_FIELD_FAIL_REASON_BITS] = int( + consts.FailReason.WRITE_POSITION_MISMATCH + ) + violation_rows.append(row) + running_prev_position = position + + buf_i64[slot, consts.CANARY_FIELD_TOKEN] = token + buf_i64[slot, consts.CANARY_FIELD_POSITION] = position + buf_i64[slot, consts.CANARY_FIELD_PREV_HASH] = _to_signed_int64( + running_prev_hash + ) + buf_i64[slot, consts.CANARY_FIELD_REAL_KV_HASH] = 0 + + running_prev_hash = splitmix64_mix3(running_prev_hash, token, position) + + total_slots_written += 1 + + canary_buf.view(torch.int64).copy_( + buf_i64.to(canary_buf.device).view(canary_buf.shape[0], slot_stride_i64) + ) + + slot_run_counter.add_(total_slots_written) + + if len(violation_rows) == 0: + return + + base_idx = int(violation_write_index.detach().to("cpu").item()) + ring_capacity = int(violation_ring.shape[0]) + new_rows = torch.tensor(violation_rows, dtype=torch.int64, device=work_device) + write_count_in_ring = max(0, min(len(violation_rows), ring_capacity - base_idx)) + if write_count_in_ring > 0: + ring_host = violation_ring.detach().to(device=work_device) + ring_host[base_idx : base_idx + write_count_in_ring, :] = new_rows[ + :write_count_in_ring, : + ] + violation_ring.copy_(ring_host.to(violation_ring.device)) + + violation_write_index[0] = violation_write_index[0] + len(violation_rows) diff --git a/python/sglang/jit_kernel/tests/kv_canary/_canary_helpers.py b/python/sglang/jit_kernel/tests/kv_canary/_canary_helpers.py new file mode 100644 index 000000000..42366f55c --- /dev/null +++ b/python/sglang/jit_kernel/tests/kv_canary/_canary_helpers.py @@ -0,0 +1,330 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Optional + +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 VerifyPlan +from sglang.jit_kernel.kv_canary.write import WritePlan +from sglang.jit_kernel.tests.kv_canary._constants import ( + _I64_SIGN_BIT, + _U64_MASK, + DEFAULT_NUM_SLOTS, + DEFAULT_RING_CAPACITY, + DEFAULT_SLOT_STRIDE_BYTES, +) + +__all__ = [ + "FakeViolationLog", + "assert_canary_buf_equal", + "assert_canary_state_equal", + "assert_only_bits_set", + "chain_anchor_signed", + "make_canary_buf", + "make_canary_buf_pair", + "make_log_pair", + "make_verify_plan", + "make_verify_plan_pair", + "make_write_plan", + "make_write_plan_pair", + "read_slot_fields", + "stamp_clean_chain", + "stamp_pair", + "to_signed_int64", + "write_slot_fields", +] + + +@dataclass(frozen=True, slots=True, kw_only=True) +class FakeViolationLog: + ring: torch.Tensor + write_index: torch.Tensor + slot_run_counter: torch.Tensor + kernel_run_counter: torch.Tensor + enable_chain_position_assert: torch.Tensor + + @classmethod + def allocate( + cls, *, capacity: int = DEFAULT_RING_CAPACITY, device: torch.device + ) -> "FakeViolationLog": + return cls( + ring=torch.zeros( + capacity, consts.VIOLATION_FIELDS, dtype=torch.int64, device=device + ), + write_index=torch.zeros(1, dtype=torch.int32, device=device), + slot_run_counter=torch.zeros(1, dtype=torch.int64, device=device), + kernel_run_counter=torch.zeros(1, dtype=torch.int64, device=device), + enable_chain_position_assert=torch.ones( + 1, dtype=torch.int32, device=device + ), + ) + + +def make_canary_buf( + *, + num_slots: int = DEFAULT_NUM_SLOTS, + slot_stride_bytes: int = DEFAULT_SLOT_STRIDE_BYTES, + device: torch.device, +) -> torch.Tensor: + return torch.zeros(num_slots, slot_stride_bytes, dtype=torch.uint8, device=device) + + +def make_canary_buf_pair( + *, + num_slots: int = DEFAULT_NUM_SLOTS, + slot_stride_bytes: int = DEFAULT_SLOT_STRIDE_BYTES, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + cuda_buf = make_canary_buf( + num_slots=num_slots, slot_stride_bytes=slot_stride_bytes, device=device + ) + return cuda_buf, cuda_buf.clone() + + +def make_log_pair( + *, + capacity: int = DEFAULT_RING_CAPACITY, + device: torch.device, +) -> tuple[FakeViolationLog, FakeViolationLog]: + return ( + FakeViolationLog.allocate(capacity=capacity, device=device), + FakeViolationLog.allocate(capacity=capacity, device=device), + ) + + +def make_verify_plan( + *, + slot_indices: list[int], + positions: list[int], + prev_slot_indices: list[int], + expected_input_ids: Optional[list[int]] = None, + capacity: Optional[int] = None, + device: torch.device, +) -> VerifyPlan: + """Build a VerifyPlan whose active prefix matches the three input lists. + + Active prefix mirrors the input lists. Tail entries are left at the + allocate-time defaults; ``verify_num_valid = len(slot_indices)``. + + ``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. + """ + n_active = len(slot_indices) + if not (len(positions) == n_active and len(prev_slot_indices) == n_active): + raise ValueError( + "make_verify_plan: slot_indices, positions, and prev_slot_indices must all have the same length" + ) + if expected_input_ids is None: + expected_input_ids = [-1] * n_active + if len(expected_input_ids) != n_active: + raise ValueError( + "make_verify_plan: expected_input_ids must match len(slot_indices)" + ) + cap = capacity if capacity is not None else max(n_active, 1) + plan = VerifyPlan.allocate(verify_capacity=cap, device=device) + if n_active > 0: + plan.verify_slot_indices[:n_active] = torch.tensor( + slot_indices, dtype=torch.int64, device=device + ) + plan.verify_expected_tokens[:n_active] = torch.tensor( + expected_input_ids, dtype=torch.int64, device=device + ) + plan.verify_expected_positions[:n_active] = torch.tensor( + positions, dtype=torch.int64, device=device + ) + plan.verify_prev_slot_indices[:n_active] = torch.tensor( + prev_slot_indices, dtype=torch.int64, device=device + ) + plan.verify_num_valid[0] = n_active + return plan + + +def make_verify_plan_pair( + *, + slot_indices: list[int], + positions: list[int], + prev_slot_indices: list[int], + expected_input_ids: Optional[list[int]] = None, + capacity: Optional[int] = None, + device: torch.device, +) -> tuple[VerifyPlan, VerifyPlan]: + return ( + make_verify_plan( + slot_indices=slot_indices, + positions=positions, + prev_slot_indices=prev_slot_indices, + expected_input_ids=expected_input_ids, + capacity=capacity, + device=device, + ), + make_verify_plan( + slot_indices=slot_indices, + positions=positions, + prev_slot_indices=prev_slot_indices, + expected_input_ids=expected_input_ids, + capacity=capacity, + device=device, + ), + ) + + +def make_write_plan( + *, + write_offsets: list[int], + seed_slot_indices: list[int], + num_valid_reqs: int, + req_capacity: Optional[int] = None, + device: torch.device, +) -> WritePlan: + """Build a WritePlan from raw offsets and seed slot lists. + + ``write_offsets`` must have length ``len(seed_slot_indices) + 1`` (the trailing total entry count). + """ + n_active = len(seed_slot_indices) + if len(write_offsets) != n_active + 1: + raise ValueError( + "make_write_plan: write_offsets must have length len(seed_slot_indices) + 1" + ) + cap = req_capacity if req_capacity is not None else max(n_active, 1) + plan = WritePlan.allocate(write_req_capacity=cap, device=device) + if n_active > 0: + plan.write_seed_slot_indices[:n_active] = torch.tensor( + seed_slot_indices, dtype=torch.int64, device=device + ) + plan.write_offsets[: n_active + 1] = torch.tensor( + write_offsets, dtype=torch.int64, device=device + ) + plan.write_num_valid_reqs[0] = num_valid_reqs + return plan + + +def make_write_plan_pair( + *, + write_offsets: list[int], + seed_slot_indices: list[int], + num_valid_reqs: int, + req_capacity: Optional[int] = None, + device: torch.device, +) -> tuple[WritePlan, WritePlan]: + return ( + make_write_plan( + write_offsets=write_offsets, + seed_slot_indices=seed_slot_indices, + num_valid_reqs=num_valid_reqs, + req_capacity=req_capacity, + device=device, + ), + make_write_plan( + write_offsets=write_offsets, + seed_slot_indices=seed_slot_indices, + num_valid_reqs=num_valid_reqs, + req_capacity=req_capacity, + device=device, + ), + ) + + +def to_signed_int64(value: int) -> int: + value &= _U64_MASK + if value >= _I64_SIGN_BIT: + value -= 1 << 64 + return value + + +def chain_anchor_signed() -> int: + return to_signed_int64(splitmix64(consts.CANARY_CHAIN_ANCHOR)) + + +def write_slot_fields( + *, + canary_buf: torch.Tensor, + slot_idx: int, + token: int, + position: int, + prev_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 + + +def stamp_pair( + buf_pair: tuple[torch.Tensor, torch.Tensor], + *, + slot_idx: int, + token: int, + position: int, + prev_hash: int, +) -> None: + """Stamp the same slot fields into both (cuda, ref) canary buffers.""" + for buf in buf_pair: + write_slot_fields( + canary_buf=buf, + slot_idx=slot_idx, + token=token, + position=position, + prev_hash=prev_hash, + ) + + +def read_slot_fields( + *, canary_buf: torch.Tensor, slot_idx: int +) -> tuple[int, int, int, int]: + row = canary_buf.view(torch.int64)[slot_idx, :4].detach().cpu().tolist() + return int(row[0]), int(row[1]), int(row[2]), int(row[3]) + + +def stamp_clean_chain( + *, + cuda_buf: torch.Tensor, + ref_buf: torch.Tensor, + slot_indices: list[int], + tokens: list[int], + positions: list[int], +) -> list[int]: + running_prev_hash = splitmix64(consts.CANARY_CHAIN_ANCHOR) + stored_prev_hashes: list[int] = [] + for slot_idx, token, position in zip(slot_indices, tokens, positions): + signed_prev = to_signed_int64(running_prev_hash) + for buf in (cuda_buf, ref_buf): + write_slot_fields( + canary_buf=buf, + slot_idx=slot_idx, + token=token, + position=position, + prev_hash=signed_prev, + ) + stored_prev_hashes.append(signed_prev) + running_prev_hash = splitmix64_mix3(running_prev_hash, token, position) + return stored_prev_hashes + + +def assert_canary_state_equal( + *, log_a: FakeViolationLog, log_b: FakeViolationLog +) -> None: + for name in ("ring", "write_index", "slot_run_counter", "kernel_run_counter"): + assert torch.equal( + getattr(log_a, name), getattr(log_b, name) + ), f"{name} diverged (CUDA vs ref)" + + +def assert_canary_buf_equal(*, buf_a: torch.Tensor, buf_b: torch.Tensor) -> None: + assert torch.equal(buf_a, buf_b), "canary_buf diverged (CUDA vs ref)" + + +def assert_only_bits_set(fail_bits: int, expected_bits: int) -> None: + assert ( + fail_bits & expected_bits + ) == expected_bits, ( + f"missing expected bits: expected {expected_bits:#b} got {fail_bits:#b}" + ) + assert ( + fail_bits & ~expected_bits + ) == 0, f"unexpected extra bits: got {fail_bits:#b} extras {fail_bits & ~expected_bits:#b}" diff --git a/python/sglang/jit_kernel/tests/kv_canary/_fixtures.py b/python/sglang/jit_kernel/tests/kv_canary/_fixtures.py new file mode 100644 index 000000000..60d8ac46e --- /dev/null +++ b/python/sglang/jit_kernel/tests/kv_canary/_fixtures.py @@ -0,0 +1,150 @@ +from __future__ import annotations + +import random +from typing import Literal, Optional + +import torch + +from sglang.jit_kernel.kv_canary.verify import VerifyPlan +from sglang.jit_kernel.kv_canary.write import WritePlan + +_DEVICE = torch.device("cuda") + + +LutKind = Literal["identity", "shift", "permutation", "with_oob"] + + +def make_lut( + *, + kind: LutKind, + pool_size: int, + device: torch.device, + rng: Optional[random.Random] = None, +) -> torch.Tensor: + base = torch.arange(pool_size + 1, dtype=torch.int64, device=device) + if kind == "identity": + return base.contiguous() + if kind == "shift": + return (base + 100).contiguous() + if kind in ("permutation", "with_oob"): + if rng is None: + rng = random.Random(0) + perm = list(range(pool_size + 1)) + rng.shuffle(perm) + out = torch.tensor(perm, dtype=torch.int64, device=device) + if kind == "with_oob": + out[-1] = pool_size + 999 + return out.contiguous() + raise ValueError(f"unknown LutKind: {kind}") + + +ReqToTokenKind = Literal["linear", "sparse_permuted"] + + +def make_req_to_token( + *, + kind: ReqToTokenKind, + max_reqs: int, + max_seq_len: int, + device: torch.device, + rng: Optional[random.Random] = None, +) -> torch.Tensor: + if kind == "linear": + rp_axis = torch.arange(max_reqs, device=device, dtype=torch.int32).unsqueeze(1) + pos_axis = torch.arange( + max_seq_len, device=device, dtype=torch.int32 + ).unsqueeze(0) + return (rp_axis * max_seq_len + pos_axis).contiguous() + if rng is None: + rng = random.Random(0) + pool_size = max_reqs * max_seq_len + # Slots index into a full_to_swa LUT sized [pool_size + 1], so values must stay + # in [0, pool_size]. The universe spans [1, pool_size] (skipping 0 as reserved), + # giving exactly max_reqs * max_seq_len unique slots — one per (rp, pos) cell. + slot_universe = list(range(1, pool_size + 1)) + rng.shuffle(slot_universe) + rtt = torch.zeros((max_reqs, max_seq_len), dtype=torch.int32, device=device) + cursor = 0 + for rp in range(max_reqs): + per_req = slot_universe[cursor : cursor + max_seq_len] + cursor += max_seq_len + rtt[rp, :] = torch.tensor(per_req, dtype=torch.int32, device=device) + return rtt.contiguous() + + +PaddingKind = Literal["none", "trailing", "interleaved"] + + +def make_padding_mask( + *, + bs: int, + kind: PaddingKind, + rng: Optional[random.Random] = None, + padding_fraction: float = 0.25, +) -> list[bool]: + if bs == 0: + return [] + if kind == "none": + return [False] * bs + n_pad = max(1, int(bs * padding_fraction)) if bs > 0 else 0 + n_pad = min(n_pad, bs) + if kind == "trailing": + return [False] * (bs - n_pad) + [True] * n_pad + if kind == "interleaved": + if rng is None: + rng = random.Random(0) + mask = [False] * bs + chosen = rng.sample(range(bs), k=n_pad) + for idx in chosen: + mask[idx] = True + return mask + raise ValueError(f"unknown PaddingKind: {kind}") + + +CapacityKind = Literal["loose", "tight_match", "under_by_one"] + + +def derive_plan_capacity( + *, + kind: CapacityKind, + total_verify: int, + extras_count: int, + bs: int, +) -> tuple[int, int]: + needed = total_verify + extras_count + if kind == "loose": + return max(needed + 64, 128), max(bs + 4, 8) + if kind == "tight_match": + return max(needed, 1), max(bs + 4, 8) + if kind == "under_by_one": + return max(needed - 1, 1), max(bs + 4, 8) + raise ValueError(f"unknown CapacityKind: {kind}") + + +def allocate_plan_pair( + *, + verify_capacity: int, + write_req_capacity: int, +) -> tuple[VerifyPlan, WritePlan, VerifyPlan, WritePlan]: + return ( + VerifyPlan.allocate(verify_capacity=verify_capacity, device=_DEVICE), + WritePlan.allocate(write_req_capacity=write_req_capacity, device=_DEVICE), + VerifyPlan.allocate(verify_capacity=verify_capacity, device=_DEVICE), + WritePlan.allocate(write_req_capacity=write_req_capacity, device=_DEVICE), + ) + + +def empty_extras() -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + return ( + torch.zeros(1, dtype=torch.int64, device=_DEVICE), + torch.zeros(1, dtype=torch.int64, device=_DEVICE), + torch.zeros(1, dtype=torch.int64, device=_DEVICE), + torch.zeros(1, dtype=torch.int32, device=_DEVICE), + ) + + +def dummy_pseudo_tensors(num_tokens: int) -> tuple[torch.Tensor, torch.Tensor]: + return ( + torch.zeros(num_tokens, dtype=torch.int64, device=_DEVICE), + torch.zeros(num_tokens, dtype=torch.int64, device=_DEVICE), + ) diff --git a/python/sglang/jit_kernel/tests/kv_canary/_invariants.py b/python/sglang/jit_kernel/tests/kv_canary/_invariants.py new file mode 100644 index 000000000..55a7a6387 --- /dev/null +++ b/python/sglang/jit_kernel/tests/kv_canary/_invariants.py @@ -0,0 +1,520 @@ +"""Ref/real-independent invariant assertions for kv_canary kernel tests. + +Each invariant only looks at the kernel's inputs and outputs (shape relationships, monotonicity, tail +positions, etc.) — it must never re-implement the reference algorithm. Hand and fuzz tests both call +into this module so a single contract violation surfaces consistently. +""" + +from __future__ import annotations + +from typing import Optional + +import torch + +from sglang.jit_kernel.kv_canary import consts +from sglang.jit_kernel.kv_canary.verify import CanaryLaunchTag, VerifyPlan +from sglang.jit_kernel.kv_canary.write import WritePlan +from sglang.jit_kernel.tests.kv_canary._canary_helpers import FakeViolationLog + + +class PlanInvariants: + @staticmethod + def assert_all( + *, + verify_plan: VerifyPlan, + write_plan: WritePlan, + req_pool_indices: torch.Tensor, + prefix_lens: torch.Tensor, + extend_seq_lens: torch.Tensor, + swa_window_size: int, + extras_slot_indices: torch.Tensor, + extras_positions: torch.Tensor, + extras_prev_slot_indices: torch.Tensor, + extras_count: int, + ) -> None: + PlanInvariants._assert_write_offsets_monotone(write_plan) + PlanInvariants._assert_write_offsets_total_matches_active_extend_sum( + write_plan=write_plan, + extend_seq_lens=extend_seq_lens, + req_pool_indices=req_pool_indices, + ) + derived = PlanInvariants._assert_verify_num_valid_equals_derived_plus_extras( + verify_plan=verify_plan, + prefix_lens=prefix_lens, + req_pool_indices=req_pool_indices, + swa_window_size=swa_window_size, + extras_count=extras_count, + ) + PlanInvariants._assert_padding_row_seed_is_minus_one( + write_plan=write_plan, + req_pool_indices=req_pool_indices, + ) + # In overflow (derived + extras > verify_capacity) the plan kernel disables + # verify (enable=0) and the verify entries buffer is partially populated; the + # downstream entry-shape invariants only hold when the kernel actually emitted + # the full set, so guard them on enable=1. + verify_enabled = int(verify_plan.enable[0].item()) == 1 + if verify_enabled: + PlanInvariants._assert_extras_land_at_tail( + verify_plan=verify_plan, + derived_verify_count=derived, + extras_slot_indices=extras_slot_indices, + extras_positions=extras_positions, + extras_prev_slot_indices=extras_prev_slot_indices, + extras_count=extras_count, + ) + PlanInvariants._assert_prev_slot_minus_one_iff_chain_head( + verify_plan=verify_plan, + swa_window_size=swa_window_size, + derived_verify_count=derived, + ) + + @staticmethod + def _assert_write_offsets_monotone(write_plan: WritePlan) -> None: + n_active = int(write_plan.write_num_valid_reqs[0].item()) + if n_active < 0: + raise AssertionError(f"write_num_valid_reqs negative: {n_active}") + offsets = write_plan.write_offsets[: n_active + 1].detach().cpu().tolist() + for i in range(len(offsets) - 1): + assert ( + offsets[i] <= offsets[i + 1] + ), f"write_offsets non-monotone at {i}: {offsets[i]} > {offsets[i + 1]}" + + @staticmethod + def _assert_write_offsets_total_matches_active_extend_sum( + *, + write_plan: WritePlan, + extend_seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + ) -> None: + n_active = int(write_plan.write_num_valid_reqs[0].item()) + total = int(write_plan.write_offsets[n_active].item()) + rpi_cpu = req_pool_indices.detach().cpu().tolist() + ext_cpu = extend_seq_lens.detach().cpu().tolist() + expected_total = sum(ext for rpi, ext in zip(rpi_cpu, ext_cpu) if rpi != 0) + assert ( + total == expected_total + ), f"write_offsets total {total} != active extend sum {expected_total}" + + @staticmethod + def _assert_extras_land_at_tail( + *, + verify_plan: VerifyPlan, + derived_verify_count: int, + extras_slot_indices: torch.Tensor, + extras_positions: torch.Tensor, + extras_prev_slot_indices: torch.Tensor, + extras_count: int, + ) -> None: + if extras_count == 0: + return + tail_start = derived_verify_count + tail_end = derived_verify_count + extras_count + n_valid = int(verify_plan.verify_num_valid[0].item()) + assert ( + tail_end <= n_valid + ), f"extras tail {tail_end} exceeds verify_num_valid {n_valid}" + plan_slots = verify_plan.verify_slot_indices[tail_start:tail_end] + plan_positions = verify_plan.verify_expected_positions[tail_start:tail_end] + plan_prevs = verify_plan.verify_prev_slot_indices[tail_start:tail_end] + assert torch.equal(plan_slots, extras_slot_indices[:extras_count]) + assert torch.equal(plan_positions, extras_positions[:extras_count]) + assert torch.equal(plan_prevs, extras_prev_slot_indices[:extras_count]) + + @staticmethod + def _assert_padding_row_seed_is_minus_one( + *, + write_plan: WritePlan, + req_pool_indices: torch.Tensor, + ) -> None: + n_active = int(write_plan.write_num_valid_reqs[0].item()) + if n_active == 0: + return + rpi_cpu = req_pool_indices.detach().cpu().tolist() + seeds_cpu = ( + write_plan.write_seed_slot_indices[:n_active].detach().cpu().tolist() + ) + for r in range(min(n_active, len(rpi_cpu))): + if rpi_cpu[r] == 0: + assert ( + seeds_cpu[r] == -1 + ), f"padding row {r} has seed {seeds_cpu[r]} != -1" + + @staticmethod + def _assert_prev_slot_minus_one_iff_chain_head( + *, + verify_plan: VerifyPlan, + swa_window_size: int, + derived_verify_count: int, + ) -> None: + if derived_verify_count == 0: + return + positions_cpu = ( + verify_plan.verify_expected_positions[:derived_verify_count] + .detach() + .cpu() + .tolist() + ) + prevs_cpu = ( + verify_plan.verify_prev_slot_indices[:derived_verify_count] + .detach() + .cpu() + .tolist() + ) + for i, (pos, prev) in enumerate(zip(positions_cpu, prevs_cpu)): + if pos == 0: + assert ( + prev == -1 + ), f"entry {i} at position 0 must have prev=-1, got {prev}" + else: + if swa_window_size == 0: + assert ( + prev != -1 + ), f"FULL entry {i} at position {pos} must have prev != -1, got {prev}" + + @staticmethod + def _assert_verify_num_valid_equals_derived_plus_extras( + *, + verify_plan: VerifyPlan, + prefix_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + swa_window_size: int, + extras_count: int, + ) -> int: + rpi_cpu = req_pool_indices.detach().cpu().tolist() + pfx_cpu = prefix_lens.detach().cpu().tolist() + derived = 0 + for rpi, pfx in zip(rpi_cpu, pfx_cpu): + if rpi == 0: + continue + if swa_window_size > 0: + window_start = max(0, pfx - swa_window_size) + derived += max(0, pfx - window_start) + else: + derived += max(0, pfx) + # The plan kernel clamps verify_num_valid to verify_capacity and turns enable + # off when (derived + extras) overflows the slot indices buffer. The invariant + # must match that: on overflow the kernel records the capacity, on no overflow + # it records the exact derived total (so the verify kernel scans every row). + verify_capacity = int(verify_plan.verify_slot_indices.shape[0]) + expected_unclamped = derived + extras_count + expected = min(expected_unclamped, verify_capacity) + overflow = expected_unclamped > verify_capacity + actual = int(verify_plan.verify_num_valid[0].item()) + assert actual == expected, ( + f"verify_num_valid {actual} != min(derived {derived} + extras {extras_count}, " + f"verify_capacity {verify_capacity}) = {expected}" + ) + enable = int(verify_plan.enable[0].item()) + expected_enable = 0 if overflow else 1 + assert enable == expected_enable, ( + f"verify_plan.enable {enable} != expected {expected_enable} " + f"(overflow={overflow}; derived+extras={expected_unclamped}, " + f"verify_capacity={verify_capacity})" + ) + return derived + + +class VerifyInvariants: + @staticmethod + def assert_all( + *, + canary_buf_before: torch.Tensor, + canary_buf_after: torch.Tensor, + log_before: FakeViolationLog, + log_after: FakeViolationLog, + plan: VerifyPlan, + kernel_kind: CanaryLaunchTag, + ) -> None: + VerifyInvariants._assert_canary_buf_unchanged( + canary_buf_before=canary_buf_before, canary_buf_after=canary_buf_after + ) + VerifyInvariants._assert_violation_count_le_active_entries( + log_after=log_after, log_before=log_before, plan=plan + ) + VerifyInvariants._assert_violation_rows_have_valid_slot_and_kernel_kind( + log_after=log_after, + log_before=log_before, + plan=plan, + kernel_kind=kernel_kind, + ) + VerifyInvariants._assert_slot_run_counter_incremented_by_active_entries( + log_before=log_before, log_after=log_after, plan=plan + ) + VerifyInvariants._assert_kernel_run_counter_incremented_by_one( + log_before=log_before, log_after=log_after + ) + + @staticmethod + def _assert_canary_buf_unchanged( + *, + canary_buf_before: torch.Tensor, + canary_buf_after: torch.Tensor, + ) -> None: + assert torch.equal( + canary_buf_before, canary_buf_after + ), "verify kernel mutated canary_buf (must be read-only)" + + @staticmethod + def _assert_violation_count_le_active_entries( + *, + log_after: FakeViolationLog, + log_before: FakeViolationLog, + plan: VerifyPlan, + ) -> None: + delta = int(log_after.write_index[0].item()) - int( + log_before.write_index[0].item() + ) + n_active = int(plan.verify_num_valid[0].item()) + assert ( + 0 <= delta <= n_active + ), f"violation_write_index delta {delta} out of [0, {n_active}]" + + @staticmethod + def _assert_violation_rows_have_valid_slot_and_kernel_kind( + *, + log_after: FakeViolationLog, + log_before: FakeViolationLog, + plan: VerifyPlan, + kernel_kind: CanaryLaunchTag, + ) -> None: + write_idx_after = int(log_after.write_index[0].item()) + write_idx_before = int(log_before.write_index[0].item()) + if write_idx_after == write_idx_before: + return + ring_capacity = log_after.ring.shape[0] + visible_start = write_idx_before + visible_end = min(write_idx_after, ring_capacity) + if visible_end <= visible_start: + return + n_active = int(plan.verify_num_valid[0].item()) + plan_slots = set(plan.verify_slot_indices[:n_active].detach().cpu().tolist()) + rows = log_after.ring[visible_start:visible_end].detach().cpu() + for i in range(rows.shape[0]): + kind = int(rows[i, consts.VIOLATION_FIELD_KERNEL_KIND].item()) + assert kind == int( + kernel_kind + ), f"row {visible_start + i} kernel_kind {kind} != expected {int(kernel_kind)}" + slot = int(rows[i, consts.VIOLATION_FIELD_SLOT_IDX].item()) + assert ( + slot in plan_slots + ), f"row {visible_start + i} slot {slot} not in plan_slots" + + @staticmethod + def _assert_slot_run_counter_incremented_by_active_entries( + *, + log_before: FakeViolationLog, + log_after: FakeViolationLog, + plan: VerifyPlan, + ) -> None: + n_active = int(plan.verify_num_valid[0].item()) + delta = int(log_after.slot_run_counter[0].item()) - int( + log_before.slot_run_counter[0].item() + ) + assert ( + delta == n_active + ), f"slot_run_counter delta {delta} != active entries {n_active}" + + @staticmethod + def _assert_kernel_run_counter_incremented_by_one( + *, + log_before: FakeViolationLog, + log_after: FakeViolationLog, + ) -> None: + delta = int(log_after.kernel_run_counter[0].item()) - int( + log_before.kernel_run_counter[0].item() + ) + assert delta == 1, f"kernel_run_counter delta {delta} != 1" + + +class WriteInvariants: + @staticmethod + def assert_all( + *, + canary_buf_before: torch.Tensor, + canary_buf_after: torch.Tensor, + plan: WritePlan, + input_ids: torch.Tensor, + positions: torch.Tensor, + out_cache_loc: torch.Tensor, + enable_write_verify_inputs: bool, + expected_input_tokens: Optional[torch.Tensor], + expected_input_positions: Optional[torch.Tensor], + log_before: FakeViolationLog, + log_after: FakeViolationLog, + ) -> None: + WriteInvariants._assert_written_slots_token_position_match_input( + canary_buf_after=canary_buf_after, + plan=plan, + input_ids=input_ids, + positions=positions, + out_cache_loc=out_cache_loc, + ) + WriteInvariants._assert_slot_minus_one_skipped( + canary_buf_before=canary_buf_before, + canary_buf_after=canary_buf_after, + plan=plan, + out_cache_loc=out_cache_loc, + ) + WriteInvariants._assert_pseudo_violation_only_on_mismatch( + enable_write_verify_inputs=enable_write_verify_inputs, + log_before=log_before, + log_after=log_after, + expected_input_tokens=expected_input_tokens, + expected_input_positions=expected_input_positions, + input_ids=input_ids, + positions=positions, + out_cache_loc=out_cache_loc, + plan=plan, + ) + WriteInvariants._assert_write_slot_run_counter_incremented( + log_before=log_before, + log_after=log_after, + plan=plan, + out_cache_loc=out_cache_loc, + ) + WriteInvariants._assert_write_kernel_run_counter_incremented_by_one( + log_before=log_before, log_after=log_after + ) + + @staticmethod + def _assert_written_slots_token_position_match_input( + *, + canary_buf_after: torch.Tensor, + plan: WritePlan, + input_ids: torch.Tensor, + positions: torch.Tensor, + out_cache_loc: torch.Tensor, + ) -> None: + n_active = int(plan.write_num_valid_reqs[0].item()) + if n_active == 0: + return + offsets = plan.write_offsets[: n_active + 1].detach().cpu().tolist() + total = offsets[n_active] + slots_cpu = out_cache_loc[:total].detach().cpu().tolist() + tokens_cpu = input_ids[:total].detach().cpu().tolist() + pos_cpu = positions[:total].detach().cpu().tolist() + view = canary_buf_after.view(torch.int64) + for i in range(total): + slot = slots_cpu[i] + if slot < 0: + continue + stored_token = int(view[slot, 0].item()) + stored_position = int(view[slot, 1].item()) + assert ( + stored_token == tokens_cpu[i] + ), f"slot {slot}: stored token {stored_token} != input {tokens_cpu[i]}" + assert ( + stored_position == pos_cpu[i] + ), f"slot {slot}: stored position {stored_position} != input {pos_cpu[i]}" + + @staticmethod + def _assert_slot_minus_one_skipped( + *, + canary_buf_before: torch.Tensor, + canary_buf_after: torch.Tensor, + plan: WritePlan, + out_cache_loc: torch.Tensor, + ) -> None: + n_active = int(plan.write_num_valid_reqs[0].item()) + if n_active == 0: + return + total = int(plan.write_offsets[n_active].item()) + slots_cpu = out_cache_loc[:total].detach().cpu().tolist() + written_slots = {s for s in slots_cpu if s >= 0} + view_before = canary_buf_before.view(torch.int64) + view_after = canary_buf_after.view(torch.int64) + num_slots = canary_buf_after.shape[0] + for slot in range(num_slots): + if slot in written_slots: + continue + assert torch.equal( + view_before[slot], view_after[slot] + ), f"slot {slot} not in out_cache_loc but canary_buf changed" + + @staticmethod + def _assert_pseudo_violation_only_on_mismatch( + *, + enable_write_verify_inputs: bool, + log_before: FakeViolationLog, + log_after: FakeViolationLog, + expected_input_tokens: Optional[torch.Tensor], + expected_input_positions: Optional[torch.Tensor], + input_ids: torch.Tensor, + positions: torch.Tensor, + out_cache_loc: torch.Tensor, + plan: WritePlan, + ) -> None: + delta = int(log_after.write_index[0].item()) - int( + log_before.write_index[0].item() + ) + if not enable_write_verify_inputs: + assert ( + delta == 0 + ), f"enable_write_verify_inputs=OFF must produce no violations, got {delta}" + return + if expected_input_tokens is None or expected_input_positions is None: + return + n_active = int(plan.write_num_valid_reqs[0].item()) + if n_active == 0: + assert delta == 0, f"empty plan produced {delta} violations" + return + total = int(plan.write_offsets[n_active].item()) + tok = input_ids[:total].detach().cpu().tolist() + pos = positions[:total].detach().cpu().tolist() + exp_tok = expected_input_tokens[:total].detach().cpu().tolist() + exp_pos = expected_input_positions[:total].detach().cpu().tolist() + slots_cpu = out_cache_loc[:total].detach().cpu().tolist() + mismatch_entries = sum( + 1 + for i in range(total) + if slots_cpu[i] >= 0 and (tok[i] != exp_tok[i] or pos[i] != exp_pos[i]) + ) + no_mismatch = mismatch_entries == 0 + if no_mismatch: + assert ( + delta == 0 + ), f"enable_write_verify_inputs=ON with no mismatch produced {delta} violations" + else: + assert ( + delta == mismatch_entries + ), f"write input mismatch count {mismatch_entries} produced {delta} violations" + + @staticmethod + def _assert_write_slot_run_counter_incremented( + *, + log_before: FakeViolationLog, + log_after: FakeViolationLog, + plan: WritePlan, + out_cache_loc: torch.Tensor, + ) -> None: + n_active = int(plan.write_num_valid_reqs[0].item()) + if n_active == 0: + delta = int(log_after.slot_run_counter[0].item()) - int( + log_before.slot_run_counter[0].item() + ) + assert delta == 0, f"empty plan incremented slot_run_counter by {delta}" + return + total = int(plan.write_offsets[n_active].item()) + # The write kernel skips entries where out_cache_loc < 0 (the documented "mark + # skip" path used by SWA-translated callers), so the slot_run_counter delta + # tracks the count of writeable entries, not the planned total. + writeable = int((out_cache_loc[:total] >= 0).sum().item()) + delta = int(log_after.slot_run_counter[0].item()) - int( + log_before.slot_run_counter[0].item() + ) + assert delta == writeable, ( + f"slot_run_counter delta {delta} != writeable entries {writeable} " + f"(total={total}, skipped={total - writeable})" + ) + + @staticmethod + def _assert_write_kernel_run_counter_incremented_by_one( + *, + log_before: FakeViolationLog, + log_after: FakeViolationLog, + ) -> None: + delta = int(log_after.kernel_run_counter[0].item()) - int( + log_before.kernel_run_counter[0].item() + ) + assert delta == 1, f"kernel_run_counter delta {delta} != 1"