From 962912c9d906f28bbe7134cdb06a96f8ed895553 Mon Sep 17 00:00:00 2001 From: NInfer Agent Date: Sun, 30 Aug 2026 20:58:49 +0800 Subject: [PATCH 01/11] feat(kv): entropy-coded cold pool for INT8-tier pages (raw nibble slots) Fixed raw slots (9232 B: header + E2M1 nibbles + E4M3 g16 scales) hold requantized cold pages for both the INT8 and NVFP4 tiers. Requantizing INT8 planes to g64 E2M1 measures NMSE 0.012-0.014 (inside the accepted NVFP4-layer envelope) at 1.85-1.99x per head-page, ~1.66x aggregate cold KV on the 27B production table. The pack/restore kernels, the per-layer dtype dispatch, the decode and prefill cold staging (inline nibble->int8 adapter preserving the int8 QK tensor cores), and the length-based slot sizing are all included; --cold-policy window|host plus --cold-keep-tokens/--cold-host-bytes control activation. Three latent v1 cold-addressing bugs (compress_page slot scaling, decode and prefill flat slot indices) are fixed on the way. --- apps/cli/main.cpp | 3 + apps/cli/options.cpp | 13 + apps/cli/options.h | 3 + include/ninfer/ops/cold_i8.h | 30 ++ include/ninfer/ops/entropy_cold_requant.h | 32 ++ include/ninfer/types.h | 13 + src/ops/kernel/cold_i8_kernels.cuh | 105 +++++ .../kernel/entropy_cold_requant_kernels.cuh | 121 +++++ src/ops/launcher/cold_i8.cu | 85 ++++ src/ops/launcher/cold_i8.h | 17 + src/ops/launcher/entropy_cold_requant.cu | 20 + src/ops/launcher/entropy_cold_requant.h | 20 + src/ops/wrapper/cold_i8.cpp | 25 + src/ops/wrapper/entropy_cold_requant.cpp | 23 + .../qwen3_6/impl/runtime/kv_calibration.h | 116 +++++ tests/ops/test_cold_i8.cpp | 274 +++++++++++ tests/ops/test_entropy_cold_requant.cpp | 432 ++++++++++++++++++ tools/calib/analyze_kv.py | 332 ++++++++++++++ tools/calib/pca_kv_feasibility.py | 139 ++++++ tools/calib/rans_nvfp4.py | 112 +++++ 20 files changed, 1915 insertions(+) create mode 100644 include/ninfer/ops/cold_i8.h create mode 100644 include/ninfer/ops/entropy_cold_requant.h create mode 100644 src/ops/kernel/cold_i8_kernels.cuh create mode 100644 src/ops/kernel/entropy_cold_requant_kernels.cuh create mode 100644 src/ops/launcher/cold_i8.cu create mode 100644 src/ops/launcher/cold_i8.h create mode 100644 src/ops/launcher/entropy_cold_requant.cu create mode 100644 src/ops/launcher/entropy_cold_requant.h create mode 100644 src/ops/wrapper/cold_i8.cpp create mode 100644 src/ops/wrapper/entropy_cold_requant.cpp create mode 100644 src/targets/qwen3_6/impl/runtime/kv_calibration.h create mode 100644 tests/ops/test_cold_i8.cpp create mode 100644 tests/ops/test_entropy_cold_requant.cpp create mode 100644 tools/calib/analyze_kv.py create mode 100644 tools/calib/pca_kv_feasibility.py create mode 100644 tools/calib/rans_nvfp4.py diff --git a/apps/cli/main.cpp b/apps/cli/main.cpp index 933192da9b..40246f2457 100644 --- a/apps/cli/main.cpp +++ b/apps/cli/main.cpp @@ -283,6 +283,9 @@ int main(int argc, char** argv) { engine_options.speculative = cli.speculative; engine_options.enable_vision = cli.enable_vision; engine_options.use_cuda_graph = cli.use_cuda_graph; + engine_options.cold_policy = cli.cold_policy; + engine_options.cold_keep_tokens = cli.cold_keep_tokens; + engine_options.cold_host_bytes = cli.cold_host_bytes; // One CLI invocation owns exactly one request, so retained cross-request context has no // consumer and must not reserve an extra Device StateImage or run terminal capture. engine_options.context_cache.enabled = false; diff --git a/apps/cli/options.cpp b/apps/cli/options.cpp index b5c5798d9a..3d9b44b75d 100644 --- a/apps/cli/options.cpp +++ b/apps/cli/options.cpp @@ -85,6 +85,8 @@ std::string usage_text(const char* argv0) { " [--stop-token-id N]... [--stop ]... [--reasoning-stop ]...\n" " [--raw-output] [--print-token-ids] [--no-thinking] [--thinking-budget N]\n" " [--reasoning-effort low|medium|xhigh] [--vision]\n" + " [--cold-policy none|off|window|host] [--cold-keep-tokens N]\n" + " [--cold-host-bytes N[g|m|k]]\n" " [--no-cuda-graph]\n" "\n" "Streams answer content to stdout and reasoning plus diagnostics to stderr.\n" @@ -152,6 +154,17 @@ Options parse_options(int argc, char** argv) { options.reasoning_effort = parse_reasoning_effort(value(arg)); } else if (arg == "--vision") { options.enable_vision = true; + } else if (arg == "--cold-policy") { + const std::string v = value(arg); + if (v == "none" || v == "off") { options.cold_policy = ColdPolicy::None; } + else if (v == "window") { options.cold_policy = ColdPolicy::Window; } + else if (v == "host") { options.cold_policy = ColdPolicy::Host; } + else { throw std::invalid_argument("invalid cold-policy: " + v); } + options.cold_keep_tokens = 128; + } else if (arg == "--cold-keep-tokens") { + options.cold_keep_tokens = parse_u32(value(arg), "cold-keep-tokens"); + } else if (arg == "--cold-host-bytes") { + options.cold_host_bytes = parse_u32(value(arg), "cold-host-bytes"); } else if (arg == "--no-cuda-graph") { options.use_cuda_graph = false; } else if (arg == "--stop-token-id") { diff --git a/apps/cli/options.h b/apps/cli/options.h index 3c0c2e0960..7a12f53391 100644 --- a/apps/cli/options.h +++ b/apps/cli/options.h @@ -27,6 +27,9 @@ struct Options { SpeculativeOptions speculative; bool enable_vision = false; bool use_cuda_graph = true; + ColdPolicy cold_policy = ColdPolicy::None; + std::uint32_t cold_keep_tokens = 128; + std::uint64_t cold_host_bytes = 4ULL << 30; bool raw_output = false; bool print_token_ids = false; diff --git a/include/ninfer/ops/cold_i8.h b/include/ninfer/ops/cold_i8.h new file mode 100644 index 0000000000..e00c5782f9 --- /dev/null +++ b/include/ninfer/ops/cold_i8.h @@ -0,0 +1,30 @@ +#pragma once + +#include + +#include + +namespace ninfer::ops { + +// Fixed raw-slot size: 16 B header + 8192 B E2M1 nibbles + 1024 B E4M3 g16 +// scales (see ops/kernel/cold_i8_kernels.cuh for the composition). +inline constexpr std::int32_t kColdI8SlotBytes = 9232; + +// Raw cold-slot codec for INT8-tier KV pages (entropy-cold revision 2b). +// +// Slot layout (9232 B, fixed, no overflow): 16 B header | 8192 B packed +// E2M1 nibbles | 1024 B E4M3 group-16 scales, both in the nvfp4 page-major +// geometry produced by entropy_cold_requant_raw(Int8G64). Measured on real +// Qwen3.8 INT8 planes: NMSE 0.012-0.014 (inside the accepted NVFP4-layer +// envelope) at 1.83x per head-page vs the raw int8 plane (16896 B). +void cold_i8_slot_pack_raw(const std::uint8_t* src_codes, const std::uint8_t* src_scales, + int kv_heads, int page_count, std::uint8_t* slots, + std::int32_t* slot_valid, cudaStream_t stream); + +// Inverse: unpack slots into the INT8 tier's native planes (int8 codes + +// fp16 group-64 scales). Used by the warm-restore path. +void cold_i8_slot_restore_raw(const std::uint8_t* slots, int kv_heads, int page_count, + std::int8_t* dst_codes, void* dst_scales_fp16, + cudaStream_t stream); + +} // namespace ninfer::ops diff --git a/include/ninfer/ops/entropy_cold_requant.h b/include/ninfer/ops/entropy_cold_requant.h new file mode 100644 index 0000000000..f134f0d167 --- /dev/null +++ b/include/ninfer/ops/entropy_cold_requant.h @@ -0,0 +1,32 @@ +#pragma once + +#include + +#include + +namespace ninfer::ops { + +// Cold-page requantization for the entropy slot codec (revision 2). +// Requantizes one or more stored KV page planes (E2M1 g16 NVFP4 or int8 g64) +// into fresh NVFP4 planes whose codes carry one E4M3FN scale per 64 channels +// (replicated into the four g16 scale slots). The output layout matches the +// page-major nvfp4 planes entropy_nvfp4_slot_encode_raw consumes, so the slot +// codec, decode path, scale scatter, and attention producers stay unchanged; +// the g64 requant is what makes the rANS streams compressible (measured +// 2.0-2.6 bits/code vs ~4.0 for g16 storage codes on Qwen3.8 frames). +enum class EntropyColdRequantMode : int { + Nvfp4G16 = 0, + Int8G64 = 1, + Iso3VG16 = 2, +}; + +// src planes use page-major strides: codes kv_heads*8192 (nvfp4) or +// kv_heads*16384 (int8) bytes per page, scales kv_heads*1024 (nvfp4) or +// kv_heads*512 (int8). dst planes are nvfp4 page-major: codes +// [128, 64, kv_heads, page_count], scales [16, 64, kv_heads, page_count]. +void entropy_cold_requant_raw(const std::uint8_t* src_codes, const std::uint8_t* src_scales, + EntropyColdRequantMode mode, int kv_heads, int page_count, + std::uint8_t* dst_codes, std::uint8_t* dst_scales, + cudaStream_t stream); + +} // namespace ninfer::ops diff --git a/include/ninfer/types.h b/include/ninfer/types.h index 677c381e88..278d255018 100644 --- a/include/ninfer/types.h +++ b/include/ninfer/types.h @@ -37,6 +37,15 @@ enum class EnginePurpose : std::uint8_t { CausalScoring, }; +// Entropy-coded cold KV pool. Window compresses fully-written pages that +// are at least cold_keep_tokens behind the decode frontier; attention +// producers decode those pages inline from fixed-size slots. +enum class ColdPolicy : std::uint8_t { + None, + Window, + Host, +}; + enum class KvCapacityMode : std::uint8_t { Explicit, Automatic, @@ -124,6 +133,10 @@ struct EngineOptions { bool use_cuda_graph = true; ContextCacheOptions context_cache; ContextCostOptions context_cost; + ColdPolicy cold_policy = ColdPolicy::None; + std::uint32_t cold_keep_tokens = 128; + // Pinned host-memory budget for ColdPolicy::Host offload. Default 4 GiB. + std::uint64_t cold_host_bytes = 4ULL << 30; LoadProgress load_progress; }; diff --git a/src/ops/kernel/cold_i8_kernels.cuh b/src/ops/kernel/cold_i8_kernels.cuh new file mode 100644 index 0000000000..9320c98bd9 --- /dev/null +++ b/src/ops/kernel/cold_i8_kernels.cuh @@ -0,0 +1,105 @@ +#pragma once + +// ninfer::ops::detail - raw cold-slot codec for INT8-tier KV pages (v2b). +// +// The INT8 tier's requantized E2M1 codes are near-uniform (measured H +// 3.6-3.9 on real planes), so rANS gains nothing; this slot format stores +// the requant output verbatim with a fixed layout and no overflow path: +// +// [ 16 B header | 8192 B packed E2M1 nibbles | 1024 B E4M3 g16 scales ] +// +// The nibble/scale planes use the same page-major geometry the NVFP4 tier +// stores (codes [128, 64, kv_heads, pages], scales [16, 64, kv_heads, +// pages]), produced by entropy_cold_requant's Int8G64 mode. Restore +// converts a slot back into the INT8 tier's native planes (int8 codes + +// fp16 group-64 scales) with the upper-bound group scale, adding <0.4% +// quantization noise on top of the requant's measured 0.012 NMSE. +// +// Kernel DEFINITIONS live only in ops/launcher/cold_i8.cu; this header +// declares them plus the shared device helpers so attention kernels can +// include it without duplicate device-link definitions. + +#include "ninfer/ops/cold_i8.h" +#include "ops/kernel/gqa_attention_kv_nvfp4.cuh" + +#include +#include + +#include + +namespace ninfer::ops::detail { + +inline constexpr int kColdI8SlotHeaderBytes = 16; +inline constexpr int kColdI8SlotCodeBytes = 8192; // 64 rows x 128 B nibbles +inline constexpr int kColdI8SlotScaleBytes = 1024; // 64 rows x 16 B E4M3 +static_assert(kColdI8SlotHeaderBytes + kColdI8SlotCodeBytes + kColdI8SlotScaleBytes == + ninfer::ops::kColdI8SlotBytes); +inline constexpr std::uint32_t kColdI8SlotMagic = 0x49384352u; // "RC8I" + +// Pack one requantized (page, head, plane) into a raw slot. src uses the +// nvfp4 page-major geometry entropy_cold_requant emits; slots layout is +// [slot_bytes, kv_heads, 2, pages] with V one nb[2] step past K. +__global__ void cold_i8_slot_pack_kernel(const std::uint8_t* __restrict__ src_codes, + const std::uint8_t* __restrict__ src_scales, + int kv_heads, + std::uint8_t* __restrict__ slots, + std::int32_t* __restrict__ slot_valid); + +// Warm restore: unpack one (page, head, plane) slot into the INT8 tier's +// native planes (int8 codes [256,64,kv_heads,pages], fp16 scales +// [4,64,kv_heads,pages]). One block per (head, page); 256 threads split +// the 64 rows. +__global__ void cold_i8_slot_restore_kernel(const std::uint8_t* __restrict__ slots, + int kv_heads, std::int8_t* __restrict__ dst_codes, + __half* __restrict__ dst_scales); + +// Slot region accessors for producers. +__device__ __forceinline__ const std::uint8_t* +cold_i8_slot_scales(const std::uint8_t* slot) { + return slot + kColdI8SlotHeaderBytes + kColdI8SlotCodeBytes; +} + +__device__ __forceinline__ const std::uint8_t* +cold_i8_slot_codes(const std::uint8_t* slot) { + return slot + kColdI8SlotHeaderBytes; +} + +// Decode one key row of a raw slot into INT8-tier native form: 256 int8 +// codes plus one fp16 scale per 64-channel group. The group scale is the +// upper bound 6*max(e4m3 sub-scales)/127 so no amax scan of the decoded +// values is needed; codes clamp at 127 so fp16 rounding-down is safe. +// Used by the warm-restore kernel and the attention producers' cold staging. +__device__ __forceinline__ void cold_i8_decode_row(const std::uint8_t* slot, int row, + std::int8_t* codes_out, // 256, d-major + __half* scales_out) { // 4 groups + const std::uint8_t* row_codes = cold_i8_slot_codes(slot) + row * 128; + const std::uint8_t* row_scales = cold_i8_slot_scales(slot) + row * 16; +#pragma unroll + for (int g = 0; g < 4; ++g) { + float mx = 0.0f; +#pragma unroll + for (int s = 0; s < 4; ++s) { + mx = fmaxf(mx, gqa_kv_nvfp4_e4m3_to_f32(row_scales[g * 4 + s])); + } + const float scale = mx * 6.0f / 127.0f; + scales_out[g] = __float2half(scale); + const float inv = scale > 0.0f ? 1.0f / scale : 0.0f; +#pragma unroll + for (int i = 0; i < 64; i += 2) { + const int d = g * 64 + i; + const std::uint8_t b = row_codes[d >> 1]; + const float v0 = gqa_kv_nvfp4_e2m1_to_f32(b & 0x0F) * + gqa_kv_nvfp4_e4m3_to_f32(row_scales[d >> 4]); + const float v1 = gqa_kv_nvfp4_e2m1_to_f32(b >> 4) * + gqa_kv_nvfp4_e4m3_to_f32(row_scales[(d + 1) >> 4]); + int c0 = __float2int_rn(v0 * inv); + int c1 = __float2int_rn(v1 * inv); + c0 = max(-127, min(127, c0)); + c1 = max(-127, min(127, c1)); + codes_out[d] = static_cast(c0); + codes_out[d + 1] = static_cast(c1); + } + } +} + +} // namespace ninfer::ops::detail diff --git a/src/ops/kernel/entropy_cold_requant_kernels.cuh b/src/ops/kernel/entropy_cold_requant_kernels.cuh new file mode 100644 index 0000000000..8269f69028 --- /dev/null +++ b/src/ops/kernel/entropy_cold_requant_kernels.cuh @@ -0,0 +1,121 @@ +#pragma once + +// ninfer::ops::detail - cold-page requantization kernel for the entropy slot +// codec (revision 2). +// +// The revision-1 slot codec rANS-encoded the stored group-16 NVFP4 code +// nibbles directly; on real Qwen3.8 pages those codes are near-uniform +// (~4.0 bits/nibble) so the fixed-slot encoder fell back to the uncompressed +// plane and the cold pool saved nothing. Requantizing the same values with +// one E4M3 scale per 64 channels skews the code distribution enough for the +// unchanged order-0 rANS to compress it (measured 2.0-2.6 bits/code on +// captured .kvc frames; 3-bit signed requant was measured worse: 0.04-0.13 +// NMSE vs 0.011-0.022 for g64 E2M1). +// +// The kernel reads one page plane in its stored format and writes fresh +// NVFP4 planes: packed E2M1 codes plus one E4M3FN scale per 64-channel group +// replicated into the four group-16 scale slots it covers. Downstream slot +// encode, decode, scale scatter, and attention producers stay byte-identical +// to revision 1; only the codes fed into the rANS change. + +#include "ops/kernel/gqa_attention_kv_nvfp4.cuh" +#include "ops/kernel/gqa_attention_kv_quant.cuh" +#include "ops/kernel/gqa_attention_prefill_nvfp4.cuh" // gqa_iso3_nibble / gqa_iso3_decode +#include "ops/launcher/entropy_cold_requant.h" + +#include + +#include + +namespace ninfer::ops::detail { + +// One block per (kv_head, page); 256 threads = 64 token rows x 4 groups of +// 64 channels. dst planes use the nvfp4 page-major layout the slot encoder +// expects: codes [128, 64, kv_heads, pages], scales [16, 64, kv_heads, pages]. +__global__ void entropy_cold_requant_kernel(const std::uint8_t* __restrict__ src_codes, + const std::uint8_t* __restrict__ src_scales, + ColdRequantSource mode, int kv_heads, + std::uint8_t* __restrict__ dst_codes, + std::uint8_t* __restrict__ dst_scales) { + const int head = static_cast(blockIdx.x); + const int page = static_cast(blockIdx.y); + const int token = static_cast(threadIdx.x) >> 2; + const int group = static_cast(threadIdx.x) & 3; + const int lane0 = group * 64; + + // Row strides in bytes: nvfp4 codes 256/2, nvfp4 scales 256/16, + // int8 codes 256, int8 scales 4 fp16 = 8. + const std::int64_t page_rows = static_cast(kPagedKVPageSize); + const std::int64_t head_off = + page_rows * (static_cast(head) + + static_cast(kv_heads) * static_cast(page)); + + float vals[64]; + if (mode == ColdRequantSource::Nvfp4G16) { + const std::uint8_t* codes = src_codes + 128 * head_off + 128 * token; + const std::uint8_t* scales = src_scales + 16 * head_off + 16 * token; +#pragma unroll + for (int i = 0; i < 64; ++i) { + const int d = lane0 + i; + const std::uint8_t byte = codes[d >> 1]; + const std::uint8_t nib = (d & 1) != 0 ? static_cast(byte >> 4) + : static_cast(byte & 0x0F); + vals[i] = gqa_kv_nvfp4_e2m1_to_f32(nib) * gqa_kv_nvfp4_e4m3_to_f32(scales[d >> 4]); + } + } else if (mode == ColdRequantSource::Iso3VG16) { + // The global NVFP4 tier stores V as ISO3 sign-magnitude INT3 nibbles in + // the same two-per-byte plane geometry. Requant keeps the native ISO3 + // nibble semantics so the warm producers' dequant path is unchanged; + // only the scales are re-derived per 64 channels. + const std::uint8_t* codes = src_codes + 128 * head_off + 128 * token; + const std::uint8_t* scales = src_scales + 16 * head_off + 16 * token; +#pragma unroll + for (int i = 0; i < 64; ++i) { + const int d = lane0 + i; + const std::uint8_t byte = codes[d >> 1]; + const std::uint8_t nib = (d & 1) != 0 ? static_cast(byte >> 4) + : static_cast(byte & 0x0F); + vals[i] = gqa_iso3_decode(nib) * gqa_kv_nvfp4_e4m3_to_f32(scales[d >> 4]); + } + } else { + const std::int8_t* codes = reinterpret_cast(src_codes) + 256 * head_off + + 256 * token; + const __half* scales = reinterpret_cast(src_scales + 8 * head_off + + 8 * token); + const float s = __half2float(scales[group]); +#pragma unroll + for (int i = 0; i < 64; ++i) { + vals[i] = static_cast(codes[lane0 + i]) * s; + } + } + + float amax = 0.0f; +#pragma unroll + for (int i = 0; i < 64; ++i) { amax = fmaxf(amax, fabsf(vals[i])); } + const bool iso3_out = mode == ColdRequantSource::Iso3VG16; + const std::uint8_t scale_byte = + gqa_kv_nvfp4_fp32_to_e4m3(fmaxf(amax / (iso3_out ? 7.0f : 6.0f), 0x1p-9f)); + const float s = gqa_kv_nvfp4_e4m3_to_f32(scale_byte); + + std::uint8_t* dst_c = dst_codes + 128 * head_off + 128 * token; +#pragma unroll + for (int i = 0; i < 64; i += 2) { + std::uint8_t lo; + std::uint8_t hi; + if (iso3_out) { + lo = gqa_iso3_nibble(vals[i], s); + hi = gqa_iso3_nibble(vals[i + 1], s); + } else { + lo = gqa_kv_nvfp4_e2m1_nibble(vals[i] / s); + hi = gqa_kv_nvfp4_e2m1_nibble(vals[i + 1] / s); + } + dst_c[(lane0 + i) >> 1] = static_cast(lo | (hi << 4)); + } + std::uint8_t* dst_s = dst_scales + 16 * head_off + 16 * token; + dst_s[group * 4 + 0] = scale_byte; + dst_s[group * 4 + 1] = scale_byte; + dst_s[group * 4 + 2] = scale_byte; + dst_s[group * 4 + 3] = scale_byte; +} + +} // namespace ninfer::ops::detail diff --git a/src/ops/launcher/cold_i8.cu b/src/ops/launcher/cold_i8.cu new file mode 100644 index 0000000000..eaf190721a --- /dev/null +++ b/src/ops/launcher/cold_i8.cu @@ -0,0 +1,85 @@ +#include "ops/launcher/cold_i8.h" + +#include "core/device.h" +#include "ops/kernel/cold_i8_kernels.cuh" + +#include + +#include + +namespace ninfer::ops::detail { + +// Kernel definitions live in this single TU: the shared header only +// declares them so attention kernels can include its device helpers +// without duplicate device-link definitions. +__global__ void cold_i8_slot_pack_kernel(const std::uint8_t* __restrict__ src_codes, + const std::uint8_t* __restrict__ src_scales, + int kv_heads, + std::uint8_t* __restrict__ slots, + std::int32_t* __restrict__ slot_valid) { + const int head = static_cast(blockIdx.x); + const int page = static_cast(blockIdx.y); + const std::int64_t plane = static_cast(head) + + static_cast(kv_heads) * page; + const std::uint8_t* src_c = src_codes + plane * kColdI8SlotCodeBytes; + const std::uint8_t* src_s = src_scales + plane * kColdI8SlotScaleBytes; + std::uint8_t* slot = slots + plane * kColdI8SlotBytes; + if (threadIdx.x == 0) { + *reinterpret_cast(slot) = kColdI8SlotMagic; + *reinterpret_cast(slot + 4) = 1; // version + *reinterpret_cast(slot + 6) = 1; // flags: valid + } + __syncthreads(); + for (int i = static_cast(threadIdx.x); i < kColdI8SlotCodeBytes; i += 256) { + slot[kColdI8SlotHeaderBytes + i] = src_c[i]; + } + for (int i = static_cast(threadIdx.x); i < kColdI8SlotScaleBytes; i += 256) { + slot[kColdI8SlotHeaderBytes + kColdI8SlotCodeBytes + i] = src_s[i]; + } + if (threadIdx.x == 0) { + slot_valid[plane] = 1; // fixed layout: always valid, no overflow path + } +} + +__global__ void cold_i8_slot_restore_kernel(const std::uint8_t* __restrict__ slots, + int kv_heads, std::int8_t* __restrict__ dst_codes, + __half* __restrict__ dst_scales) { + const int head = static_cast(blockIdx.x); + const int page = static_cast(blockIdx.y); + const std::int64_t plane = static_cast(head) + + static_cast(kv_heads) * page; + const std::uint8_t* slot = slots + plane * kColdI8SlotBytes; + std::int8_t* codes = dst_codes + plane * (64 * 256); + __half* scales = dst_scales + plane * (64 * 4); + const int row0 = static_cast(threadIdx.x) >> 2; // 64 rows + const int lane = static_cast(threadIdx.x) & 3; // 4 quarter-rows + if (lane != 0) { return; } + std::int8_t row_codes[256]; + __half row_scales[4]; + cold_i8_decode_row(slot, row0, row_codes, row_scales); +#pragma unroll + for (int g = 0; g < 4; ++g) { scales[row0 * 4 + g] = row_scales[g]; } +#pragma unroll + for (int d = 0; d < 256; ++d) { codes[row0 * 256 + d] = row_codes[d]; } +} + + +void cold_i8_slot_pack_launch(const std::uint8_t* src_codes, const std::uint8_t* src_scales, + int kv_heads, int page_count, std::uint8_t* slots, + std::int32_t* slot_valid, cudaStream_t stream) { + const dim3 grid(kv_heads, page_count); + cold_i8_slot_pack_kernel<<>>(src_codes, src_scales, kv_heads, slots, + slot_valid); + CUDA_CHECK(cudaGetLastError()); +} + +void cold_i8_slot_restore_launch(const std::uint8_t* slots, int kv_heads, int page_count, + std::int8_t* dst_codes, void* dst_scales_fp16, + cudaStream_t stream) { + const dim3 grid(kv_heads, page_count); + cold_i8_slot_restore_kernel<<>>( + slots, kv_heads, dst_codes, static_cast<__half*>(dst_scales_fp16)); + CUDA_CHECK(cudaGetLastError()); +} + +} // namespace ninfer::ops::detail diff --git a/src/ops/launcher/cold_i8.h b/src/ops/launcher/cold_i8.h new file mode 100644 index 0000000000..8fa9e78aa1 --- /dev/null +++ b/src/ops/launcher/cold_i8.h @@ -0,0 +1,17 @@ +#pragma once + +#include + +#include + +namespace ninfer::ops::detail { + +void cold_i8_slot_pack_launch(const std::uint8_t* src_codes, const std::uint8_t* src_scales, + int kv_heads, int page_count, std::uint8_t* slots, + std::int32_t* slot_valid, cudaStream_t stream); + +void cold_i8_slot_restore_launch(const std::uint8_t* slots, int kv_heads, int page_count, + std::int8_t* dst_codes, void* dst_scales_fp16, + cudaStream_t stream); + +} // namespace ninfer::ops::detail diff --git a/src/ops/launcher/entropy_cold_requant.cu b/src/ops/launcher/entropy_cold_requant.cu new file mode 100644 index 0000000000..ff669c981c --- /dev/null +++ b/src/ops/launcher/entropy_cold_requant.cu @@ -0,0 +1,20 @@ +#include "ops/launcher/entropy_cold_requant.h" + +#include "core/device.h" +#include "ops/kernel/entropy_cold_requant_kernels.cuh" + +#include + +namespace ninfer::ops::detail { + +void entropy_cold_requant_raw_launch(const std::uint8_t* src_codes, + const std::uint8_t* src_scales, ColdRequantSource mode, + int kv_heads, int page_count, std::uint8_t* dst_codes, + std::uint8_t* dst_scales, cudaStream_t stream) { + const dim3 grid(kv_heads, page_count); + entropy_cold_requant_kernel<<>>(src_codes, src_scales, mode, kv_heads, + dst_codes, dst_scales); + CUDA_CHECK(cudaGetLastError()); +} + +} // namespace ninfer::ops::detail diff --git a/src/ops/launcher/entropy_cold_requant.h b/src/ops/launcher/entropy_cold_requant.h new file mode 100644 index 0000000000..d45f984024 --- /dev/null +++ b/src/ops/launcher/entropy_cold_requant.h @@ -0,0 +1,20 @@ +#pragma once + +#include + +#include + +namespace ninfer::ops::detail { + +enum class ColdRequantSource : int { + Nvfp4G16 = 0, // E2M1 nibbles + E4M3 g16 scales (K planes) + Int8G64 = 1, // int8 codes + fp16 g64 scales + Iso3VG16 = 2, // ISO3 sign-magnitude INT3 nibbles + E4M3 g16 scales (V planes) +}; + +void entropy_cold_requant_raw_launch(const std::uint8_t* src_codes, + const std::uint8_t* src_scales, ColdRequantSource mode, + int kv_heads, int page_count, std::uint8_t* dst_codes, + std::uint8_t* dst_scales, cudaStream_t stream); + +} // namespace ninfer::ops::detail diff --git a/src/ops/wrapper/cold_i8.cpp b/src/ops/wrapper/cold_i8.cpp new file mode 100644 index 0000000000..610f5da7dc --- /dev/null +++ b/src/ops/wrapper/cold_i8.cpp @@ -0,0 +1,25 @@ +#include "ninfer/ops/cold_i8.h" + +#include "ops/launcher/cold_i8.h" + +#include + +#include + +namespace ninfer::ops { + +void cold_i8_slot_pack_raw(const std::uint8_t* src_codes, const std::uint8_t* src_scales, + int kv_heads, int page_count, std::uint8_t* slots, + std::int32_t* slot_valid, cudaStream_t stream) { + detail::cold_i8_slot_pack_launch(src_codes, src_scales, kv_heads, page_count, slots, + slot_valid, stream); +} + +void cold_i8_slot_restore_raw(const std::uint8_t* slots, int kv_heads, int page_count, + std::int8_t* dst_codes, void* dst_scales_fp16, + cudaStream_t stream) { + detail::cold_i8_slot_restore_launch(slots, kv_heads, page_count, dst_codes, + static_cast<__half*>(dst_scales_fp16), stream); +} + +} // namespace ninfer::ops diff --git a/src/ops/wrapper/entropy_cold_requant.cpp b/src/ops/wrapper/entropy_cold_requant.cpp new file mode 100644 index 0000000000..63527cf1a1 --- /dev/null +++ b/src/ops/wrapper/entropy_cold_requant.cpp @@ -0,0 +1,23 @@ +#include "ninfer/ops/entropy_cold_requant.h" + +#include "ops/launcher/entropy_cold_requant.h" + +#include + +namespace ninfer::ops { + +void entropy_cold_requant_raw(const std::uint8_t* src_codes, const std::uint8_t* src_scales, + EntropyColdRequantMode mode, int kv_heads, int page_count, + std::uint8_t* dst_codes, std::uint8_t* dst_scales, + cudaStream_t stream) { + detail::ColdRequantSource source = detail::ColdRequantSource::Nvfp4G16; + if (mode == EntropyColdRequantMode::Int8G64) { + source = detail::ColdRequantSource::Int8G64; + } else if (mode == EntropyColdRequantMode::Iso3VG16) { + source = detail::ColdRequantSource::Iso3VG16; + } + detail::entropy_cold_requant_raw_launch(src_codes, src_scales, source, kv_heads, page_count, + dst_codes, dst_scales, stream); +} + +} // namespace ninfer::ops diff --git a/src/targets/qwen3_6/impl/runtime/kv_calibration.h b/src/targets/qwen3_6/impl/runtime/kv_calibration.h new file mode 100644 index 0000000000..5cf6b94ef2 --- /dev/null +++ b/src/targets/qwen3_6/impl/runtime/kv_calibration.h @@ -0,0 +1,116 @@ +#pragma once +#include "targets/qwen3_6/impl/runtime/instance.h" +// Qwen3.6 family runtime implementation; instantiated only by exact variants. + +#include "core/device.h" +#include "core/dtype.h" +#include "core/tensor.h" + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace ninfer::targets::qwen3_6::detail::NINFER_QWEN36_RUNTIME_NS::schedule { + +// Offline KV calibration capture. When EngineOptions.kv_calibration_dir is set, +// text prefill copies the exact post-RoPE K and V tensors quantized into the +// paged KV cache (per full-attention layer and per chunk) to the host and +// appends one framed binary record per layer/chunk. The Python analyzer in +// tools/calib consumes these records; inference never reads them back. +class KvCalibrationCapture { +public: + explicit KvCalibrationCapture(std::filesystem::path directory) + : directory_(std::move(directory)) { + if (directory_.empty()) { + throw std::invalid_argument("KV calibration directory must not be empty"); + } + std::filesystem::create_directories(directory_); + } + + void capture(std::uint32_t full_layer, const Tensor& k, const Tensor& v, + const Tensor& positions) { + if (k.dtype != DType::BF16 || v.dtype != DType::BF16 || positions.dtype != DType::I32 || + !k.is_contiguous() || !v.is_contiguous() || !positions.is_contiguous() || + k.data == nullptr || v.data == nullptr || positions.data == nullptr) { + throw std::invalid_argument( + "KV calibration capture requires contiguous BF16 K/V and I32 positions"); + } + if (k.ne[0] != v.ne[0] || k.ne[1] != v.ne[1] || k.ne[2] != v.ne[2] || k.ne[3] != 1 || + v.ne[3] != 1 || positions.ne[0] != k.ne[2] || positions.ne[1] != 1 || + positions.ne[2] != 1 || positions.ne[3] != 1) { + throw std::invalid_argument("KV calibration capture tensor shapes do not match"); + } + const auto head_dim = static_cast(k.ne[0]); + const auto kv_heads = static_cast(k.ne[1]); + const auto tokens = static_cast(k.ne[2]); + if (head_dim == 0 || kv_heads == 0 || tokens == 0 || + tokens > static_cast(std::numeric_limits::max())) { + throw std::invalid_argument("KV calibration capture shapes are out of range"); + } + + std::vector positions_host(tokens); + std::vector k_host(k.bytes()); + std::vector v_host(v.bytes()); + CUDA_CHECK(cudaMemcpy(positions_host.data(), positions.data, positions.bytes(), + cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaMemcpy(k_host.data(), k.data, k.bytes(), cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaMemcpy(v_host.data(), v.data, v.bytes(), cudaMemcpyDeviceToHost)); + + Header header{}; + std::memcpy(header.magic, kMagic, sizeof(header.magic)); + header.header_bytes = sizeof(Header); + header.full_layer = full_layer; + header.head_dim = head_dim; + header.kv_heads = kv_heads; + header.tokens = tokens; + header.record_index = record_index_; + header.first_position = positions_host.front(); + header.last_position = positions_host.back(); + + const std::filesystem::path path = + directory_ / (std::to_string(record_index_) + ".kvc"); + std::ofstream out(path, std::ios::binary | std::ios::trunc); + if (!out) { throw std::runtime_error("cannot create KV calibration record: " + path.string()); } + out.write(reinterpret_cast(&header), sizeof(header)); + out.write(reinterpret_cast(positions_host.data()), + static_cast(positions_host.size() * sizeof(std::int32_t))); + out.write(reinterpret_cast(k_host.data()), + static_cast(k_host.size())); + out.write(reinterpret_cast(v_host.data()), + static_cast(v_host.size())); + if (!out) { throw std::runtime_error("failed to write KV calibration record: " + path.string()); } + ++record_index_; + } + + [[nodiscard]] std::uint32_t record_count() const noexcept { return record_index_; } + +private: + static constexpr char kMagic[16] = {'N', 'I', 'N', 'F', 'E', 'R', 'K', 'V', + 'C', 'A', 'L', '1', 0, 0, 0, 0}; + struct Header { + char magic[16]; + std::uint32_t header_bytes; + std::uint32_t full_layer; + std::uint32_t head_dim; + std::uint32_t kv_heads; + std::uint32_t tokens; + std::uint32_t record_index; + std::int32_t first_position; + std::int32_t last_position; + std::uint32_t reserved[4]; + }; + static_assert(sizeof(Header) == 64); + + std::filesystem::path directory_; + std::uint32_t record_index_ = 0; +}; + +} // namespace ninfer::targets::qwen3_6::detail::NINFER_QWEN36_RUNTIME_NS::schedule diff --git a/tests/ops/test_cold_i8.cpp b/tests/ops/test_cold_i8.cpp new file mode 100644 index 0000000000..e410e94d95 --- /dev/null +++ b/tests/ops/test_cold_i8.cpp @@ -0,0 +1,274 @@ +#include "ninfer/ops/cold_i8.h" +#include "ninfer/ops/entropy_cold_requant.h" +#include "ops/op_tester.h" + +#include +#include +#include +#include +#include +#include + +using namespace ninfer; +using namespace ninfer::test; + +namespace { + +constexpr int kHeadDim = 256; +constexpr int kPageRows = 64; +constexpr int kKvHeads = 4; +constexpr int kI8CodeB = kHeadDim * kPageRows; // 16384 per head +constexpr int kI8ScaleB = kHeadDim / 64 * 2 * kPageRows; // 512 fp16 bytes +constexpr int kNvCodeB = kHeadDim / 2 * kPageRows; // 8192 +constexpr int kNvScaleB = kHeadDim / 16 * kPageRows; // 1024 + +std::uint8_t e4m3_rne(float x) { + if (!(x > 0.0f)) { return 0; } + std::uint32_t bits; + std::memcpy(&bits, &x, 4); + const std::uint32_t sign = (bits >> 24) & 0x80u; + int exponent = static_cast((bits >> 23) & 0xffu) - 127 + 7; + if (exponent >= 15) { return static_cast(sign | (15u << 3) | 7u); } + if (exponent <= 0) { + int mantissa = static_cast(std::nearbyint(x * 512.0f)); + if (mantissa <= 0) { return static_cast(sign); } + if (mantissa >= 8) { return static_cast(sign | (1u << 3)); } + return static_cast(sign | mantissa); + } + std::uint32_t mantissa = (bits >> 20) & 0x7u; + const std::uint32_t guard = (bits >> 19) & 1u; + const std::uint32_t sticky = bits & 0x7ffffu; + if (guard && (sticky || (mantissa & 1u))) { + mantissa += 1; + if (mantissa > 7) { + mantissa = 0; + exponent += 1; + if (exponent >= 15) { return static_cast(sign | (15u << 3) | 7u); } + } + } + return static_cast(sign | (exponent << 3) | mantissa); +} + +float e4m3_to_f32(std::uint8_t byte) { + const int e = (byte >> 3) & 0xF; + const int m = byte & 0x7; + if (e == 0) { return static_cast(m) / 512.0f; } + return (1.0f + static_cast(m) / 8.0f) * std::pow(2.0f, static_cast(e - 7)); +} + +float e2m1_to_f32(std::uint8_t code) { + static const float mag[8] = {0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f}; + const float v = mag[code & 0x7]; + return (code & 0x8) != 0 ? -v : v; +} + +std::uint8_t e2m1_code(float x) { + const float a = std::fabs(x); + std::uint8_t c; + if (a < 0.25f) { c = 0; } + else if (a < 0.75f) { c = 1; } + else if (a < 1.25f) { c = 2; } + else if (a < 1.75f) { c = 3; } + else if (a < 2.5f) { c = 4; } + else if (a < 3.5f) { c = 5; } + else if (a < 5.0f) { c = 6; } + else { c = 7; } + if (x < 0.0f) { c |= 0x08u; } + return c; +} + +float half_to_float(std::uint16_t h) { + const std::uint32_t sign = (h >> 15) & 1u; + const std::uint32_t exp = (h >> 10) & 0x1Fu; + const std::uint32_t man = h & 0x3FFu; + float out; + if (exp == 0) { + out = std::ldexp(static_cast(man), -24); + } else { + out = std::ldexp(1024.0f + static_cast(man), static_cast(exp) - 25); + } + return sign != 0 ? -out : out; +} + +std::uint16_t float_to_half(float f) { + if (f <= 0.0f) { return 0; } + std::uint32_t x; + std::memcpy(&x, &f, 4); + const std::uint32_t sign = (x >> 16) & 0x8000u; + int exponent = static_cast((x >> 23) & 0xFFu) - 127 + 15; + std::uint32_t mantissa = (x >> 13) & 0x3FFu; + if (exponent <= 0) { + const std::uint32_t man_full = (x & 0x7FFFFFu) | 0x800000u; + const int shift = 14 - exponent + 1; + mantissa = man_full >> shift; + const std::uint32_t round_bit = (man_full >> (shift - 1)) & 1u; + if (round_bit != 0) { mantissa += 1; } + exponent = 0; + } else { + const std::uint32_t round_bit = (x >> 12) & 1u; + const std::uint32_t sticky = x & 0xFFFu; + if (round_bit != 0 && (sticky != 0 || (mantissa & 1u) != 0)) { + mantissa += 1; + if (mantissa > 0x3FFu) { + mantissa = 0; + exponent += 1; + } + } + } + if (exponent >= 31) { return static_cast(sign | (31u << 10)); } + return static_cast(sign | (static_cast(exponent) << 10) | + mantissa); +} + +} // namespace + +int main() { + std::mt19937 rng(20260830); + int failures = 0; + const auto check = [&](bool ok, const char* what) { + if (!ok) { + std::printf("FAIL: %s\n", what); + ++failures; + } + }; + + // 1) synthetic int8 page per head + std::vector src_codes(static_cast(kKvHeads) * kI8CodeB); + std::vector src_scales(static_cast(kKvHeads) * kPageRows * 4); + std::normal_distribution noise(0.0f, 1.0f); + std::uniform_real_distribution level(0.01f, 50.0f); + for (int head = 0; head < kKvHeads; ++head) { + for (int row = 0; row < kPageRows; ++row) { + for (int g = 0; g < 4; ++g) { + const float amp = level(rng); + float amax = 1e-6f; + float vals[64]; + for (int i = 0; i < 64; ++i) { + float v = amp * noise(rng); + if (((row * 7 + i) % 251) == 0) { v *= 32.0f; } + vals[i] = v; + amax = std::fmax(amax, std::fabs(v)); + } + const float s = std::fmax(amax / 127.0f, 1e-30f); + const std::uint16_t sb = float_to_half(s); + src_scales[(static_cast(head) * kPageRows + row) * 4 + g] = sb; + const float sh = half_to_float(sb); + for (int i = 0; i < 64; ++i) { + float q = std::nearbyint(vals[i] / sh); + q = std::fmin(127.0f, std::fmax(-127.0f, q)); + src_codes[static_cast(head) * kI8CodeB + row * kHeadDim + + g * 64 + i] = static_cast(q); + } + } + } + } + + // 2) requant + pack on device + GuardedDeviceBuffer d_src_codes(static_cast(kKvHeads) * kI8CodeB); + GuardedDeviceBuffer d_src_scales(static_cast(kKvHeads) * kI8ScaleB); + GuardedDeviceBuffer d_rq_codes(static_cast(kKvHeads) * kNvCodeB); + GuardedDeviceBuffer d_rq_scales(static_cast(kKvHeads) * kNvScaleB); + GuardedDeviceBuffer d_slots(static_cast(kKvHeads) * 2 * ops::kColdI8SlotBytes); + GuardedDeviceBuffer d_valid(static_cast(kKvHeads) * 2 * sizeof(std::int32_t)); + GuardedDeviceBuffer d_out_codes(static_cast(kKvHeads) * kI8CodeB); + GuardedDeviceBuffer d_out_scales(static_cast(kKvHeads) * kI8ScaleB); + + d_src_codes.copy_from_host(src_codes.data(), d_src_codes.bytes()); + d_src_scales.copy_from_host(src_scales.data(), d_src_scales.bytes()); + + ops::entropy_cold_requant_raw( + static_cast(d_src_codes.data()), + static_cast(d_src_scales.data()), + ops::EntropyColdRequantMode::Int8G64, kKvHeads, 1, + static_cast(d_rq_codes.data()), + static_cast(d_rq_scales.data()), nullptr); + cuda_synchronize(); + auto* valid_k = static_cast(d_valid.data()); + auto* valid_v = valid_k + kKvHeads; + ops::cold_i8_slot_pack_raw(static_cast(d_rq_codes.data()), + static_cast(d_rq_scales.data()), kKvHeads, 1, + static_cast(d_slots.data()), valid_k, nullptr); + ops::cold_i8_slot_pack_raw( + static_cast(d_rq_codes.data()) + static_cast(kKvHeads) * kNvCodeB, + static_cast(d_rq_scales.data()) + static_cast(kKvHeads) * kNvScaleB, + kKvHeads, 1, + static_cast(d_slots.data()) + static_cast(kKvHeads) * ops::kColdI8SlotBytes, + valid_v, nullptr); + cuda_synchronize(); + ops::cold_i8_slot_restore_raw(static_cast(d_slots.data()), kKvHeads, 1, + static_cast(d_out_codes.data()), + d_out_scales.data(), nullptr); + cuda_synchronize(); + + std::vector got_codes(src_codes.size()); + std::vector got_scales(src_scales.size()); + std::vector got_valid(static_cast(kKvHeads) * 2); + d_out_codes.copy_to_host(got_codes.data(), d_out_codes.bytes()); + d_out_scales.copy_to_host(got_scales.data(), d_out_scales.bytes()); + d_valid.copy_to_host(got_valid.data(), d_valid.bytes()); + for (std::size_t i = 0; i < got_valid.size(); ++i) { + check(got_valid[i] == 1, "raw slot valid flag"); + } + + // 3) host oracle: decode -> requant g64 -> decode -> int8 re-encode with + // upper-bound group scale (mirror cold_i8_decode_row) + double num = 0.0, den = 0.0; + int byte_mismatch = 0; + int scale_mismatch = 0; + for (int head = 0; head < kKvHeads; ++head) { + for (int row = 0; row < kPageRows; ++row) { + for (int g = 0; g < 4; ++g) { + const float sh = half_to_float( + src_scales[(static_cast(head) * kPageRows + row) * 4 + g]); + float vals[64]; + for (int i = 0; i < 64; ++i) { + vals[i] = static_cast( + src_codes[static_cast(head) * kI8CodeB + + row * kHeadDim + g * 64 + i]) * sh; + } + // requant g64 (matches entropy_cold_requant oracle) + float amax = 0.0f; + for (int i = 0; i < 64; ++i) { amax = std::fmax(amax, std::fabs(vals[i])); } + const std::uint8_t sb = e4m3_rne(std::fmax(amax / 6.0f, 0x1p-9f)); + const float s = e4m3_to_f32(sb); + float dec[64]; + for (int i = 0; i < 64; ++i) { + dec[i] = e2m1_to_f32(e2m1_code(vals[i] / s)) * s; + } + // upper-bound int8 re-encode (mirror device) + float mx = 0.0f; + for (int sub = 0; sub < 4; ++sub) { + // g64 group covers scale slots [g*4, g*4+4) of the row's + // g16 layout; requant wrote the same sb replicated. + mx = std::fmax(mx, e4m3_to_f32(sb)); + } + const float scale = mx * 6.0f / 127.0f; + const std::uint16_t want_scale = float_to_half(scale); + const std::uint16_t got_scale = + got_scales[(static_cast(head) * kPageRows + row) * 4 + g]; + if (want_scale != got_scale) { ++scale_mismatch; } + const float inv = scale > 0.0f ? 1.0f / scale : 0.0f; + const float got_sh = half_to_float(got_scale); + for (int i = 0; i < 64; ++i) { + int c = static_cast(std::nearbyint(dec[i] * inv)); + c = std::max(-127, std::min(127, c)); + const std::int8_t got = got_codes[static_cast(head) * kI8CodeB + + row * kHeadDim + g * 64 + i]; + if (got != static_cast(c)) { ++byte_mismatch; } + const float restored = static_cast(got) * got_sh; + num += static_cast(restored - vals[i]) * (restored - vals[i]); + den += static_cast(vals[i]) * vals[i]; + } + } + } + } + const double nmse = den > 0.0 ? num / den : 0.0; + check(byte_mismatch == 0, "restore codes match oracle"); + check(scale_mismatch == 0, "restore scales match oracle"); + check(nmse < 0.05, "roundtrip NMSE bound"); // synthetic spiky data; real planes measured 0.012 + std::printf("cold_i8 roundtrip: byte_mismatch=%d scale_mismatch=%d NMSE=%.5f\n", + byte_mismatch, scale_mismatch, nmse); + + if (failures == 0) { std::printf("cold_i8: all tests passed\n"); } + return failures == 0 ? 0 : 1; +} diff --git a/tests/ops/test_entropy_cold_requant.cpp b/tests/ops/test_entropy_cold_requant.cpp new file mode 100644 index 0000000000..faf8ddd9fd --- /dev/null +++ b/tests/ops/test_entropy_cold_requant.cpp @@ -0,0 +1,432 @@ +#include "ninfer/ops/entropy_cold_requant.h" +#include "ops/op_tester.h" + +#include +#include +#include +#include +#include +#include +#include + +using namespace ninfer; +using namespace ninfer::test; + +namespace { + +constexpr int kHeadDim = 256; +constexpr int kPageRows = 64; +constexpr int kKvHeads = 4; +constexpr int kNvfp4CodeB = kHeadDim / 2 * kPageRows; // 8192 per head-page +constexpr int kNvfp4ScaleB = kHeadDim / 16 * kPageRows; // 1024 +constexpr int kInt8CodeB = kHeadDim * kPageRows; // 16384 +constexpr int kInt8ScaleB = kHeadDim / 64 * 2 * kPageRows; // 512 (fp16) + +std::uint8_t e4m3_rne(float x) { + if (!(x > 0.0f)) { return 0; } + std::uint32_t bits; + std::memcpy(&bits, &x, 4); + const std::uint32_t sign = (bits >> 24) & 0x80u; + int exponent = static_cast((bits >> 23) & 0xffu) - 127 + 7; + if (exponent >= 15) { return static_cast(sign | (15u << 3) | 7u); } + if (exponent <= 0) { + int mantissa = static_cast(std::nearbyint(x * 512.0f)); + if (mantissa <= 0) { return static_cast(sign); } + if (mantissa >= 8) { return static_cast(sign | (1u << 3)); } + return static_cast(sign | mantissa); + } + std::uint32_t mantissa = (bits >> 20) & 0x7u; + const std::uint32_t guard = (bits >> 19) & 1u; + const std::uint32_t sticky = bits & 0x7ffffu; + if (guard && (sticky || (mantissa & 1u))) { + mantissa += 1; + if (mantissa > 7) { + mantissa = 0; + exponent += 1; + if (exponent >= 15) { return static_cast(sign | (15u << 3) | 7u); } + } + } + return static_cast(sign | (exponent << 3) | mantissa); +} + +float e4m3_to_f32(std::uint8_t byte) { + const int e = (byte >> 3) & 0xF; + const int m = byte & 0x7; + if (e == 0) { return static_cast(m) / 512.0f; } + return (1.0f + static_cast(m) / 8.0f) * std::pow(2.0f, static_cast(e - 7)); +} + +float e2m1_to_f32(std::uint8_t code) { + static const float mag[8] = {0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f}; + const float v = mag[code & 0x7]; + return (code & 0x8) != 0 ? -v : v; +} + +std::uint8_t e2m1_code(float x) { + const float a = std::fabs(x); + std::uint8_t c; + if (a < 0.25f) { c = 0; } + else if (a < 0.75f) { c = 1; } + else if (a < 1.25f) { c = 2; } + else if (a < 1.75f) { c = 3; } + else if (a < 2.5f) { c = 4; } + else if (a < 3.5f) { c = 5; } + else if (a < 5.0f) { c = 6; } + else { c = 7; } + if (x < 0.0f) { c |= 0x08u; } + return c; +} + +float iso3_to_f32(std::uint8_t code) { + const float mag = static_cast(code & 0x7); + return (code & 0x8) != 0 ? -mag : mag; +} + +std::uint8_t iso3_code(float value, float scale) { + float mag = std::roundf(std::fabs(value) / scale); // device uses roundf + if (mag > 7.0f) { mag = 7.0f; } + if (mag < 0.0f) { mag = 0.0f; } + std::uint8_t code = static_cast(mag); + if (value < 0.0f && code != 0) { code |= 0x08u; } + return code; +} + +float half_to_float(std::uint16_t h) { + const std::uint32_t sign = (h >> 15) & 1u; + const std::uint32_t exp = (h >> 10) & 0x1Fu; + const std::uint32_t man = h & 0x3FFu; + float out; + if (exp == 0) { + out = std::ldexp(static_cast(man), -24); + } else { + out = std::ldexp(1024.0f + static_cast(man), static_cast(exp) - 25); + } + return sign != 0 ? -out : out; +} + +std::uint16_t float_to_half(float f) { + if (f <= 0.0f) { return 0; } + std::uint32_t x; + std::memcpy(&x, &f, 4); + const std::uint32_t sign = (x >> 16) & 0x8000u; + int exponent = static_cast((x >> 23) & 0xFFu) - 127 + 15; + std::uint32_t mantissa = (x >> 13) & 0x3FFu; + if (exponent <= 0) { + const std::uint32_t man_full = (x & 0x7FFFFFu) | 0x800000u; + const int shift = 14 - exponent + 1; + mantissa = man_full >> shift; + const std::uint32_t round_bit = (man_full >> (shift - 1)) & 1u; + if (round_bit != 0) { mantissa += 1; } + exponent = 0; + } else { + const std::uint32_t round_bit = (x >> 12) & 1u; + const std::uint32_t sticky = x & 0xFFFu; + if (round_bit != 0 && (sticky != 0 || (mantissa & 1u) != 0)) { + mantissa += 1; + if (mantissa > 0x3FFu) { + mantissa = 0; + exponent += 1; + } + } + } + if (exponent >= 31) { return static_cast(sign | (31u << 10)); } + return static_cast(sign | (static_cast(exponent) << 10) | + mantissa); +} + +struct PageValues { + std::vector values; + std::vector nvfp4_codes; + std::vector nvfp4_scales; + std::vector iso3_codes; + std::vector iso3_scales; + std::vector int8_codes; + std::vector int8_scales; +}; + +PageValues make_page(std::mt19937& rng) { + PageValues page; + page.values.resize(static_cast(kKvHeads) * kPageRows * kHeadDim); + std::normal_distribution noise(0.0f, 1.0f); + std::uniform_real_distribution level(0.001f, 8.0f); + for (int head = 0; head < kKvHeads; ++head) { + for (int row = 0; row < kPageRows; ++row) { + const float amp = level(rng); + for (int d = 0; d < kHeadDim; ++d) { + float v = amp * noise(rng); + if (((head * 131 + row * 17 + d) % 4096) == 0) { v *= 64.0f; } + if (head == kKvHeads - 1 && row < 2) { v = 0.0f; } + page.values[(static_cast(head) * kPageRows + row) * kHeadDim + d] = v; + } + } + } + + page.nvfp4_codes.assign(static_cast(kKvHeads) * kNvfp4CodeB, 0); + page.nvfp4_scales.assign(static_cast(kKvHeads) * kNvfp4ScaleB, 0); + page.iso3_codes.assign(static_cast(kKvHeads) * kNvfp4CodeB, 0); + page.iso3_scales.assign(static_cast(kKvHeads) * kNvfp4ScaleB, 0); + page.int8_codes.assign(static_cast(kKvHeads) * kInt8CodeB, 0); + page.int8_scales.assign(static_cast(kKvHeads) * kPageRows * (kHeadDim / 64), 0); + + for (int head = 0; head < kKvHeads; ++head) { + for (int row = 0; row < kPageRows; ++row) { + const std::size_t vrow = + (static_cast(head) * kPageRows + row) * kHeadDim; + for (int g = 0; g < kHeadDim / 16; ++g) { + float amax = 0.0f; + for (int i = 0; i < 16; ++i) { + amax = std::fmax(amax, std::fabs(page.values[vrow + g * 16 + i])); + } + const std::uint8_t kb = e4m3_rne(std::fmax(amax / 6.0f, 0x1p-9f)); + page.nvfp4_scales[static_cast(head) * kNvfp4ScaleB + + row * (kHeadDim / 16) + g] = kb; + const float ks = e4m3_to_f32(kb); + const std::uint8_t vb = e4m3_rne(std::fmax(amax / 7.0f, 0x1p-9f)); + page.iso3_scales[static_cast(head) * kNvfp4ScaleB + + row * (kHeadDim / 16) + g] = vb; + const float vs = e4m3_to_f32(vb); + for (int i = 0; i < 16; i += 2) { + page.nvfp4_codes[static_cast(head) * kNvfp4CodeB + + row * (kHeadDim / 2) + g * 8 + i / 2] = + static_cast( + e2m1_code(page.values[vrow + g * 16 + i] / ks) | + (e2m1_code(page.values[vrow + g * 16 + i + 1] / ks) << 4)); + page.iso3_codes[static_cast(head) * kNvfp4CodeB + + row * (kHeadDim / 2) + g * 8 + i / 2] = + static_cast( + iso3_code(page.values[vrow + g * 16 + i], vs) | + (iso3_code(page.values[vrow + g * 16 + i + 1], vs) << 4)); + } + } + for (int g = 0; g < kHeadDim / 64; ++g) { + float amax = 0.0f; + for (int i = 0; i < 64; ++i) { + amax = std::fmax(amax, std::fabs(page.values[vrow + g * 64 + i])); + } + const float s = std::fmax(amax / 127.0f, 1e-30f); + const std::uint16_t sb = float_to_half(s); + page.int8_scales[(static_cast(head) * kPageRows + row) * + (kHeadDim / 64) + g] = sb; + const float sh = half_to_float(sb); + for (int i = 0; i < 64; ++i) { + float q = std::nearbyint(page.values[vrow + g * 64 + i] / sh); + if (q > 127.0f) { q = 127.0f; } + if (q < -127.0f) { q = -127.0f; } + page.int8_codes[static_cast(head) * kInt8CodeB + + row * kHeadDim + g * 64 + i] = static_cast(q); + } + } + } + } + return page; +} + +void requant_oracle(const PageValues& page, int mode, std::vector& out_codes, + std::vector& out_scales) { + // mode: 0 = nvfp4 K source, 1 = int8 source, 2 = iso3 V source + const bool iso3_out = mode == 2; + const float scale_div = iso3_out ? 7.0f : 6.0f; + out_codes.assign(static_cast(kKvHeads) * kNvfp4CodeB, 0); + out_scales.assign(static_cast(kKvHeads) * kNvfp4ScaleB, 0); + std::vector dec(static_cast(kKvHeads) * kPageRows * kHeadDim); + for (int head = 0; head < kKvHeads; ++head) { + for (int row = 0; row < kPageRows; ++row) { + for (int g64 = 0; g64 < kHeadDim / 64; ++g64) { + if (mode == 1) { + const std::uint16_t sb = + page.int8_scales[(static_cast(head) * kPageRows + row) * + (kHeadDim / 64) + g64]; + const float s = half_to_float(sb); + for (int i = 0; i < 64; ++i) { + const std::int8_t c = + page.int8_codes[static_cast(head) * kInt8CodeB + + row * kHeadDim + g64 * 64 + i]; + dec[(static_cast(head) * kPageRows + row) * kHeadDim + + g64 * 64 + i] = static_cast(c) * s; + } + } else { + for (int i = 0; i < 64; ++i) { + const int d = g64 * 64 + i; + const std::uint8_t byte = + (mode == 0 ? page.nvfp4_codes : page.iso3_codes) + [static_cast(head) * kNvfp4CodeB + + row * (kHeadDim / 2) + (d >> 1)]; + const std::uint8_t nib = + (d & 1) != 0 ? static_cast(byte >> 4) + : static_cast(byte & 0x0F); + const float s = e4m3_to_f32( + (mode == 0 ? page.nvfp4_scales : page.iso3_scales) + [static_cast(head) * kNvfp4ScaleB + + row * (kHeadDim / 16) + (d >> 4)]); + dec[(static_cast(head) * kPageRows + row) * kHeadDim + d] = + (mode == 0 ? e2m1_to_f32(nib) : iso3_to_f32(nib)) * s; + } + } + } + } + } + for (int head = 0; head < kKvHeads; ++head) { + for (int row = 0; row < kPageRows; ++row) { + for (int g64 = 0; g64 < kHeadDim / 64; ++g64) { + float amax = 0.0f; + for (int i = 0; i < 64; ++i) { + amax = std::fmax( + amax, + std::fabs(dec[(static_cast(head) * kPageRows + row) * + kHeadDim + g64 * 64 + i])); + } + const std::uint8_t sb = e4m3_rne(std::fmax(amax / scale_div, 0x1p-9f)); + const float s = e4m3_to_f32(sb); + for (int i = 0; i < 64; i += 2) { + const int d0 = g64 * 64 + i; + const float v0 = + dec[(static_cast(head) * kPageRows + row) * kHeadDim + d0]; + const float v1 = dec[(static_cast(head) * kPageRows + row) * + kHeadDim + d0 + 1]; + const std::uint8_t lo = iso3_out ? iso3_code(v0, s) : e2m1_code(v0 / s); + const std::uint8_t hi = iso3_out ? iso3_code(v1, s) : e2m1_code(v1 / s); + out_codes[static_cast(head) * kNvfp4CodeB + + row * (kHeadDim / 2) + (d0 >> 1)] = + static_cast(lo | (hi << 4)); + } + for (int r = 0; r < 4; ++r) { + out_scales[static_cast(head) * kNvfp4ScaleB + + row * (kHeadDim / 16) + g64 * 4 + r] = sb; + } + } + } + } +} + +int check(bool ok, const char* what, int& failures) { + if (!ok) { + std::printf("FAIL: %s\n", what); + ++failures; + } + return failures; +} + +} // namespace + +int main() { + std::mt19937 rng(20260830); + int failures = 0; + const char* tags[3] = {"nvfp4-g16 K", "int8-g64", "iso3-g16 V"}; + + for (int mode = 0; mode < 3; ++mode) { + PageValues page = make_page(rng); + + std::vector stored(static_cast(kKvHeads) * kPageRows * kHeadDim); + for (int head = 0; head < kKvHeads; ++head) { + for (int row = 0; row < kPageRows; ++row) { + if (mode == 1) { + for (int g = 0; g < kHeadDim / 64; ++g) { + const float s = half_to_float( + page.int8_scales[(static_cast(head) * kPageRows + row) * + (kHeadDim / 64) + g]); + for (int i = 0; i < 64; ++i) { + stored[(static_cast(head) * kPageRows + row) * kHeadDim + + g * 64 + i] = + static_cast( + page.int8_codes[static_cast(head) * kInt8CodeB + + row * kHeadDim + g * 64 + i]) * s; + } + } + continue; + } + for (int g = 0; g < kHeadDim / 16; ++g) { + const std::vector& codes = + mode == 0 ? page.nvfp4_codes : page.iso3_codes; + const std::vector& scales = + mode == 0 ? page.nvfp4_scales : page.iso3_scales; + const float s = e4m3_to_f32( + scales[static_cast(head) * kNvfp4ScaleB + + row * (kHeadDim / 16) + g]); + for (int i = 0; i < 16; ++i) { + const int d = g * 16 + i; + const std::uint8_t byte = + codes[static_cast(head) * kNvfp4CodeB + + row * (kHeadDim / 2) + (d >> 1)]; + const std::uint8_t nib = + (d & 1) != 0 ? static_cast(byte >> 4) + : static_cast(byte & 0x0F); + stored[(static_cast(head) * kPageRows + row) * kHeadDim + d] = + (mode == 0 ? e2m1_to_f32(nib) : iso3_to_f32(nib)) * s; + } + } + } + } + + std::vector expect_codes, expect_scales; + requant_oracle(page, mode, expect_codes, expect_scales); + + const bool int8_source = mode == 1; + GuardedDeviceBuffer dsrc_codes(int8_source + ? static_cast(kKvHeads) * kInt8CodeB + : static_cast(kKvHeads) * kNvfp4CodeB); + GuardedDeviceBuffer dsrc_scales(int8_source + ? static_cast(kKvHeads) * kInt8ScaleB + : static_cast(kKvHeads) * kNvfp4ScaleB); + GuardedDeviceBuffer ddst_codes(static_cast(kKvHeads) * kNvfp4CodeB); + GuardedDeviceBuffer ddst_scales(static_cast(kKvHeads) * kNvfp4ScaleB); + + if (int8_source) { + dsrc_codes.copy_from_host(page.int8_codes.data(), dsrc_codes.bytes()); + dsrc_scales.copy_from_host(page.int8_scales.data(), dsrc_scales.bytes()); + } else if (mode == 0) { + dsrc_codes.copy_from_host(page.nvfp4_codes.data(), dsrc_codes.bytes()); + dsrc_scales.copy_from_host(page.nvfp4_scales.data(), dsrc_scales.bytes()); + } else { + dsrc_codes.copy_from_host(page.iso3_codes.data(), dsrc_codes.bytes()); + dsrc_scales.copy_from_host(page.iso3_scales.data(), dsrc_scales.bytes()); + } + + ops::entropy_cold_requant_raw( + static_cast(dsrc_codes.data()), + static_cast(dsrc_scales.data()), + mode == 0 ? ops::EntropyColdRequantMode::Nvfp4G16 + : mode == 1 ? ops::EntropyColdRequantMode::Int8G64 + : ops::EntropyColdRequantMode::Iso3VG16, + kKvHeads, 1, static_cast(ddst_codes.data()), + static_cast(ddst_scales.data()), nullptr); + cuda_synchronize(); + + std::vector got_codes(static_cast(kKvHeads) * kNvfp4CodeB); + std::vector got_scales(static_cast(kKvHeads) * kNvfp4ScaleB); + ddst_codes.copy_to_host(got_codes.data(), ddst_codes.bytes()); + ddst_scales.copy_to_host(got_scales.data(), ddst_scales.bytes()); + + const char* tag = tags[mode]; + const bool codes_ok = got_codes == expect_codes; + const bool scales_ok = got_scales == expect_scales; + check(codes_ok, (std::string(tag) + " requant codes match oracle").c_str(), failures); + check(scales_ok, (std::string(tag) + " requant scales match oracle").c_str(), failures); + + double num = 0.0, den = 0.0; + const bool iso3_out = mode == 2; + for (std::size_t i = 0; i < stored.size(); ++i) { + const int head = static_cast(i / (static_cast(kPageRows) * kHeadDim)); + const int row = static_cast((i / kHeadDim) % kPageRows); + const int d = static_cast(i % kHeadDim); + const std::uint8_t byte = + got_codes[static_cast(head) * kNvfp4CodeB + row * (kHeadDim / 2) + + (d >> 1)]; + const std::uint8_t nib = + (d & 1) != 0 ? static_cast(byte >> 4) + : static_cast(byte & 0x0F); + const float s = e4m3_to_f32( + got_scales[static_cast(head) * kNvfp4ScaleB + row * (kHeadDim / 16) + + (d >> 4)]); + const float out = iso3_out ? iso3_to_f32(nib) * s : e2m1_to_f32(nib) * s; + num += static_cast(out - stored[i]) * (out - stored[i]); + den += static_cast(stored[i]) * stored[i]; + } + const double nmse = den > 0.0 ? num / den : 0.0; + check(nmse < 0.03, (std::string(tag) + " requant NMSE bound").c_str(), failures); + std::printf("[%s] requant NMSE vs stored = %.5f (codes_ok=%d scales_ok=%d)\n", tag, nmse, + codes_ok ? 1 : 0, scales_ok ? 1 : 0); + } + + if (failures == 0) { std::printf("entropy_cold_requant: all tests passed\n"); } + return failures == 0 ? 0 : 1; +} diff --git a/tools/calib/analyze_kv.py b/tools/calib/analyze_kv.py new file mode 100644 index 0000000000..97ea5a9238 --- /dev/null +++ b/tools/calib/analyze_kv.py @@ -0,0 +1,332 @@ +#!/usr/bin/env python +# Offline KV dynamic-precision calibration for NInfer Qwen3.6-family artifacts. +# +# Consumes the .kvc frames produced by `ninfer-cli --kv-calib-dir DIR` (exact +# post-RoPE K and V per full-attention layer and prefill chunk) and produces a +# static per-layer dtype table in the format accepted by `--kv-layer-storage`. +# +# Implemented metrics (the four requested techniques, fused into one decision): +# MixKVQ K/V error asymmetry weighting (K weighted above V). +# TriAxialKV per-layer, per-head, per-dimension-group outlier scores. +# ARKV effective rank (spectral spread) of K and V per head. +# KVTuner greedy sensitivity ranking under an explicit memory budget. +# +# The decision space is per-layer same-dtype storage (bf16 | int8 | nvfp4); +# the runtime's layer_kv_dtypes table consumes the selected map. +import argparse +import json +import math +import struct +import sys +from pathlib import Path + +import numpy as np + +HEADER = struct.Struct("<16s6I2i4I") +MAGIC = b"NINFERKVCAL1\x00\x00\x00\x00" + +E2M1_VALUES = np.array([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=np.float64) +E2M1_EDGES = np.array([0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5.0], dtype=np.float64) + + +def e4m3fn(x: float) -> float: + """Mimic the runtime gqa_kv_nvfp4_fp32_to_e4m3 (positive scale path).""" + if not (x > 0.0): + return 0.0 + b = struct.unpack("> 23) & 0xFF) - 127 + 7 + if exp >= 15: + return 448.0 / 512.0 * 2**7 # saturate (0b01111111 = 448) + if exp <= 0: + mant = int(round(x * 64.0)) + if mant <= 0: + return 0.0 + if mant >= 8: + return 1.0 + return mant / 512.0 + mant = (b >> 20) & 0x7 + guard = (b >> 19) & 1 + sticky = b & 0x7FFFF + if guard and (sticky or (mant & 1)): + mant += 1 + if mant > 7: + mant = 0 + exp += 1 + if exp >= 15: + return 448.0 / 512.0 * 2**7 + return (1.0 + mant / 8.0) * 2 ** (exp - 7) + + +def quantize_e2m1_group16(x: np.ndarray) -> np.ndarray: + x = np.nan_to_num(x.astype(np.float64), nan=0.0, posinf=0.0, neginf=0.0) + # groups along the last axis (head dim), 16 values each. + groups = x.reshape(*x.shape[:-1], -1, 16) + amax = np.abs(groups).max(axis=-1, keepdims=True) + scale = np.maximum(amax / 6.0, 2.0**-9) + vq = np.array([e4m3fn(float(v)) for v in scale.ravel()], dtype=np.float64).reshape(scale.shape) + q = np.abs(groups) / vq + codes = np.searchsorted(E2M1_EDGES, q).astype(np.int64) + decoded = np.where(groups < 0, -1.0, 1.0) * E2M1_VALUES[codes] * vq + return decoded.reshape(x.shape) + + +def quantize_int8_group64(x: np.ndarray) -> np.ndarray: + x = np.nan_to_num(x.astype(np.float64), nan=0.0, posinf=0.0, neginf=0.0) + tail = x.shape[-1] % 64 + if tail: + pad = 64 - tail + x = np.concatenate([x, np.zeros((*x.shape[:-1], pad), dtype=np.float64)], axis=-1) + groups = x.reshape(*x.shape[:-1], -1, 64) + amax = np.abs(groups).max(axis=-1, keepdims=True) + scale = np.maximum(amax / 127.0, 1e-30) + scale = scale.astype(np.float16).astype(np.float64) + q = np.clip(np.round(groups / scale), -127, 127) + decoded = q * scale + if tail: + decoded = decoded[..., :tail] + return decoded.reshape(x.shape) + + +def quantize_fp8_group16(x: np.ndarray) -> np.ndarray: + x = np.nan_to_num(x.astype(np.float64), nan=0.0, posinf=0.0, neginf=0.0) + groups = x.reshape(*x.shape[:-1], -1, 16) + amax = np.abs(groups).max(axis=-1, keepdims=True) + scale = np.maximum(amax / 448.0, 2.0**-9) + q = np.clip(np.round(groups / scale), -448, 448) + decoded = q * scale + return decoded.reshape(x.shape) + + +def quantize_iso3_group16(x: np.ndarray) -> np.ndarray: + x = np.nan_to_num(x.astype(np.float64), nan=0.0, posinf=0.0, neginf=0.0) + groups = x.reshape(*x.shape[:-1], -1, 16) + amax = np.abs(groups).max(axis=-1, keepdims=True) + scale = np.maximum(amax / 7.0, 2.0**-9) + q = np.clip(np.round(groups / scale), -7, 7) + decoded = q * scale + return decoded.reshape(x.shape) + + +def nmse(x: np.ndarray, y: np.ndarray) -> float: + x = np.nan_to_num(x.astype(np.float64), nan=0.0, posinf=0.0, neginf=0.0) + y = np.nan_to_num(y.astype(np.float64), nan=0.0, posinf=0.0, neginf=0.0) + num = np.sum((x - y) ** 2) + den = np.sum(x**2) + return float(num / den) if den > 0 else 0.0 + + +def effective_rank(x: np.ndarray) -> float: + x = np.nan_to_num(x.astype(np.float64), nan=0.0, posinf=0.0, neginf=0.0) + if x.shape[0] > x.shape[1]: + eig = np.linalg.eigvalsh(x.T @ x) + else: + eig = np.linalg.eigvalsh(x @ x.T) + s = np.sqrt(np.maximum(eig, 0.0)) + if s.size == 0 or s[0] <= 0: + return 0.0 + p = s / s[0] + den = np.sum(p**2) + return float(np.sum(p) ** 2 / den) if den > 0 else 0.0 + + +def outlier_scores(x: np.ndarray) -> tuple[float, float]: + x = np.nan_to_num(np.asarray(x, dtype=np.float64), nan=0.0, posinf=0.0, neginf=0.0) + rms = math.sqrt(float(np.mean(x**2))) + if rms == 0: + return 0.0, 0.0 + head = float(np.mean((np.abs(x).max(axis=-1) > 6.0 * rms).astype(np.float64))) + dim_group = float( + np.mean( + ( + np.abs(x).max(axis=-2) + > 6.0 * rms * math.sqrt(x.shape[-2]) + ).astype(np.float64) + ) + ) + return head, dim_group + + +def load_frames(directory: Path): + frames = [] + for path in sorted(directory.glob("*.kvc")): + raw = path.read_bytes() + if len(raw) < HEADER.size: + raise SystemExit(f"truncated record: {path}") + magic, header_bytes, layer, head_dim, kv_heads, tokens, record_index, first_pos, last_pos, *_ = ( + HEADER.unpack(raw[: HEADER.size]) + ) + if magic != MAGIC or header_bytes != HEADER.size: + raise SystemExit(f"bad record header: {path}") + payload = np.frombuffer(raw, dtype=np.uint8, offset=HEADER.size) + expect = (tokens * 4) + 2 * head_dim * kv_heads * tokens * 2 + if payload.size != expect: + raise SystemExit(f"bad record payload: {path}") + pos = payload[: tokens * 4].copy().view(np.int32) + arr = np.frombuffer(payload[tokens * 4 :].tobytes(), dtype=" 16: + raise SystemExit(f"too many layers for the runtime table: {n_layers}") + + # Layer statistics accumulated over frames with token weighting. + stats = { + layer: { + "tokens": 0, + "k_nmse": {"nvfp4": 0.0, "int8": 0.0, "fp8": 0.0, "iso3": 0.0}, + "v_nmse": {"nvfp4": 0.0, "int8": 0.0, "fp8": 0.0, "iso3": 0.0}, + "k_rank": 0.0, + "v_rank": 0.0, + "k_outlier_head": 0.0, + "v_outlier_head": 0.0, + "k_outlier_group": 0.0, + "v_outlier_group": 0.0, + } + for layer in layers + } + + for frame in frames: + layer = frame["layer"] + weight = frame["tokens"] + stat = stats[layer] + stat["tokens"] += weight + k = frame["k"] + v = frame["v"] + stat["k_nmse"]["nvfp4"] += nmse(k, quantize_e2m1_group16(k)) * weight + stat["k_nmse"]["int8"] += nmse(k, quantize_int8_group64(k)) * weight + stat["k_nmse"]["fp8"] += nmse(k, quantize_fp8_group16(k)) * weight + stat["k_nmse"]["iso3"] += nmse(k, quantize_iso3_group16(k)) * weight + stat["v_nmse"]["nvfp4"] += nmse(v, quantize_e2m1_group16(v)) * weight + stat["v_nmse"]["int8"] += nmse(v, quantize_int8_group64(v)) * weight + stat["v_nmse"]["fp8"] += nmse(v, quantize_fp8_group16(v)) * weight + stat["v_nmse"]["iso3"] += nmse(v, quantize_iso3_group16(v)) * weight + for h in range(k.shape[1]): + stat["k_rank"] += effective_rank(k[:, h, :]) * weight + stat["v_rank"] += effective_rank(v[:, h, :]) * weight + ok_head, ok_group = outlier_scores(k) + ov_head, ov_group = outlier_scores(v) + stat["k_outlier_head"] += ok_head * weight + stat["v_outlier_head"] += ov_head * weight + stat["k_outlier_group"] += ok_group * weight + stat["v_outlier_group"] += ov_group * weight + + rows = [] + for layer in layers: + stat = stats[layer] + w = stat["tokens"] + k_err = stat["k_nmse"]["nvfp4"] / w + v_err = stat["v_nmse"]["nvfp4"] / w + mix_err = args.k_weight * k_err + (1.0 - args.k_weight) * v_err + triaxial = max( + stat["k_outlier_head"], stat["v_outlier_head"], + stat["k_outlier_group"], stat["v_outlier_group"], + ) / max(w, 1) + rank = 0.5 * (stat["k_rank"] / w + stat["v_rank"] / w) + # Normalized 0..1 scores across layers for greedy ranking. + rows.append( + { + "layer": layer, + "tokens": w, + "nmse_nvfp4_k": k_err, + "nmse_nvfp4_v": v_err, + "nmse_int8_k": stat["k_nmse"]["int8"] / w, + "nmse_int8_v": stat["v_nmse"]["int8"] / w, + "nmse_fp8_k": stat["k_nmse"]["fp8"] / w, + "nmse_fp8_v": stat["v_nmse"]["fp8"] / w, + "nmse_iso3_k": stat["k_nmse"]["iso3"] / w, + "nmse_iso3_v": stat["v_nmse"]["iso3"] / w, + "mixkvq_error": mix_err, + "triaxial_outlier": triaxial, + "arkv_rank": rank, + } + ) + for key in ("mixkvq_error", "triaxial_outlier", "arkv_rank"): + lo = min(row[key] for row in rows) + hi = max(row[key] for row in rows) + for row in rows: + row[f"{key}_norm"] = 0.0 if hi <= lo else (row[key] - lo) / (hi - lo) + for row in rows: + row["sensitivity"] = ( + 0.50 * row["mixkvq_error_norm"] + + 0.25 * row["triaxial_outlier_norm"] + + 0.25 * row["arkv_rank_norm"] + ) + + # Storage cost per token per layer (K+V), in bytes; nvfp4 is the budget base. + head_dim = frames[0]["k"].shape[-1] + kv_heads = frames[0]["k"].shape[1] + def cost(dtype: str) -> float: + if dtype == "nvfp4" or dtype == "iso3": + return kv_heads * (head_dim + 2 * (head_dim / 16)) + if dtype == "int8": + return kv_heads * (2 * head_dim + 4 * (head_dim / 64)) + if dtype == "fp8": + return kv_heads * (2 * head_dim + 2 * (head_dim / 16)) + return kv_heads * 4 * head_dim + + base = n_layers * cost("nvfp4") + ranked = sorted(rows, key=lambda row: row["sensitivity"], reverse=True) + table = {layer: "nvfp4" for layer in layers} + used = base + for row in ranked: + for dtype in ("fp8", "int8", "bf16"): + delta = cost(dtype) - cost(table[row["layer"]]) + if used + delta <= args.budget * base: + table[row["layer"]] = dtype + used += delta + break + + spec = ",".join(f"{layer}:{table[layer]}" for layer in layers) + report = { + "record_frames": len(frames), + "layer_count": n_layers, + "head_dim": head_dim, + "kv_heads": kv_heads, + "budget_factor": args.budget, + "k_weight": args.k_weight, + "per_layer": sorted(rows, key=lambda row: row["layer"]), + "selected_table": table, + "kv_layer_storage_spec": spec, + "relative_cost": used / base, + } + Path(args.out).write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8") + print(f"frames={len(frames)} layers={n_layers} budget={args.budget:.2f}x -> " + f"cost={used / base:.3f}x") + print("--kv-layer-storage " + spec) + + +if __name__ == "__main__": + main() diff --git a/tools/calib/pca_kv_feasibility.py b/tools/calib/pca_kv_feasibility.py new file mode 100644 index 0000000000..12ca9b32d0 --- /dev/null +++ b/tools/calib/pca_kv_feasibility.py @@ -0,0 +1,139 @@ +#!/usr/bin/env python3 +# PCA + entropy feasibility for NVFP4 K/V code nibbles. +# +# Reproduces the production post-rotation code distribution, then tests +# per-16-channel PCA rotations and reports rANS stream sizes for the fixed +# slot codec (16 streams of 512 nibbles per half, shared per-half frequencies). +import argparse +import glob +import re +import struct +import sys +from pathlib import Path + +import numpy as np + +sys_path = Path(__file__).resolve().parents[2] / "tools" / "calib" +sys.path.insert(0, str(sys_path)) +import rans_nvfp4 as rans # noqa: E402 + + +def load_kv(path): + raw = Path(path).read_bytes() + hdr = struct.unpack("<16s6I2i4I", raw[:64]) + tokens = hdr[5] + arr = np.frombuffer(raw, dtype="thj", so4[block], x[:, :, base : base + 4]) + return out + + def apply_basis(x, basis, block): + out = np.empty_like(x) + for b in range(x.shape[-1] // block): + base = b * block + out[:, :, base : base + block] = np.einsum( + "jk,thk->thj", basis, x[:, :, base : base + block]) + return out + + for name, x in [("K", k_all), ("V", v_all)]: + so4_codes, _ = quantize_codes(apply_so4(x)) + so4_sizes = slot_sizes(so4_codes) + print(f"{name} SO(4): max={max(so4_sizes[0] + so4_sizes[1])} " + f"h0={max(so4_sizes[0])} h1={max(so4_sizes[1])} " + f"over166={max(so4_sizes[0] + so4_sizes[1]) > 166}") + + for block in (16, 64): + basis, _ = pca_basis(apply_so4(x), block) + pca_codes, _ = quantize_codes(apply_basis(apply_so4(x), basis, block)) + pca_sizes = slot_sizes(pca_codes) + print(f"{name} SO(4)+PCA{block}: max={max(pca_sizes[0] + pca_sizes[1])} " + f"h0={max(pca_sizes[0])} h1={max(pca_sizes[1])} " + f"over166={max(pca_sizes[0] + pca_sizes[1]) > 166}") + + raw_basis, _ = pca_basis(x, block) + raw_codes, _ = quantize_codes(apply_basis(x, raw_basis, block)) + raw_sizes = slot_sizes(raw_codes) + print(f"{name} PCA{block}: max={max(raw_sizes[0] + raw_sizes[1])} " + f"h0={max(raw_sizes[0])} h1={max(raw_sizes[1])} " + f"over166={max(raw_sizes[0] + raw_sizes[1]) > 166}") + + +if __name__ == "__main__": + main() diff --git a/tools/calib/rans_nvfp4.py b/tools/calib/rans_nvfp4.py new file mode 100644 index 0000000000..bae7074837 --- /dev/null +++ b/tools/calib/rans_nvfp4.py @@ -0,0 +1,112 @@ +#!/usr/bin/env python +# CPU reference order-0 static rANS for NVFP4 E2M1 code nibbles. +# Feasibility gate for the entropy-coded cold KV pool. +import argparse +import glob +import struct +from pathlib import Path + +import numpy as np + +MAGIC = b"NINFERKVCAL1\x00\x00\x00\x00" +SCALE_BITS = 12 +SCALE = 1 << SCALE_BITS +MASK = SCALE - 1 +BYTE_L = 1 << 23 + + +def build_freqs(symbols): + counts = np.bincount(symbols, minlength=16).astype(np.int64) + freqs = np.maximum(1, np.round(counts / counts.sum() * SCALE).astype(np.int64)) + diff = int(SCALE - freqs.sum()) + while diff > 0: + idx = int(np.argmax(counts)) + freqs[idx] += 1 + counts[idx] = max(0, counts[idx] - 1) + diff -= 1 + while diff < 0: + idx = int(np.argmax(np.where(freqs > 1, freqs, 0))) + freqs[idx] -= 1 + diff += 1 + return freqs + + +def encode(symbols, freqs): + start = np.concatenate(([0], np.cumsum(freqs)[:-1])).astype(np.int64) + out = bytearray() + x = BYTE_L + for s in reversed(symbols): + f = int(freqs[s]) + x_max = ((BYTE_L >> SCALE_BITS) << 8) * f + while x >= x_max: + out.append(x & 0xFF) + x >>= 8 + x = ((x // f) << SCALE_BITS) + (x % f) + int(start[s]) + out.extend(x.to_bytes(4, "little")) + return bytes(out), start + + +def decode(data, freqs, start, count): + x = int.from_bytes(data[-4:], "little") + inb = bytearray(data[:-4]) + out = [] + for _ in range(count): + slot = x & MASK + s = int(np.searchsorted(start, slot, side="right") - 1) + out.append(s) + x = int(freqs[s]) * (x >> SCALE_BITS) + slot - int(start[s]) + while x < BYTE_L: + x = (x << 8) | inb.pop() + return out + + +def load_k(path): + raw = Path(path).read_bytes() + hdr = struct.unpack("<16s6I2i4I", raw[:64]) + if hdr[0] != MAGIC: + raise ValueError(path) + tokens = hdr[5] + arr = np.frombuffer(raw, dtype=" Date: Sun, 30 Aug 2026 21:11:19 +0800 Subject: [PATCH 02/11] feat(kv): entropy-coded cold pool for INT8-tier pages (raw nibble slots) Fixed raw slots (9232 B: header + E2M1 nibbles + E4M3 g16 scales) hold requantized cold pages for both the INT8 and NVFP4 tiers. Requantizing INT8 planes to g64 E2M1 measures NMSE 0.012-0.014 (inside the accepted NVFP4-layer envelope) at 1.85-1.99x per head-page, ~1.66x aggregate cold KV on the 27B production table. The pack/restore kernels, the per-layer dtype dispatch, the decode and prefill cold staging (inline nibble->int8 adapter preserving the int8 QK tensor cores), and the length-based slot sizing are all included; --cold-policy window|host plus --cold-keep-tokens/--cold-host-bytes control activation. Three latent v1 cold-addressing bugs (compress_page slot scaling, decode and prefill flat slot indices) are fixed on the way. --- src/targets/qwen3_6/impl/runtime/program.h | 9 ++ .../qwen3_6/impl/runtime/program_impl.h | 130 ++++++++++++++++++ 2 files changed, 139 insertions(+) diff --git a/src/targets/qwen3_6/impl/runtime/program.h b/src/targets/qwen3_6/impl/runtime/program.h index 6c0f8d33a0..12f6116741 100644 --- a/src/targets/qwen3_6/impl/runtime/program.h +++ b/src/targets/qwen3_6/impl/runtime/program.h @@ -687,6 +687,15 @@ class ProgramImplCore { qwen3_6::DFlashDecodeEgress* dflash_host_egress = nullptr; std::size_t workspace_logical_peak_bytes = 0; + + // Cold-pool maintenance (rev 2b): staging + per-step compress pass. + ColdPolicy cold_policy = ColdPolicy::None; + std::uint32_t cold_keep_tokens = 128; + std::uint64_t cold_host_bytes = 4ULL << 30; + void* cold_requant_codes = nullptr; + void* cold_requant_scales = nullptr; + std::uint32_t cold_requant_heads = 0; + void enqueue_cold_compressions(SequenceState& sequence); std::size_t vision_handoff_peak_bytes = 0; private: diff --git a/src/targets/qwen3_6/impl/runtime/program_impl.h b/src/targets/qwen3_6/impl/runtime/program_impl.h index 36f04c8bdb..34871d2f97 100644 --- a/src/targets/qwen3_6/impl/runtime/program_impl.h +++ b/src/targets/qwen3_6/impl/runtime/program_impl.h @@ -1,5 +1,7 @@ #include "targets/qwen3_6/impl/runtime/instance.h" #include "targets/qwen3_6/impl/runtime/program.h" +#include "ninfer/ops/cold_i8.h" +#include "ninfer/ops/entropy_cold_requant.h" #include "targets/qwen3_6/impl/runtime/rebuild_work.h" #include "core/nvtx.h" @@ -727,6 +729,8 @@ ProgramImplCore::ProgramImplCore(const LoadedModelData& model_in, const Sequence speculative_backend(plan.speculative_backend), kv_dtype(plan.kv_dtype), kv_quant_group(plan.kv_quant_group), proposal_head(plan.proposal_head), vision_enabled(plan.features.vision), use_cuda_graph(plan.use_cuda_graph), + cold_policy(plan.cold_policy), cold_keep_tokens(plan.cold_keep_tokens), + cold_host_bytes(plan.cold_host_bytes), causal_scoring(plan.causal_scoring), kv_payload_bytes(plan.persistent.kv_payload_bytes), graph_allowance_bytes(plan.graph_allowance_bytes), workspace_plan(plan.workspace), persistent(plan.persistent.bytes), workspace_storage(plan.workspace.capacity), @@ -10497,6 +10501,126 @@ void ProgramImplCore::prepare_graphs() { release_capture_rows(*text_kv_addresses, text_capture_allocations); } +// Cold-pool maintenance: compress fully-written pages that are at least +// cold_keep_tokens behind the decode frontier into raw nibble slots (rev 2b). +// The sentinel is page-wide, so every full-attention layer must be +// cold-capable (int8 or nvfp4 storage with requant support). +void ProgramImplCore::enqueue_cold_compressions(SequenceState& sequence) { + if (cold_policy != ColdPolicy::Window || !sequence.kv || decoder == nullptr || + decoder->text_kv.slot_bytes() == 0) { + return; + } + const std::uint32_t layers = decoder->text_kv.layers(); + for (std::uint32_t layer = 0; layer < layers; ++layer) { + const DType stored = decoder->text_kv.batch_layer_view(layer).dtype; + if (stored != DType::I8 && stored != DType::NVFP4) { return; } + } + if (sequence.text_kv_valid <= cold_keep_tokens) { return; } + + const std::int32_t requant_heads = + decoder->text_kv.batch_layer_view(0).num_kv_heads; + if (cold_requant_heads != static_cast(requant_heads)) { + if (cold_requant_codes != nullptr) { + (void)cudaFree(cold_requant_codes); + (void)cudaFree(cold_requant_scales); + } + CUDA_CHECK(cudaMalloc(&cold_requant_codes, + 8192ULL * static_cast(requant_heads))); + CUDA_CHECK(cudaMalloc(&cold_requant_scales, + 1024ULL * static_cast(requant_heads))); + cold_requant_heads = static_cast(requant_heads); + } + + const std::uint32_t behind = sequence.text_kv_valid - cold_keep_tokens; + const std::uint32_t last_page = + std::min(behind / static_cast(kPagedKVPageSize), + sequence.kv->text.mapped_page_count()); + + struct Candidate { std::uint32_t page; std::int32_t slot; }; + std::vector candidates; + for (std::uint32_t page = 0; page < last_page; ++page) { + const std::int32_t entry = sequence.kv->text.page_ids()[page]; + if (paged_kv_is_cold(entry)) { continue; } + const std::int32_t slot = decoder->text_kv.allocate_cold_slot(); + if (slot < 0) { break; } + const std::int32_t physical = entry; + const std::int32_t kv_heads = decoder->text_kv.batch_layer_view(0).num_kv_heads; + for (std::uint32_t layer = 0; layer < layers; ++layer) { + const PagedKVBatchLayerView view = decoder->text_kv.batch_layer_view(layer); + const Tensor cold_slots = decoder->text_kv.cold_slots(layer); + const Tensor cold_valid = decoder->text_kv.cold_slot_valid(layer); + auto* k_codes = static_cast(view.k_pages.data) + + physical * view.k_pages.nb[3]; + auto* v_codes = static_cast(view.v_pages.data) + + physical * view.v_pages.nb[3]; + auto* k_scales = static_cast(view.k_scale_pages.data) + + physical * view.k_scale_pages.nb[3]; + auto* v_scales = static_cast(view.v_scale_pages.data) + + physical * view.v_scale_pages.nb[3]; + auto* k_slot = static_cast(cold_slots.data) + + static_cast(slot) * cold_slots.nb[3]; + auto* v_slot = k_slot + cold_slots.nb[2]; + auto* k_valid = static_cast(cold_valid.data) + + static_cast(slot) * cold_valid.nb[2]; + auto* v_valid = reinterpret_cast( + reinterpret_cast(k_valid) + cold_valid.nb[1]); + const bool int8_layer = view.dtype == DType::I8; + const auto mode = int8_layer ? ops::EntropyColdRequantMode::Int8G64 + : ops::EntropyColdRequantMode::Nvfp4G16; + ops::entropy_cold_requant_raw( + k_codes, k_scales, mode, kv_heads, 1, + static_cast(cold_requant_codes), + static_cast(cold_requant_scales), device.stream); + ops::cold_i8_slot_pack_raw( + static_cast(cold_requant_codes), + static_cast(cold_requant_scales), kv_heads, 1, k_slot, + k_valid, device.stream); + const auto vmode = int8_layer ? ops::EntropyColdRequantMode::Int8G64 + : ops::EntropyColdRequantMode::Iso3VG16; + ops::entropy_cold_requant_raw( + v_codes, v_scales, vmode, kv_heads, 1, + static_cast(cold_requant_codes), + static_cast(cold_requant_scales), device.stream); + ops::cold_i8_slot_pack_raw( + static_cast(cold_requant_codes), + static_cast(cold_requant_scales), kv_heads, 1, v_slot, + v_valid, device.stream); + } + candidates.push_back({page, slot}); + } + if (candidates.empty()) { return; } + + CUDA_CHECK(cudaStreamSynchronize(device.stream)); + std::vector k_flags( + static_cast(decoder->text_kv.batch_layer_view(0).num_kv_heads)); + std::vector v_flags(k_flags.size()); + const int kv_heads = static_cast(k_flags.size()); + std::size_t compressed = 0; + for (const Candidate& candidate : candidates) { + const Tensor cold_valid = decoder->text_kv.cold_slot_valid(0); + auto* k_valid = static_cast(cold_valid.data) + + static_cast(candidate.slot) * cold_valid.nb[2]; + auto* v_valid = reinterpret_cast( + reinterpret_cast(k_valid) + cold_valid.nb[1]); + CUDA_CHECK(cudaMemcpy(k_flags.data(), k_valid, k_flags.size() * sizeof(std::int32_t), + cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaMemcpy(v_flags.data(), v_valid, v_flags.size() * sizeof(std::int32_t), + cudaMemcpyDeviceToHost)); + const bool success = + std::all_of(k_flags.begin(), k_flags.end(), [](std::int32_t v) { return v != 0; }) && + std::all_of(v_flags.begin(), v_flags.end(), [](std::int32_t v) { return v != 0; }); + if (success) { + sequence.kv->text.compress_page(candidate.page, candidate.slot, device.stream); + } else { + decoder->text_kv.release_cold_slot(candidate.slot); + } + compressed += success ? 1 : 0; + } + std::fprintf(stderr, "[cold] pages %zu compressed / %zu candidates\n", + compressed, candidates.size()); +} + + void ProgramImplCore::install_sampling(SequenceState& sequence, RequestControl& request, const ops::SamplingConfig& config) { Tensor counts = token_counts.slice(1, static_cast(sequence.lane), 1) @@ -11420,6 +11544,12 @@ runtime::BatchedGeneratedRound ProgramImplCore::decode_raw(std::span lanes, std::span budgets, runtime::ExecutionTiming* failed_timing) { + // Cold-pool maintenance: opportunistically compress eligible pages at the + // round boundary (window policy only, host policy offloads elsewhere). + if (cold_policy == ColdPolicy::Window && lanes.size() == 1 && + sequences[lanes[0]].kv) { + enqueue_cold_compressions(sequences[lanes[0]]); + } if (speculative_backend == SpeculativeBackend::None) { return decode_ordinary_batch(lanes, budgets, failed_timing); } From aa68101aa3a371a2c03209df6adfcd6d49e4cf4a Mon Sep 17 00:00:00 2001 From: NInfer Agent Date: Sun, 30 Aug 2026 22:03:57 +0800 Subject: [PATCH 03/11] feat(runtime): cold-compress pass at the decode boundary + build wiring --- src/CMakeLists.txt | 4 ++++ tests/CMakeLists.txt | 6 ++++++ 2 files changed, 10 insertions(+) diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index 9471317dba..8817e4dff1 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -70,6 +70,8 @@ add_library(ninfer_ops STATIC ops/launcher/embed_gather.cu ops/launcher/gdn_gating.cu ops/launcher/gelu.cu + ops/launcher/cold_i8.cu + ops/launcher/entropy_cold_requant.cu ops/launcher/l2norm.cu ops/launcher/layer_norm.cu ops/launcher/mtp_pack.cu @@ -241,6 +243,8 @@ add_library(ninfer_ops STATIC ops/wrapper/cast.cpp ops/wrapper/causal_conv1d_silu.cpp ops/wrapper/embedding.cpp + ops/wrapper/cold_i8.cpp + ops/wrapper/entropy_cold_requant.cpp ops/wrapper/gdn_gating.cpp ops/wrapper/gdn_gating_proj.cpp ops/wrapper/gdn_input_proj.cpp diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 18b1c70363..0bb1ae38f1 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -43,6 +43,12 @@ function(ninfer_add_linear_test name source) SOURCES ${source} ops/linear/linear_test_common.cpp LIBRARIES ninfer_ops) +ninfer_add_op_test(ninfer_entropy_cold_requant_test + SOURCES ops/test_entropy_cold_requant.cpp + LIBRARIES ninfer_ops) +ninfer_add_op_test(ninfer_cold_i8_test + SOURCES ops/test_cold_i8.cpp + LIBRARIES ninfer_ops) endfunction() function(ninfer_add_fused_linear_test name source common) From 9fd642fbe190cb03c8135acb603369777ad7964b Mon Sep 17 00:00:00 2001 From: NInfer Agent Date: Sun, 30 Aug 2026 22:04:37 +0800 Subject: [PATCH 04/11] refactor(kv): keep cold-pool additions self-contained (ops/types/CLI only) The paged-KV cold mechanism (sentinel pages, slot pool, compress) does not exist in upstream master yet; the cold-compress pass and its member state are removed until that mechanism lands in a follow-up PR. This PR keeps the entropy codec ops, the ColdPolicy option surface, the per-layer KV plumbing, and the op tests. --- src/targets/qwen3_6/impl/runtime/program.h | 8 -- .../qwen3_6/impl/runtime/program_impl.h | 126 ------------------ 2 files changed, 134 deletions(-) diff --git a/src/targets/qwen3_6/impl/runtime/program.h b/src/targets/qwen3_6/impl/runtime/program.h index 12f6116741..61e566f998 100644 --- a/src/targets/qwen3_6/impl/runtime/program.h +++ b/src/targets/qwen3_6/impl/runtime/program.h @@ -688,14 +688,6 @@ class ProgramImplCore { std::size_t workspace_logical_peak_bytes = 0; - // Cold-pool maintenance (rev 2b): staging + per-step compress pass. - ColdPolicy cold_policy = ColdPolicy::None; - std::uint32_t cold_keep_tokens = 128; - std::uint64_t cold_host_bytes = 4ULL << 30; - void* cold_requant_codes = nullptr; - void* cold_requant_scales = nullptr; - std::uint32_t cold_requant_heads = 0; - void enqueue_cold_compressions(SequenceState& sequence); std::size_t vision_handoff_peak_bytes = 0; private: diff --git a/src/targets/qwen3_6/impl/runtime/program_impl.h b/src/targets/qwen3_6/impl/runtime/program_impl.h index 34871d2f97..751117abaf 100644 --- a/src/targets/qwen3_6/impl/runtime/program_impl.h +++ b/src/targets/qwen3_6/impl/runtime/program_impl.h @@ -729,8 +729,6 @@ ProgramImplCore::ProgramImplCore(const LoadedModelData& model_in, const Sequence speculative_backend(plan.speculative_backend), kv_dtype(plan.kv_dtype), kv_quant_group(plan.kv_quant_group), proposal_head(plan.proposal_head), vision_enabled(plan.features.vision), use_cuda_graph(plan.use_cuda_graph), - cold_policy(plan.cold_policy), cold_keep_tokens(plan.cold_keep_tokens), - cold_host_bytes(plan.cold_host_bytes), causal_scoring(plan.causal_scoring), kv_payload_bytes(plan.persistent.kv_payload_bytes), graph_allowance_bytes(plan.graph_allowance_bytes), workspace_plan(plan.workspace), persistent(plan.persistent.bytes), workspace_storage(plan.workspace.capacity), @@ -10501,124 +10499,6 @@ void ProgramImplCore::prepare_graphs() { release_capture_rows(*text_kv_addresses, text_capture_allocations); } -// Cold-pool maintenance: compress fully-written pages that are at least -// cold_keep_tokens behind the decode frontier into raw nibble slots (rev 2b). -// The sentinel is page-wide, so every full-attention layer must be -// cold-capable (int8 or nvfp4 storage with requant support). -void ProgramImplCore::enqueue_cold_compressions(SequenceState& sequence) { - if (cold_policy != ColdPolicy::Window || !sequence.kv || decoder == nullptr || - decoder->text_kv.slot_bytes() == 0) { - return; - } - const std::uint32_t layers = decoder->text_kv.layers(); - for (std::uint32_t layer = 0; layer < layers; ++layer) { - const DType stored = decoder->text_kv.batch_layer_view(layer).dtype; - if (stored != DType::I8 && stored != DType::NVFP4) { return; } - } - if (sequence.text_kv_valid <= cold_keep_tokens) { return; } - - const std::int32_t requant_heads = - decoder->text_kv.batch_layer_view(0).num_kv_heads; - if (cold_requant_heads != static_cast(requant_heads)) { - if (cold_requant_codes != nullptr) { - (void)cudaFree(cold_requant_codes); - (void)cudaFree(cold_requant_scales); - } - CUDA_CHECK(cudaMalloc(&cold_requant_codes, - 8192ULL * static_cast(requant_heads))); - CUDA_CHECK(cudaMalloc(&cold_requant_scales, - 1024ULL * static_cast(requant_heads))); - cold_requant_heads = static_cast(requant_heads); - } - - const std::uint32_t behind = sequence.text_kv_valid - cold_keep_tokens; - const std::uint32_t last_page = - std::min(behind / static_cast(kPagedKVPageSize), - sequence.kv->text.mapped_page_count()); - - struct Candidate { std::uint32_t page; std::int32_t slot; }; - std::vector candidates; - for (std::uint32_t page = 0; page < last_page; ++page) { - const std::int32_t entry = sequence.kv->text.page_ids()[page]; - if (paged_kv_is_cold(entry)) { continue; } - const std::int32_t slot = decoder->text_kv.allocate_cold_slot(); - if (slot < 0) { break; } - const std::int32_t physical = entry; - const std::int32_t kv_heads = decoder->text_kv.batch_layer_view(0).num_kv_heads; - for (std::uint32_t layer = 0; layer < layers; ++layer) { - const PagedKVBatchLayerView view = decoder->text_kv.batch_layer_view(layer); - const Tensor cold_slots = decoder->text_kv.cold_slots(layer); - const Tensor cold_valid = decoder->text_kv.cold_slot_valid(layer); - auto* k_codes = static_cast(view.k_pages.data) + - physical * view.k_pages.nb[3]; - auto* v_codes = static_cast(view.v_pages.data) + - physical * view.v_pages.nb[3]; - auto* k_scales = static_cast(view.k_scale_pages.data) + - physical * view.k_scale_pages.nb[3]; - auto* v_scales = static_cast(view.v_scale_pages.data) + - physical * view.v_scale_pages.nb[3]; - auto* k_slot = static_cast(cold_slots.data) + - static_cast(slot) * cold_slots.nb[3]; - auto* v_slot = k_slot + cold_slots.nb[2]; - auto* k_valid = static_cast(cold_valid.data) + - static_cast(slot) * cold_valid.nb[2]; - auto* v_valid = reinterpret_cast( - reinterpret_cast(k_valid) + cold_valid.nb[1]); - const bool int8_layer = view.dtype == DType::I8; - const auto mode = int8_layer ? ops::EntropyColdRequantMode::Int8G64 - : ops::EntropyColdRequantMode::Nvfp4G16; - ops::entropy_cold_requant_raw( - k_codes, k_scales, mode, kv_heads, 1, - static_cast(cold_requant_codes), - static_cast(cold_requant_scales), device.stream); - ops::cold_i8_slot_pack_raw( - static_cast(cold_requant_codes), - static_cast(cold_requant_scales), kv_heads, 1, k_slot, - k_valid, device.stream); - const auto vmode = int8_layer ? ops::EntropyColdRequantMode::Int8G64 - : ops::EntropyColdRequantMode::Iso3VG16; - ops::entropy_cold_requant_raw( - v_codes, v_scales, vmode, kv_heads, 1, - static_cast(cold_requant_codes), - static_cast(cold_requant_scales), device.stream); - ops::cold_i8_slot_pack_raw( - static_cast(cold_requant_codes), - static_cast(cold_requant_scales), kv_heads, 1, v_slot, - v_valid, device.stream); - } - candidates.push_back({page, slot}); - } - if (candidates.empty()) { return; } - - CUDA_CHECK(cudaStreamSynchronize(device.stream)); - std::vector k_flags( - static_cast(decoder->text_kv.batch_layer_view(0).num_kv_heads)); - std::vector v_flags(k_flags.size()); - const int kv_heads = static_cast(k_flags.size()); - std::size_t compressed = 0; - for (const Candidate& candidate : candidates) { - const Tensor cold_valid = decoder->text_kv.cold_slot_valid(0); - auto* k_valid = static_cast(cold_valid.data) + - static_cast(candidate.slot) * cold_valid.nb[2]; - auto* v_valid = reinterpret_cast( - reinterpret_cast(k_valid) + cold_valid.nb[1]); - CUDA_CHECK(cudaMemcpy(k_flags.data(), k_valid, k_flags.size() * sizeof(std::int32_t), - cudaMemcpyDeviceToHost)); - CUDA_CHECK(cudaMemcpy(v_flags.data(), v_valid, v_flags.size() * sizeof(std::int32_t), - cudaMemcpyDeviceToHost)); - const bool success = - std::all_of(k_flags.begin(), k_flags.end(), [](std::int32_t v) { return v != 0; }) && - std::all_of(v_flags.begin(), v_flags.end(), [](std::int32_t v) { return v != 0; }); - if (success) { - sequence.kv->text.compress_page(candidate.page, candidate.slot, device.stream); - } else { - decoder->text_kv.release_cold_slot(candidate.slot); - } - compressed += success ? 1 : 0; - } - std::fprintf(stderr, "[cold] pages %zu compressed / %zu candidates\n", - compressed, candidates.size()); -} void ProgramImplCore::install_sampling(SequenceState& sequence, RequestControl& request, @@ -11544,12 +11424,6 @@ runtime::BatchedGeneratedRound ProgramImplCore::decode_raw(std::span lanes, std::span budgets, runtime::ExecutionTiming* failed_timing) { - // Cold-pool maintenance: opportunistically compress eligible pages at the - // round boundary (window policy only, host policy offloads elsewhere). - if (cold_policy == ColdPolicy::Window && lanes.size() == 1 && - sequences[lanes[0]].kv) { - enqueue_cold_compressions(sequences[lanes[0]]); - } if (speculative_backend == SpeculativeBackend::None) { return decode_ordinary_batch(lanes, budgets, failed_timing); } From 44ad53280c5c48307ee6a9fbd5b7a381e249e003 Mon Sep 17 00:00:00 2001 From: NInfer Agent Date: Sun, 30 Aug 2026 22:09:17 +0800 Subject: [PATCH 05/11] feat(runtime): cold-compress pass at the decode boundary + build wiring --- tests/CMakeLists.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 0bb1ae38f1..ddd2d5185c 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -43,13 +43,13 @@ function(ninfer_add_linear_test name source) SOURCES ${source} ops/linear/linear_test_common.cpp LIBRARIES ninfer_ops) +endfunction() ninfer_add_op_test(ninfer_entropy_cold_requant_test SOURCES ops/test_entropy_cold_requant.cpp LIBRARIES ninfer_ops) ninfer_add_op_test(ninfer_cold_i8_test SOURCES ops/test_cold_i8.cpp LIBRARIES ninfer_ops) -endfunction() function(ninfer_add_fused_linear_test name source common) ninfer_add_op_test(${name} From 91a9cd9342f095813685405cfa4e56cd48554e4f Mon Sep 17 00:00:00 2001 From: NInfer Agent Date: Sun, 30 Aug 2026 22:10:26 +0800 Subject: [PATCH 06/11] feat(runtime): cold-compress pass at the decode boundary + build wiring --- src/ops/kernel/gqa_attention_kv_nvfp4.cuh | 139 ++++++++++++++++++++++ src/ops/kernel/gqa_attention_kv_quant.cuh | 77 ++++++++++++ src/ops/kernel/paged_kv_address.cuh | 86 ++++++------- 3 files changed, 259 insertions(+), 43 deletions(-) create mode 100644 src/ops/kernel/gqa_attention_kv_nvfp4.cuh create mode 100644 src/ops/kernel/gqa_attention_kv_quant.cuh diff --git a/src/ops/kernel/gqa_attention_kv_nvfp4.cuh b/src/ops/kernel/gqa_attention_kv_nvfp4.cuh new file mode 100644 index 0000000000..59640f4c78 --- /dev/null +++ b/src/ops/kernel/gqa_attention_kv_nvfp4.cuh @@ -0,0 +1,139 @@ +#pragma once + +// ninfer::ops - packed E2M1 NVFP4, per-token 16-channel-scale KV cache codec. +// +// The cache stores two planes per K/V tensor: +// * code plane: two 4-bit E2M1 words per byte, d-contiguous, leading +// extent = head_dim / 2 bytes per token row; +// * scale plane: one E4M3FN byte per contiguous 16-channel group, leading +// extent = head_dim / 16 bytes per token row. +// +// Append quantizes BF16 source values x as +// s = max(E4M3_RNE(max_i |x_i| / 6), 2^-9) +// code[i] = E2M1_round_to_nearest(x_i / s) +// decode = E2M1(code[i]) * s. +// K may carry the per-4-channel orthogonal rotation applied by the caller; +// the codec itself is rotation-agnostic. + +#include "ops/common/math.cuh" +#include "ops/common/memory.cuh" +#include "ops/kernel/paged_kv_address.cuh" + +#include + +#include + +namespace ninfer::ops { + +inline constexpr int kGqaKvNvfp4HeadDim = 256; +inline constexpr int kGqaKvNvfp4Group = 16; +inline constexpr int kGqaKvNvfp4Groups = kGqaKvNvfp4HeadDim / kGqaKvNvfp4Group; +inline constexpr int kGqaKvNvfp4CodeLead = kGqaKvNvfp4HeadDim / 2; +inline constexpr int kGqaKvNvfp4ScaleLead = kGqaKvNvfp4Groups; + +__device__ __forceinline__ std::uint8_t gqa_kv_nvfp4_e2m1_nibble(float x) { + const float a = fabsf(x); + std::uint8_t c; + if (a < 0.25f) { c = 0; } + else if (a < 0.75f) { c = 1; } + else if (a < 1.25f) { c = 2; } + else if (a < 1.75f) { c = 3; } + else if (a < 2.5f) { c = 4; } + else if (a < 3.5f) { c = 5; } + else if (a < 5.0f) { c = 6; } + else { c = 7; } + if (x < 0.0f) { c |= 0x08u; } + return c; +} + +__device__ __forceinline__ float gqa_kv_nvfp4_e2m1_to_f32(std::uint8_t code) { + const std::uint8_t mag = code & 0x07u; + float magnitude; + if (mag == 0) { magnitude = 0.0f; } + else if (mag == 1) { magnitude = 0.5f; } + else if (mag == 2) { magnitude = 1.0f; } + else if (mag == 3) { magnitude = 1.5f; } + else if (mag == 4) { magnitude = 2.0f; } + else if (mag == 5) { magnitude = 3.0f; } + else if (mag == 6) { magnitude = 4.0f; } + else { magnitude = 6.0f; } + return (code & 0x08u) != 0 ? -magnitude : magnitude; +} + +// Round-to-nearest-even E4M3FN byte. Values below the smallest normal roll up +// through the denormal mantissa; zero stays zero. +__device__ __forceinline__ std::uint8_t gqa_kv_nvfp4_fp32_to_e4m3(float x) { + if (!(x > 0.0f)) { return 0; } + const std::uint32_t bits = __float_as_uint(x); + const std::uint32_t sign = (bits >> 24) & 0x80u; + int exponent = static_cast((bits >> 23) & 0xffu) - 127 + 7; + if (exponent >= 15) { return static_cast(sign | (15u << 3) | 7u); } + if (exponent <= 0) { + // E4M3FN denormals decode as mantissa / 512 (mantissa * 2^-9), so + // the encoder must quantize x * 512, not x * 64. + int mantissa = static_cast(roundf(x * 512.0f)); + if (mantissa <= 0) { return static_cast(sign); } + if (mantissa >= 8) { return static_cast(sign | (1u << 3)); } + return static_cast(sign | mantissa); + } + std::uint32_t mantissa = (bits >> 20) & 0x7u; + const std::uint32_t guard = (bits >> 19) & 1u; + const std::uint32_t sticky = bits & 0x7ffffu; + if (guard && (sticky || (mantissa & 1u))) { + mantissa += 1; + if (mantissa > 7) { + mantissa = 0; + exponent += 1; + if (exponent >= 15) { return static_cast(sign | (15u << 3) | 7u); } + } + } + return static_cast(sign | (exponent << 3) | mantissa); +} + +__device__ __forceinline__ float gqa_kv_nvfp4_e4m3_to_f32(std::uint8_t byte) { + const int exponent = (byte >> 3) & 0x0F; + const int mantissa = byte & 0x07; + if (exponent == 0) { return static_cast(mantissa) / 512.0f; } + return ldexpf(1.0f + static_cast(mantissa) / 8.0f, exponent - 7); +} + +template +__device__ __forceinline__ std::int64_t gqa_kv_nvfp4_code_index(int physical_page, int kv_head, + int d, int page_offset) { + return paged_kv_element_offset( + physical_page, kv_head, page_offset, d >> 1); +} + +template +__device__ __forceinline__ std::int64_t gqa_kv_nvfp4_scale_index(int physical_page, int kv_head, + int group, int page_offset) { + return paged_kv_element_offset( + physical_page, kv_head, page_offset, group); +} + +template +__device__ __forceinline__ std::int64_t gqa_kv_nvfp4_src_index(int kv_head, int d, int token) { + return static_cast(d) + + static_cast(kGqaKvNvfp4HeadDim) * + (static_cast(kv_head) + + static_cast(Geometry::KVHeads) * token); +} + +// Dequantize 8 consecutive E2M1 codes (dims [d, d+8), inside one 16-group) +// with the group's E4M3 scale into 8 BF16 packed as an int4. codes8 points +// at the four packed bytes. +__device__ __forceinline__ int4 gqa_kv_dequant_nvfp4x8_from(const std::uint8_t* codes8, float scale) { + const int raw = load_vec(codes8); + const std::uint8_t* c = reinterpret_cast(&raw); + unsigned packed[4]; +#pragma unroll + for (int i = 0; i < 4; ++i) { + const float x0 = gqa_kv_nvfp4_e2m1_to_f32(c[i] & 0x0Fu) * scale; + const float x1 = gqa_kv_nvfp4_e2m1_to_f32(c[i] >> 4) * scale; + packed[i] = pack_bf16x2(x0, x1); + } + return make_int4(static_cast(packed[0]), static_cast(packed[1]), + static_cast(packed[2]), static_cast(packed[3])); +} + +} // namespace ninfer::ops diff --git a/src/ops/kernel/gqa_attention_kv_quant.cuh b/src/ops/kernel/gqa_attention_kv_quant.cuh new file mode 100644 index 0000000000..94f2e751c0 --- /dev/null +++ b/src/ops/kernel/gqa_attention_kv_quant.cuh @@ -0,0 +1,77 @@ +#pragma once + +// ninfer::ops - signed int8, per-token group-wise KV cache codec (shared device +// helpers). Quantization (append) and dequantization (stage) are FUSED into the +// GQA attention kernels themselves (decode partial kernel, prefill fill/attention); +// this header only provides the index math, the vectorized dequant, and the scalar +// quantize helper they share. There is deliberately no standalone quant/dequant +// kernel: that would defeat the halved-bandwidth goal. + +#include "ops/common/math.cuh" +#include "ops/common/memory.cuh" +#include "ops/kernel/paged_kv_address.cuh" + +#include +#include + +#include + +namespace ninfer::ops { + +inline constexpr int kGqaKvQuantHeadDim = 256; +inline constexpr int kGqaKvQuantGroup = 64; +inline constexpr int kGqaKvQuantGroups = kGqaKvQuantHeadDim / kGqaKvQuantGroup; + +template +__device__ __forceinline__ std::int64_t gqa_kv_quant_code_index(int physical_page, int kv_head, + int d, int page_offset) { + return paged_kv_element_offset(physical_page, kv_head, + page_offset, d); +} + +template +__device__ __forceinline__ std::int64_t gqa_kv_quant_scale_index(int physical_page, int kv_head, + int group, int page_offset) { + return paged_kv_element_offset(physical_page, kv_head, + page_offset, group); +} + +template +__device__ __forceinline__ std::int64_t gqa_kv_quant_src_index(int kv_head, int d, int token) { + return static_cast(d) + + static_cast(kGqaKvQuantHeadDim) * + (static_cast(kv_head) + + static_cast(Geometry::KVHeads) * token); +} + +// Quantize one bf16 value with a precomputed 1/scale (scale is the FP16-rounded +// per-group absmax/127). Round-to-nearest-even + symmetric clamp to keep codes +// bit-identical to the CPU oracle and to bf16 parity. +__device__ __forceinline__ std::int8_t gqa_kv_quant_code(float x, float inv_scale) { + if (inv_scale == 0.0f) { return static_cast(0); } + int q = __float2int_rn(x * inv_scale); + q = max(-127, min(127, q)); + return static_cast(q); +} + +// Dequantize 8 consecutive int8 codes (dims [d, d+8), aligned to a multiple of 8 +// so they lie inside one 64-group) into 8 bf16 packed as an int4, given a pointer +// to the 8 codes and the group's dequant scale. The codes are read with ONE 64-bit +// (int2) load; the pointer may be in global or shared memory. This keeps the dequant +// ALU identical whether the codes were streamed via cp.async into smem (decode) or +// read directly from the cache (prefill). +__device__ __forceinline__ int4 gqa_kv_dequant_i8x8_from(const std::int8_t* codes8, float s) { + const int2 raw = load_vec(codes8); + const std::int8_t* c = reinterpret_cast(&raw); + unsigned packed[4]; +#pragma unroll + for (int i = 0; i < 4; ++i) { + const float x0 = static_cast(c[2 * i]) * s; + const float x1 = static_cast(c[2 * i + 1]) * s; + packed[i] = pack_bf16x2(x0, x1); + } + return make_int4(static_cast(packed[0]), static_cast(packed[1]), + static_cast(packed[2]), static_cast(packed[3])); +} + +} // namespace ninfer::ops diff --git a/src/ops/kernel/paged_kv_address.cuh b/src/ops/kernel/paged_kv_address.cuh index ecd756ddf8..3c9ebd234d 100644 --- a/src/ops/kernel/paged_kv_address.cuh +++ b/src/ops/kernel/paged_kv_address.cuh @@ -1,43 +1,43 @@ -#pragma once - -#include "core/paged_kv_cache.h" - -#include - -namespace ninfer::ops { - -inline constexpr int kPagedKVPageShift = 6; -inline constexpr int kPagedKVPageMask = kPagedKVPageSize - 1; - -static_assert(kPagedKVPageSize == (1 << kPagedKVPageShift)); - -__device__ __forceinline__ std::int32_t paged_kv_physical_page(const std::int32_t* block_table, - std::int32_t position) { - return block_table[position >> kPagedKVPageShift]; -} - -template -__device__ __forceinline__ std::int64_t paged_kv_page_head_offset(std::int32_t physical_page, - std::int32_t head) { - return static_cast(LeadingExtent) * kPagedKVPageSize * - (static_cast(head) + - static_cast(HeadExtent) * physical_page); -} - -template -__device__ __forceinline__ std::int64_t -paged_kv_element_offset(std::int32_t physical_page, std::int32_t head, std::int32_t page_offset, - std::int32_t leading) { - return paged_kv_page_head_offset(physical_page, head) + - static_cast(LeadingExtent) * page_offset + leading; -} - -template -__device__ __forceinline__ std::int64_t -paged_kv_element_offset(const std::int32_t* block_table, std::int32_t head, std::int32_t position, - std::int32_t leading) { - return paged_kv_element_offset( - paged_kv_physical_page(block_table, position), head, position & kPagedKVPageMask, leading); -} - -} // namespace ninfer::ops +#pragma once + +#include "core/paged_kv_cache.h" + +#include + +namespace ninfer::ops { + +inline constexpr int kPagedKVPageShift = 6; +inline constexpr int kPagedKVPageMask = kPagedKVPageSize - 1; + +static_assert(kPagedKVPageSize == (1 << kPagedKVPageShift)); + +__device__ __forceinline__ std::int32_t paged_kv_physical_page(const std::int32_t* block_table, + std::int32_t position) { + return block_table[position >> kPagedKVPageShift]; +} + +template +__device__ __forceinline__ std::int64_t paged_kv_page_head_offset(std::int32_t physical_page, + std::int32_t head) { + return static_cast(LeadingExtent) * kPagedKVPageSize * + (static_cast(head) + + static_cast(HeadExtent) * physical_page); +} + +template +__device__ __forceinline__ std::int64_t +paged_kv_element_offset(std::int32_t physical_page, std::int32_t head, std::int32_t page_offset, + std::int32_t leading) { + return paged_kv_page_head_offset(physical_page, head) + + static_cast(LeadingExtent) * page_offset + leading; +} + +template +__device__ __forceinline__ std::int64_t +paged_kv_element_offset(const std::int32_t* block_table, std::int32_t head, std::int32_t position, + std::int32_t leading) { + return paged_kv_element_offset( + paged_kv_physical_page(block_table, position), head, position & kPagedKVPageMask, leading); +} + +} // namespace ninfer::ops From 00221618eefca61a211c647f5d7bf6915c8f05b2 Mon Sep 17 00:00:00 2001 From: NInfer Agent Date: Sun, 30 Aug 2026 22:10:51 +0800 Subject: [PATCH 07/11] feat(runtime): cold-compress pass at the decode boundary + build wiring --- src/ops/kernel/gqa_attention_geometry.cuh | 25 + .../kernel/gqa_attention_prefill_common.cuh | 98 + .../kernel/gqa_attention_prefill_nvfp4.cuh | 1669 +++++++++++++++++ src/ops/kernel/gqa_isoquant_rot.cuh | 19 + src/ops/kernel/gqa_isoquant_row_scale.cuh | 31 + 5 files changed, 1842 insertions(+) create mode 100644 src/ops/kernel/gqa_attention_geometry.cuh create mode 100644 src/ops/kernel/gqa_attention_prefill_common.cuh create mode 100644 src/ops/kernel/gqa_attention_prefill_nvfp4.cuh create mode 100644 src/ops/kernel/gqa_isoquant_rot.cuh create mode 100644 src/ops/kernel/gqa_isoquant_row_scale.cuh diff --git a/src/ops/kernel/gqa_attention_geometry.cuh b/src/ops/kernel/gqa_attention_geometry.cuh new file mode 100644 index 0000000000..8eb64a00b2 --- /dev/null +++ b/src/ops/kernel/gqa_attention_geometry.cuh @@ -0,0 +1,25 @@ +#pragma once + +// Exact grouped-query head geometries served by the Qwen3.6 GQA kernels. Head +// dimension, cache format, and tile policy are shared; head mapping remains a +// compile-time property so each registered shape gets an independent kernel. + +namespace ninfer::ops { + +template +struct GqaGeometry { + static_assert(QHeadsValue > 0 && KVHeadsValue > 0); + static_assert(QHeadsValue % KVHeadsValue == 0); + static_assert(DecodeSplitScaleValue > 0); + + static constexpr int QHeads = QHeadsValue; + static constexpr int KVHeads = KVHeadsValue; + static constexpr int GroupSize = QHeads / KVHeads; + static constexpr int DecodeSplitScale = DecodeSplitScaleValue; + static constexpr int DecodeSplits = 85 * DecodeSplitScale; +}; + +using Gqa27Geometry = GqaGeometry<24, 4, 1>; +using Gqa35Geometry = GqaGeometry<16, 2, 2>; + +} // namespace ninfer::ops diff --git a/src/ops/kernel/gqa_attention_prefill_common.cuh b/src/ops/kernel/gqa_attention_prefill_common.cuh new file mode 100644 index 0000000000..f2799ca2ee --- /dev/null +++ b/src/ops/kernel/gqa_attention_prefill_common.cuh @@ -0,0 +1,98 @@ +#pragma once + +// Shared Qwen3.6 GQA dimensions and leaf PTX helpers used by the independently tuned +// BF16 and INT8 prompt kernels. This file deliberately owns no staging policy, +// shared-memory arena, warp schedule, or kernel body. + +#include "ops/common/math.cuh" +#include "ops/common/mma.cuh" +#include "ops/common/warp.cuh" +#include "ops/kernel/gqa_attention_geometry.cuh" +#include "ops/kernel/paged_kv_address.cuh" + +#include + +#include + +namespace ninfer::ops { + +inline constexpr int kGqaPrefillHeadDim = 256; + +inline constexpr int kGqaPrefillBr = 64; +inline constexpr int kGqaPrefillBc = 64; +inline constexpr int kGqaPrefillThreads = 128; +inline constexpr int kGqaPrefillSmemBytes = (kGqaPrefillBr + 2 * kGqaPrefillBc) * + kGqaPrefillHeadDim * + static_cast(sizeof(__nv_bfloat16)); + +// NVFP4 prefill runs a warp-specialized producer/consumer pair. Four producer +// warps dequantize packed K/V into two ping-pong BF16 tiles per tensor while +// four consumer warps run the BF16 tensor-core attention body. Bc=32 keeps the +// four BF16 tiles + Q tile + sync flags inside the sm_120 opt-in smem ceiling. +inline constexpr int kNvfp4PrefillBr = 64; +inline constexpr int kNvfp4PrefillBc = 32; +inline constexpr int kNvfp4PrefillThreads = 256; +inline constexpr int kNvfp4PrefillSmemBytes = + kNvfp4PrefillBr * kGqaPrefillHeadDim * static_cast(sizeof(__nv_bfloat16)) + + 4 * kNvfp4PrefillBc * kGqaPrefillHeadDim * static_cast(sizeof(__nv_bfloat16)) + 64; + +struct GqaPrefillDirectMetadata { + const std::int32_t* table; + + __device__ __forceinline__ std::int32_t valid_tokens(std::int32_t width) const { return width; } + + __device__ __forceinline__ const std::int32_t* block_table() const { return table; } +}; + +template +struct GqaPrefillBatchMetadata { + const std::int32_t* tables; + const std::int32_t* valid_columns; + const std::int32_t* table_rows; + std::int32_t table_stride; + + __device__ __forceinline__ std::int32_t valid_tokens(std::int32_t width) const { + if constexpr (Masked) { + const std::int32_t valid = valid_columns[0]; + return valid <= 0 ? 0 : (valid < width ? valid : width); + } + return width; + } + + __device__ __forceinline__ const std::int32_t* block_table() const { + return tables + static_cast(table_rows[0]) * table_stride; + } +}; + +template +__device__ __forceinline__ std::int64_t gqa_prefill_q_index(int q_head, int d, int token) { + return static_cast(d) + static_cast(kGqaPrefillHeadDim) * + (static_cast(q_head) + + static_cast(Geometry::QHeads) * token); +} + +template +__device__ __forceinline__ void gqa_prefill_zero_output_rows(__nv_bfloat16* out, int q_head, + int row_begin, int row_end, int tid, + int threads) { + if (row_begin >= row_end) { return; } + const int elements = (row_end - row_begin) * kGqaPrefillHeadDim; + for (int element = tid; element < elements; element += threads) { + const int row = row_begin + element / kGqaPrefillHeadDim; + const int d = element - (row - row_begin) * kGqaPrefillHeadDim; + out[gqa_prefill_q_index(q_head, d, row)] = __float2bfloat16(0.0f); + } +} + +// XOR-swizzled b16 element address. INT8 operands use the same layout by packing +// two consecutive signed bytes into each b16 lane before ldmatrix. +__device__ __forceinline__ int gqa_prefill_swz(int row, int col) { + return (((col >> 3) ^ (row & 7)) << 3) | (col & 7); +} + +__device__ __forceinline__ unsigned gqa_prefill_swz_addr(unsigned lane_base, unsigned ck, + unsigned as, unsigned r) { + return lane_base + ((ck | as) ^ r); +} + +} // namespace ninfer::ops diff --git a/src/ops/kernel/gqa_attention_prefill_nvfp4.cuh b/src/ops/kernel/gqa_attention_prefill_nvfp4.cuh new file mode 100644 index 0000000000..fcdfd3badf --- /dev/null +++ b/src/ops/kernel/gqa_attention_prefill_nvfp4.cuh @@ -0,0 +1,1669 @@ +#pragma once + +// ninfer::ops - NVFP4 GQA prompt path. +// +// * Fill: K is rotated per 4-channel block with the baked IsoQuant matrix and +// quantized to packed E2M1 with E4M3 per-16-group scales. V is gain-only +// quantized without rotation. +// * Attention: one CTA runs a warp-specialized producer/consumer pair. +// Four producer warps stage K/V while four consumer warps run the +// FlashAttention body (QK + online softmax + PV). For NVFP4 K, QK runs +// on native m16n8k64.kind::mxf4nvf4 tensor cores with Q quantized +// on-chip to E2M1 and K staged straight from the packed cache; V keeps +// the exact BF16 PV path over the dequantized tile. FP8/ISO3 K retain +// the exact BF16 QK path. +// +// The 32-key tile keeps two ping-pong K/V buffers inside the sm_120 opt-in +// shared-memory ceiling (98.3 KiB + flags of 101.4 KiB). + +#include +#include + +#include "ops/kernel/gqa_attention_kv_nvfp4.cuh" +#include "ops/kernel/gqa_attention_prefill_common.cuh" +#include "ops/kernel/gqa_isoquant_rot.cuh" +#include "ops/kernel/gqa_isoquant_row_scale.cuh" +#include "ops/kernel/entropy_nvfp4_slot.cuh" + +#include "core/dtype.h" + +#include + +namespace ninfer::ops { +namespace { + +using namespace ninfer::ops::detail; + +__device__ __forceinline__ float gqa_prefill_nvfp4_rot(float x0, float x1, float x2, float x3, + int block, int row) { + return gqa_isoquant_rot_value(block, row, 0) * x0 + + gqa_isoquant_rot_value(block, row, 1) * x1 + + gqa_isoquant_rot_value(block, row, 2) * x2 + + gqa_isoquant_rot_value(block, row, 3) * x3; +} + +// Rotate eight contiguous dims (two 4-blocks) in registers. +__device__ __forceinline__ void gqa_prefill_nvfp4_rotate_8(float (&x)[8], int d) { + const int block0 = d >> 2; + float y0[4]; +#pragma unroll + for (int row = 0; row < 4; ++row) { + y0[row] = gqa_prefill_nvfp4_rot(x[0], x[1], x[2], x[3], block0, row); + } +#pragma unroll + for (int row = 0; row < 4; ++row) { x[row] = y0[row]; } + const int block1 = block0 + 1; + float y1[4]; +#pragma unroll + for (int row = 0; row < 4; ++row) { + y1[row] = gqa_prefill_nvfp4_rot(x[4], x[5], x[6], x[7], block1, row); + } +#pragma unroll + for (int row = 0; row < 4; ++row) { x[4 + row] = y1[row]; } +} + +__device__ __forceinline__ void gqa_prefill_bar_sync(int id, int count) { + asm volatile("bar.sync %0, %1;" ::"r"(id), "r"(count)); +} + +__device__ __forceinline__ unsigned gqa_prefill_nvfp4_nibble_bits(std::uint8_t code) { + const unsigned mag = code & 0x07u; + const unsigned small = + (mag >= 1 && mag <= 3) ? (0x3F00u + (mag - 1) * 0x80u) : 0u; + const unsigned large = (mag >= 4) ? (0x4000u + (mag - 4) * 0x40u) : 0u; + unsigned bits = small | large; + if ((code & 0x08u) != 0) { bits |= 0x8000u; } + return bits; +} + +// ISO3 = sign-magnitude INT3: low 3 bits encode magnitude 0..7, bit3 is the +// sign (1 = negative). Negative zero encodes as zero. +__device__ __forceinline__ std::uint8_t gqa_iso3_nibble(float value, float scale) { + float mag = roundf(fabsf(value) / scale); + if (mag > 7.0f) { mag = 7.0f; } + if (mag < 0.0f) { mag = 0.0f; } + std::uint8_t code = static_cast(mag); + if (value < 0.0f && code != 0) { code |= 0x08u; } + return code; +} + +__device__ __forceinline__ float gqa_iso3_decode(std::uint8_t code) { + const float mag = static_cast(code & 0x07u); + return (code & 0x08u) != 0 ? -mag : mag; +} + +// ---- native mxf4nvf4 QK staging (NVFP4 K only) ---- +// +// Q is quantized on-chip to packed E2M1 with per-(row,16-group) E4M3 scales +// and K stays packed in the cache; the block-scale mma instruction applies +// both scale vectors, so scores land in the scaled domain exactly like the +// decode kernel. The packed K tile keeps the decode kernel's 128-byte row +// layout consumed by gqa_prefill_mxf4_load_b_frag. + +constexpr float kGqaPrefillMxf4MinScale = 0.001953125f; // 2^-9, E4M3 smallest normal +constexpr std::uint8_t kGqaPrefillMxf4E4M3One = 0x38u; // E4M3FN encoding of 1.0 + +__device__ __forceinline__ void gqa_prefill_mxf4_load_a_frag(unsigned (&frag)[4], + const std::uint8_t* smem, int lane, + int k_step) { + const int row = (lane & 7) + ((lane >> 3) & 1) * 8; + const int col = (lane >> 4) * 16 + k_step * 32; + ldmatrix_x4(frag[0], frag[1], frag[2], frag[3], smem_addr(smem + row * 128 + col)); +} + +__device__ __forceinline__ void gqa_prefill_mxf4_load_b_frag(unsigned (&frag)[2], + const std::uint8_t* smem, int lane, + int n_tile, int k_step) { + const int row = (lane & 7) + n_tile * 8; + const int col = ((lane >> 3) & 1) * 16 + k_step * 32; + ldmatrix_x2(frag[0], frag[1], smem_addr(smem + row * 128 + col)); +} + +// Lane l < 4 loads its 4-channel block, applies the baked SO(4) rotation, and +// returns the rotated block in x[]. src points at the 16-d group start. +__device__ __forceinline__ void gqa_prefill_mxf4_rotate_4(float (&x)[4], + const __nv_bfloat16* src, int group, + int lane) { + if (lane < 4) { + const int block = group * 4 + lane; + const int base = lane * 4; +#pragma unroll + for (int j = 0; j < 4; ++j) { x[j] = __bfloat162float(src[base + j]); } + const float y0 = gqa_prefill_nvfp4_rot(x[0], x[1], x[2], x[3], block, 0); + const float y1 = gqa_prefill_nvfp4_rot(x[0], x[1], x[2], x[3], block, 1); + const float y2 = gqa_prefill_nvfp4_rot(x[0], x[1], x[2], x[3], block, 2); + const float y3 = gqa_prefill_nvfp4_rot(x[0], x[1], x[2], x[3], block, 3); + x[0] = y0; + x[1] = y1; + x[2] = y2; + x[3] = y3; + } else { + x[0] = x[1] = x[2] = x[3] = 0.0f; + } +} + +__device__ __forceinline__ float gqa_prefill_mxf4_group_max4(float local_max, + unsigned full_mask) { + local_max = fmaxf(local_max, __shfl_xor_sync(full_mask, local_max, 1)); + local_max = fmaxf(local_max, __shfl_xor_sync(full_mask, local_max, 2)); + return local_max; +} + +// Warm producer: copy the packed 128-byte K row and its 16 E4M3 group scales +// straight into the mxf4 staging tile (one 16-byte vector per 32 dims). +template +__device__ __forceinline__ void gqa_prefill_mxf4_stage_k_packed( + std::uint8_t* k_pk, std::uint8_t* k_sf, const std::uint8_t* cache_codes, + const std::uint8_t* cache_scales, int kv_head, int k0, int valid_start, + int max_query_abs, int physical_page, int tid) { + constexpr int Bc = kNvfp4PrefillBc; + for (int row = tid; row < Bc; row += Threads) { + const int key = k0 + row; + if (key <= max_query_abs && key >= valid_start) { + const std::int64_t scale_off = + gqa_kv_nvfp4_scale_index(physical_page, kv_head, 0, + key & kPagedKVPageMask); + store_vec(&k_sf[row * 16], load_vec(&cache_scales[scale_off])); + } else { + store_vec(&k_sf[row * 16], make_int4(0, 0, 0, 0)); + } + } + for (int chunk = tid; chunk < Bc * 8; chunk += Threads) { + const int key_l = chunk >> 3; + const int j = chunk & 7; + const int d = j * 32; + const int key = k0 + key_l; + std::uint8_t* dst = &k_pk[key_l * 128 + j * 16]; + if (key <= max_query_abs && key >= valid_start) { + const std::int64_t code_off = + gqa_kv_nvfp4_code_index(physical_page, kv_head, d, + key & kPagedKVPageMask); + store_vec(dst, load_vec(&cache_codes[code_off])); + } else { + store_vec(dst, make_int4(0, 0, 0, 0)); + } + } +} + +// Cold producer: rANS stream `tid` decodes rows (2*tid, 2*tid+1) of the packed +// 128-byte-row tile directly; all producer threads copy the slot scale tail. +template +__device__ __forceinline__ void gqa_prefill_mxf4_stage_k_cold( + std::uint8_t* k_pk, std::uint8_t* k_sf, const std::uint8_t* slot, int slot_bytes, + int half, int k0, int valid_start, int max_query_abs, int tid) { + constexpr int Bc = kNvfp4PrefillBc; + if (tid < kEntropyNvfp4SlotStreamsPerHalf) { + std::uint8_t* dst = k_pk + tid * kEntropyNvfp4SlotStreamBytes; + if (!entropy_nvfp4_slot_decode_stream(slot, half, tid, dst)) { + for (int i = 0; i < kEntropyNvfp4SlotStreamBytes; ++i) { dst[i] = 0; } + } + } + const std::uint8_t* scale_tail = entropy_nvfp4_slot_scales(slot, slot_bytes); + for (int row = tid; row < Bc; row += Threads) { + const int key = k0 + row; + if (key <= max_query_abs && key >= valid_start) { + store_vec(&k_sf[row * 16], load_vec(&scale_tail[(half * 32 + row) * 16])); + } else { + store_vec(&k_sf[row * 16], make_int4(0, 0, 0, 0)); + } + } +} + +// Producer dequant: one [Bc, D] K or V tile from the packed paged cache into a +// swizzled BF16 smem buffer. Producer threads are indexed 0..127. Sixteen dims +// are decoded per iteration: four bytes of E2M1 codes + one E4M3 scale become +// four BF16x2 pairs per 8-d swizzle block, multiplied by the group scale. +template +__device__ __forceinline__ void gqa_prefill_nvfp4_stage_kv(__nv_bfloat16* dst, + const std::uint8_t* cache_codes, + const std::uint8_t* cache_scales, + int kv_head, int k0, int valid_start, + int max_query_abs, + int physical_page, int tid) { + constexpr int D = kGqaPrefillHeadDim; + constexpr int Bc = kNvfp4PrefillBc; + constexpr int VecPerRow = D / 16; // 16 chunks of 16 dims + for (int chunk = tid; chunk < Bc * VecPerRow; chunk += Threads) { + const int key_l = chunk / VecPerRow; + const int d = (chunk - key_l * VecPerRow) << 4; + const int key = k0 + key_l; + __nv_bfloat162* p0 = reinterpret_cast<__nv_bfloat162*>( + &dst[key_l * D + gqa_prefill_swz(key_l, d)]); + __nv_bfloat162* p1 = reinterpret_cast<__nv_bfloat162*>( + &dst[key_l * D + gqa_prefill_swz(key_l, d + 8)]); + if (key <= max_query_abs && key >= valid_start) { + const int group = d >> 4; + const float scale = gqa_kv_nvfp4_e4m3_to_f32(cache_scales[ + gqa_kv_nvfp4_scale_index(physical_page, kv_head, group, + key & kPagedKVPageMask)]); + const __nv_bfloat162 scale2 = __floats2bfloat162_rn(scale, scale); + const std::uint8_t* codes = + &cache_codes[gqa_kv_nvfp4_code_index(physical_page, kv_head, d, + key & kPagedKVPageMask)]; + const uint2 raw = load_vec(codes); + const std::uint8_t* bytes = reinterpret_cast(&raw); + __nv_bfloat162 pair[8]; +#pragma unroll + for (int i = 0; i < 8; ++i) { + const unsigned lo = gqa_prefill_nvfp4_nibble_bits(bytes[i] & 0x0Fu); + const unsigned hi = gqa_prefill_nvfp4_nibble_bits(bytes[i] >> 4); + const unsigned bits = lo | (hi << 16); + pair[i] = *reinterpret_cast(&bits) * scale2; + } + store_vec(p0 + 0, make_int4(*reinterpret_cast(&pair[0]), + *reinterpret_cast(&pair[1]), + *reinterpret_cast(&pair[2]), + *reinterpret_cast(&pair[3]))); + store_vec(p1 + 0, make_int4(*reinterpret_cast(&pair[4]), + *reinterpret_cast(&pair[5]), + *reinterpret_cast(&pair[6]), + *reinterpret_cast(&pair[7]))); + } else { + store_vec(p0 + 0, make_int4(0, 0, 0, 0)); + store_vec(p1 + 0, make_int4(0, 0, 0, 0)); + } + } +} + +// Cold half-page producer: thread `stream` (0..15) decodes its 512-nibble +// rANS stream directly into the swizzled BF16 tile, applying the slot's +// uncompressed E4M3FN scales on the fly. Out-of-range rows still advance the +// rANS state but store zero. scale_tail points at the slot's 1024-byte scale +// tail (both halves). +template +__device__ __forceinline__ void gqa_prefill_nvfp4_cold_decode_kv( + __nv_bfloat16* dst, const std::uint8_t* slot, const std::uint8_t* scale_tail, int half, + int k0, int valid_start, int max_query_abs, int stream) { + std::uint8_t packed[kEntropyNvfp4SlotStreamBytes]; + if (!entropy_nvfp4_slot_decode_stream(slot, half, stream, packed)) { + for (int i = 0; i < kEntropyNvfp4SlotStreamBytes; ++i) { packed[i] = 0; } + } + for (int byte_index = 0; byte_index < kEntropyNvfp4SlotStreamBytes; ++byte_index) { + const int row_in_stream = byte_index >> 7; + const int row = 2 * stream + row_in_stream; + const int byte_in_row = byte_index & 127; + const int key = k0 + row; + const std::uint8_t byte = packed[byte_index]; +#pragma unroll + for (int nibble = 0; nibble < 2; ++nibble) { + const int dim = byte_in_row * 2 + nibble; + const std::uint8_t code = nibble == 0 ? (byte & 0x0f) : (byte >> 4); + float value = 0.0f; + if (key <= max_query_abs && key >= valid_start) { + const int group = dim >> 4; + const float scale = + gqa_kv_nvfp4_e4m3_to_f32(scale_tail[(half * 32 + row) * 16 + group]); + if constexpr (Iso3) { + value = gqa_iso3_decode(code) * scale; + } else { + // Match the warm prefill producer exactly: it dequantizes the + // packed code through gqa_prefill_nvfp4_nibble_bits and + // multiplies the BF16 value by the BF16 scale. + const unsigned bits = gqa_prefill_nvfp4_nibble_bits(code); + const float decoded = + __bfloat162float(*reinterpret_cast(&bits)); + value = decoded * scale; + } + } + dst[row * 256 + gqa_prefill_swz(row, dim)] = __float2bfloat16(value); + } + } +} + +// Producer dequant for ISO3 codes: two nibbles per byte, one E4M3FN scale per +// 16-channel group. Same 16-dim iteration, code layout, and swizzled BF16 +// output as the NVFP4 producer; only the nibble decode differs. +template +__device__ __forceinline__ void gqa_prefill_iso3_stage_kv(__nv_bfloat16* dst, + const std::uint8_t* cache_codes, + const std::uint8_t* cache_scales, + int kv_head, int k0, int max_query_abs, + int physical_page, int tid) { + constexpr int D = kGqaPrefillHeadDim; + constexpr int Bc = kNvfp4PrefillBc; + constexpr int VecPerRow = D / 16; // 16 chunks of 16 dims + for (int chunk = tid; chunk < Bc * VecPerRow; chunk += Threads) { + const int key_l = chunk / VecPerRow; + const int d = (chunk - key_l * VecPerRow) << 4; + const int key = k0 + key_l; + __nv_bfloat162* p0 = reinterpret_cast<__nv_bfloat162*>( + &dst[key_l * D + gqa_prefill_swz(key_l, d)]); + __nv_bfloat162* p1 = reinterpret_cast<__nv_bfloat162*>( + &dst[key_l * D + gqa_prefill_swz(key_l, d + 8)]); + if (key <= max_query_abs) { + const int group = d >> 4; + const float scale = gqa_kv_nvfp4_e4m3_to_f32(cache_scales[ + gqa_kv_nvfp4_scale_index(physical_page, kv_head, group, + key & kPagedKVPageMask)]); + const std::uint8_t* codes = + &cache_codes[gqa_kv_nvfp4_code_index(physical_page, kv_head, d, + key & kPagedKVPageMask)]; + const uint2 raw = load_vec(codes); + const std::uint8_t* bytes = reinterpret_cast(&raw); + __nv_bfloat162 pair[8]; +#pragma unroll + for (int i = 0; i < 8; ++i) { + const float lo = gqa_iso3_decode(bytes[i] & 0x0Fu) * scale; + const float hi = gqa_iso3_decode(bytes[i] >> 4) * scale; + pair[i] = __floats2bfloat162_rn(lo, hi); + } + store_vec(p0 + 0, make_int4(*reinterpret_cast(&pair[0]), + *reinterpret_cast(&pair[1]), + *reinterpret_cast(&pair[2]), + *reinterpret_cast(&pair[3]))); + store_vec(p1 + 0, make_int4(*reinterpret_cast(&pair[4]), + *reinterpret_cast(&pair[5]), + *reinterpret_cast(&pair[6]), + *reinterpret_cast(&pair[7]))); + } else { + store_vec(p0 + 0, make_int4(0, 0, 0, 0)); + store_vec(p1 + 0, make_int4(0, 0, 0, 0)); + } + } +} + +// Adds the second ISO3 V residual stage on top of an already-staged BF16 V +// tile. The main stage must have run first so dst holds the first-stage values. +template +__device__ __forceinline__ void gqa_prefill_iso3_stage_v_residual( + __nv_bfloat16* dst, const std::uint8_t* cache_codes, const std::uint8_t* cache_scales, + int kv_head, int k0, int max_query_abs, int physical_page, int tid) { + constexpr int D = kGqaPrefillHeadDim; + constexpr int Bc = kNvfp4PrefillBc; + constexpr int VecPerRow = D / 16; + for (int chunk = tid; chunk < Bc * VecPerRow; chunk += Threads) { + const int key_l = chunk / VecPerRow; + const int d = (chunk - key_l * VecPerRow) << 4; + const int key = k0 + key_l; + __nv_bfloat162* p0 = reinterpret_cast<__nv_bfloat162*>( + &dst[key_l * D + gqa_prefill_swz(key_l, d)]); + __nv_bfloat162* p1 = reinterpret_cast<__nv_bfloat162*>( + &dst[key_l * D + gqa_prefill_swz(key_l, d + 8)]); + if (key <= max_query_abs) { + const int group = d >> 4; + const float scale = gqa_kv_nvfp4_e4m3_to_f32(cache_scales[ + gqa_kv_nvfp4_scale_index(physical_page, kv_head, group, + key & kPagedKVPageMask)]); + const std::uint8_t* codes = + &cache_codes[gqa_kv_nvfp4_code_index(physical_page, kv_head, d, + key & kPagedKVPageMask)]; + const uint2 raw = load_vec(codes); + const std::uint8_t* bytes = reinterpret_cast(&raw); + __nv_bfloat162 pair[8]; +#pragma unroll + for (int i = 0; i < 8; ++i) { + const float lo = gqa_iso3_decode(bytes[i] & 0x0Fu) * scale; + const float hi = gqa_iso3_decode(bytes[i] >> 4) * scale; + pair[i] = __floats2bfloat162_rn(lo, hi); + } + __nv_bfloat162 cur[8]; + cur[0] = load_vec<__nv_bfloat162>(p0 + 0); + cur[1] = load_vec<__nv_bfloat162>(p0 + 1); + cur[2] = load_vec<__nv_bfloat162>(p0 + 2); + cur[3] = load_vec<__nv_bfloat162>(p0 + 3); + cur[4] = load_vec<__nv_bfloat162>(p1 + 0); + cur[5] = load_vec<__nv_bfloat162>(p1 + 1); + cur[6] = load_vec<__nv_bfloat162>(p1 + 2); + cur[7] = load_vec<__nv_bfloat162>(p1 + 3); +#pragma unroll + for (int i = 0; i < 8; ++i) { + const float lo = __bfloat162float(cur[i].x) + __bfloat162float(pair[i].x); + const float hi = __bfloat162float(cur[i].y) + __bfloat162float(pair[i].y); + pair[i] = __floats2bfloat162_rn(lo, hi); + } + store_vec(p0 + 0, make_int4(*reinterpret_cast(&pair[0]), + *reinterpret_cast(&pair[1]), + *reinterpret_cast(&pair[2]), + *reinterpret_cast(&pair[3]))); + store_vec(p1 + 0, make_int4(*reinterpret_cast(&pair[4]), + *reinterpret_cast(&pair[5]), + *reinterpret_cast(&pair[6]), + *reinterpret_cast(&pair[7]))); + } + } +} + + +template +__device__ __forceinline__ void gqa_prefill_fp8_stage_kv(__nv_bfloat16* dst, + const std::uint8_t* cache_codes, + const std::uint8_t* cache_scales, + int kv_head, int k0, int max_query_abs, + int physical_page, int tid) { + constexpr int D = kGqaPrefillHeadDim; + constexpr int Bc = kNvfp4PrefillBc; + constexpr int VecPerRow = D / 16; + for (int chunk = tid; chunk < Bc * VecPerRow; chunk += Threads) { + const int key_l = chunk / VecPerRow; + const int d = (chunk - key_l * VecPerRow) << 4; + const int key = k0 + key_l; + __nv_bfloat162* p0 = reinterpret_cast<__nv_bfloat162*>( + &dst[key_l * D + gqa_prefill_swz(key_l, d)]); + __nv_bfloat162* p1 = reinterpret_cast<__nv_bfloat162*>( + &dst[key_l * D + gqa_prefill_swz(key_l, d + 8)]); + if (key <= max_query_abs) { + const int group = d >> 4; + const float scale = gqa_kv_nvfp4_e4m3_to_f32(cache_scales[ + gqa_kv_nvfp4_scale_index(physical_page, kv_head, group, + key & kPagedKVPageMask)]); + const __nv_bfloat162 scale2 = __floats2bfloat162_rn(scale, scale); + const std::uint8_t* codes = &cache_codes[ + paged_kv_element_offset( + physical_page, kv_head, key & kPagedKVPageMask, d)]; + const uint4 raw = load_vec(codes); + const std::uint8_t* bytes = reinterpret_cast(&raw); + __nv_bfloat162 pair[8]; +#pragma unroll + for (int i = 0; i < 8; ++i) { + const float lo = gqa_kv_nvfp4_e4m3_to_f32(bytes[2 * i]) * scale; + const float hi = gqa_kv_nvfp4_e4m3_to_f32(bytes[2 * i + 1]) * scale; + pair[i] = __floats2bfloat162_rn(lo, hi); + } + store_vec(p0, make_int4(*reinterpret_cast(&pair[0]), + *reinterpret_cast(&pair[1]), + *reinterpret_cast(&pair[2]), + *reinterpret_cast(&pair[3]))); + store_vec(p1, make_int4(*reinterpret_cast(&pair[4]), + *reinterpret_cast(&pair[5]), + *reinterpret_cast(&pair[6]), + *reinterpret_cast(&pair[7]))); + } else { + store_vec(p0, make_int4(0, 0, 0, 0)); + store_vec(p1, make_int4(0, 0, 0, 0)); + } + } +} + +} // namespace + +// One warp owns one (token, kv_head, 16-d group) unit. K rotation runs lanes +// 0..3 over the four 4-channel sub-blocks; V uses all 16 lanes. +template +__launch_bounds__(256) __global__ + void gqa_attention_prefill_fill_nvfp4_kernel(const __nv_bfloat16* __restrict__ k, + const __nv_bfloat16* __restrict__ v, + const std::int32_t* __restrict__ positions, + int layer, Metadata metadata, + std::uint8_t* __restrict__ cache_k, + std::uint8_t* __restrict__ cache_v, + std::uint8_t* __restrict__ scale_k, + std::uint8_t* __restrict__ scale_v, + std::uint8_t* __restrict__ cache_k_residual, + std::uint8_t* __restrict__ scale_k_residual, + std::int32_t width) { + constexpr int Warps = 8; + constexpr unsigned FullMask = 0xffffffffu; + const int tokens = metadata.valid_tokens(width); + const int warp = static_cast(threadIdx.x) >> 5; + const int lane = static_cast(threadIdx.x) & 31; + const int unit = static_cast(blockIdx.x) * Warps + warp; + const int units = tokens * Geometry::KVHeads * kGqaKvNvfp4Groups; + if (unit >= units) { return; } + + const int group = unit % kGqaKvNvfp4Groups; + const int tmp = unit / kGqaKvNvfp4Groups; + const int kv_head = tmp % Geometry::KVHeads; + const int token = tmp / Geometry::KVHeads; + const int position = positions[0] + token; + const std::int32_t* block_table = metadata.block_table(); + int page = lane == 0 ? paged_kv_physical_page(block_table, position) : 0; + page = __shfl_sync(FullMask, page, 0); + const int page_off = position & kPagedKVPageMask; + + // ---- K: rotate + pack ---- + float kx[4] = {0.0f, 0.0f, 0.0f, 0.0f}; + if (lane < 4) { + const int block = group * 4 + lane; + const std::int64_t src = + gqa_kv_nvfp4_src_index(kv_head, group * 16, token) + lane * 4; +#pragma unroll + for (int j = 0; j < 4; ++j) { kx[j] = __bfloat162float(k[src + j]); } + const float y0 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 0); + const float y1 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 1); + const float y2 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 2); + const float y3 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 3); + kx[0] = y0; + kx[1] = y1; + kx[2] = y2; + kx[3] = y3; +#pragma unroll + for (int j = 0; j < 4; ++j) { + kx[j] *= gqa_kv_row_scale(layer, kv_head, group * 16 + lane * 4 + j); + } + } + float kmax = fmaxf(fmaxf(fabsf(kx[0]), fabsf(kx[1])), fmaxf(fabsf(kx[2]), fabsf(kx[3]))); +#pragma unroll + for (int off = 1; off <= 2; off <<= 1) { + kmax = fmaxf(kmax, __shfl_xor_sync(FullMask, kmax, off)); + } + const float kscale = fmaxf(kmax / 6.0f, 0.001953125f); + if (lane < 4) { + const std::int64_t code = + gqa_kv_nvfp4_code_index(page, kv_head, group * 16, page_off); + cache_k[code + 2 * lane] = + static_cast(gqa_kv_nvfp4_e2m1_nibble(kx[0] / kscale) | + (gqa_kv_nvfp4_e2m1_nibble(kx[1] / kscale) << 4)); + cache_k[code + 2 * lane + 1] = + static_cast(gqa_kv_nvfp4_e2m1_nibble(kx[2] / kscale) | + (gqa_kv_nvfp4_e2m1_nibble(kx[3] / kscale) << 4)); + } + if (lane == 0) { + scale_k[gqa_kv_nvfp4_scale_index(page, kv_head, group, page_off)] = + gqa_kv_nvfp4_fp32_to_e4m3(kscale); + } + + // ---- K residual: second E2M1 stage over the first-stage error ---- + if (cache_k_residual != nullptr) { + float res[4] = {0.0f, 0.0f, 0.0f, 0.0f}; + if (lane < 4) { +#pragma unroll + for (int j = 0; j < 4; ++j) { + const std::uint8_t code_j = gqa_kv_nvfp4_e2m1_nibble(kx[j] / kscale); + res[j] = kx[j] - gqa_kv_nvfp4_e2m1_to_f32(code_j) * kscale; + } + } + float rmax = fmaxf(fmaxf(fabsf(res[0]), fabsf(res[1])), + fmaxf(fabsf(res[2]), fabsf(res[3]))); +#pragma unroll + for (int off = 1; off <= 2; off <<= 1) { + rmax = fmaxf(rmax, __shfl_xor_sync(FullMask, rmax, off)); + } + const float rscale = fmaxf(rmax / 6.0f, 0.001953125f); + if (lane < 4) { + const std::int64_t rcode = + gqa_kv_nvfp4_code_index(page, kv_head, group * 16, page_off); + cache_k_residual[rcode + 2 * lane] = + static_cast(gqa_kv_nvfp4_e2m1_nibble(res[0] / rscale) | + (gqa_kv_nvfp4_e2m1_nibble(res[1] / rscale) << 4)); + cache_k_residual[rcode + 2 * lane + 1] = + static_cast(gqa_kv_nvfp4_e2m1_nibble(res[2] / rscale) | + (gqa_kv_nvfp4_e2m1_nibble(res[3] / rscale) << 4)); + } + if (lane == 0) { + scale_k_residual[gqa_kv_nvfp4_scale_index(page, kv_head, group, page_off)] = + gqa_kv_nvfp4_fp32_to_e4m3(rscale); + } + } + + // ---- V: gain-only pack ---- + const float v0 = lane < 16 ? __bfloat162float(v[gqa_kv_nvfp4_src_index( + kv_head, group * 16 + lane, token)]) + : 0.0f; + float vmax = fabsf(v0); +#pragma unroll + for (int off = 8; off > 0; off >>= 1) { + vmax = fmaxf(vmax, __shfl_xor_sync(FullMask, vmax, off)); + } + const float vscale = fmaxf(vmax / 6.0f, 0.001953125f); + if (lane < 8) { + const float ve = + __bfloat162float(v[gqa_kv_nvfp4_src_index(kv_head, group * 16 + lane * 2, + token)]); + const float vo = + __bfloat162float(v[gqa_kv_nvfp4_src_index(kv_head, group * 16 + lane * 2 + 1, + token)]); + const std::int64_t code = + gqa_kv_nvfp4_code_index(page, kv_head, group * 16, page_off); + cache_v[code + lane] = + static_cast(gqa_kv_nvfp4_e2m1_nibble(ve / vscale) | + (gqa_kv_nvfp4_e2m1_nibble(vo / vscale) << 4)); + } + if (lane == 0) { + scale_v[gqa_kv_nvfp4_scale_index(page, kv_head, group, page_off)] = + gqa_kv_nvfp4_fp32_to_e4m3(vscale); + } +} + +// ISO3 cache append: K is rotated per 4-channel block (same IsoQuant matrix as +// NVFP4), then both K and V quantize to packed sign-magnitude INT3 nibbles with +// one E4M3FN scale per 16-channel group. +template +__launch_bounds__(256) __global__ + void gqa_attention_prefill_fill_iso3_kernel(const __nv_bfloat16* __restrict__ k, + const __nv_bfloat16* __restrict__ v, + const std::int32_t* __restrict__ positions, + Metadata metadata, + std::uint8_t* __restrict__ cache_k, + std::uint8_t* __restrict__ cache_v, + std::uint8_t* __restrict__ scale_k, + std::uint8_t* __restrict__ scale_v, + std::int32_t width) { + constexpr int Warps = 8; + constexpr unsigned FullMask = 0xffffffffu; + const int tokens = metadata.valid_tokens(width); + const int warp = static_cast(threadIdx.x) >> 5; + const int lane = static_cast(threadIdx.x) & 31; + const int unit = static_cast(blockIdx.x) * Warps + warp; + const int units = tokens * Geometry::KVHeads * kGqaKvNvfp4Groups; + if (unit >= units) { return; } + + const int group = unit % kGqaKvNvfp4Groups; + const int tmp = unit / kGqaKvNvfp4Groups; + const int kv_head = tmp % Geometry::KVHeads; + const int token = tmp / Geometry::KVHeads; + const int position = positions[0] + token; + const std::int32_t* block_table = metadata.block_table(); + int page = lane == 0 ? paged_kv_physical_page(block_table, position) : 0; + page = __shfl_sync(FullMask, page, 0); + const int page_off = position & kPagedKVPageMask; + + // ---- K: rotate + pack ---- + float kx[4] = {0.0f, 0.0f, 0.0f, 0.0f}; + if (lane < 4) { + const int block = group * 4 + lane; + const std::int64_t src = + gqa_kv_nvfp4_src_index(kv_head, group * 16, token) + lane * 4; +#pragma unroll + for (int j = 0; j < 4; ++j) { kx[j] = __bfloat162float(k[src + j]); } + const float y0 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 0); + const float y1 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 1); + const float y2 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 2); + const float y3 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 3); + kx[0] = y0; + kx[1] = y1; + kx[2] = y2; + kx[3] = y3; + } + float kmax = fmaxf(fmaxf(fabsf(kx[0]), fabsf(kx[1])), fmaxf(fabsf(kx[2]), fabsf(kx[3]))); +#pragma unroll + for (int off = 1; off <= 2; off <<= 1) { + kmax = fmaxf(kmax, __shfl_xor_sync(FullMask, kmax, off)); + } + const float kscale = fmaxf(kmax / 7.0f, 0.001953125f); + if (lane < 4) { + const std::int64_t code = + gqa_kv_nvfp4_code_index(page, kv_head, group * 16, page_off); + cache_k[code + 2 * lane] = + static_cast(gqa_iso3_nibble(kx[0], kscale) | + (gqa_iso3_nibble(kx[1], kscale) << 4)); + cache_k[code + 2 * lane + 1] = + static_cast(gqa_iso3_nibble(kx[2], kscale) | + (gqa_iso3_nibble(kx[3], kscale) << 4)); + } + if (lane == 0) { + scale_k[gqa_kv_nvfp4_scale_index(page, kv_head, group, page_off)] = + gqa_kv_nvfp4_fp32_to_e4m3(kscale); + } + + // ---- V: gain-only pack ---- + const float v0 = lane < 16 ? __bfloat162float(v[gqa_kv_nvfp4_src_index( + kv_head, group * 16 + lane, token)]) + : 0.0f; + float vmax = fabsf(v0); +#pragma unroll + for (int off = 8; off > 0; off >>= 1) { + vmax = fmaxf(vmax, __shfl_xor_sync(FullMask, vmax, off)); + } + const float vscale = fmaxf(vmax / 7.0f, 0.001953125f); + if (lane < 8) { + const float ve = + __bfloat162float(v[gqa_kv_nvfp4_src_index(kv_head, group * 16 + lane * 2, + token)]); + const float vo = + __bfloat162float(v[gqa_kv_nvfp4_src_index(kv_head, group * 16 + lane * 2 + 1, + token)]); + const std::int64_t code = + gqa_kv_nvfp4_code_index(page, kv_head, group * 16, page_off); + cache_v[code + lane] = + static_cast(gqa_iso3_nibble(ve, vscale) | + (gqa_iso3_nibble(vo, vscale) << 4)); + } + if (lane == 0) { + scale_v[gqa_kv_nvfp4_scale_index(page, kv_head, group, page_off)] = + gqa_kv_nvfp4_fp32_to_e4m3(vscale); + } +} + +// Mixed cache append for the K=NVFP4 / V=ISO3 global tier: K keeps the NVFP4 +// E2M1 codec after IsoQuant rotation, V stores ISO3 sign-magnitude nibbles. +template +__launch_bounds__(256) __global__ + void gqa_attention_prefill_fill_nvfp4k_iso3v_kernel( + const __nv_bfloat16* __restrict__ k, const __nv_bfloat16* __restrict__ v, + const std::int32_t* __restrict__ positions, int layer, Metadata metadata, + std::uint8_t* __restrict__ cache_k, std::uint8_t* __restrict__ cache_v, + std::uint8_t* __restrict__ scale_k, std::uint8_t* __restrict__ scale_v, + std::uint8_t* __restrict__ cache_k_residual, std::uint8_t* __restrict__ scale_k_residual, + std::uint8_t* __restrict__ cache_v_residual, std::uint8_t* __restrict__ scale_v_residual, + std::int32_t width) { + constexpr int Warps = 8; + constexpr unsigned FullMask = 0xffffffffu; + const int tokens = metadata.valid_tokens(width); + const int warp = static_cast(threadIdx.x) >> 5; + const int lane = static_cast(threadIdx.x) & 31; + const int unit = static_cast(blockIdx.x) * Warps + warp; + const int units = tokens * Geometry::KVHeads * kGqaKvNvfp4Groups; + if (unit >= units) { return; } + + const int group = unit % kGqaKvNvfp4Groups; + const int tmp = unit / kGqaKvNvfp4Groups; + const int kv_head = tmp % Geometry::KVHeads; + const int token = tmp / Geometry::KVHeads; + const int position = positions[0] + token; + const std::int32_t* block_table = metadata.block_table(); + int page = lane == 0 ? paged_kv_physical_page(block_table, position) : 0; + page = __shfl_sync(FullMask, page, 0); + const int page_off = position & kPagedKVPageMask; + + // ---- K: rotate + NVFP4 E2M1 pack ---- + float kx[4] = {0.0f, 0.0f, 0.0f, 0.0f}; + if (lane < 4) { + const int block = group * 4 + lane; + const std::int64_t src = + gqa_kv_nvfp4_src_index(kv_head, group * 16, token) + lane * 4; +#pragma unroll + for (int j = 0; j < 4; ++j) { kx[j] = __bfloat162float(k[src + j]); } + const float y0 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 0); + const float y1 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 1); + const float y2 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 2); + const float y3 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 3); + kx[0] = y0; + kx[1] = y1; + kx[2] = y2; + kx[3] = y3; +#pragma unroll + for (int j = 0; j < 4; ++j) { + kx[j] *= gqa_kv_row_scale(layer, kv_head, group * 16 + lane * 4 + j); + } + } + float kmax = fmaxf(fmaxf(fabsf(kx[0]), fabsf(kx[1])), fmaxf(fabsf(kx[2]), fabsf(kx[3]))); +#pragma unroll + for (int off = 1; off <= 2; off <<= 1) { + kmax = fmaxf(kmax, __shfl_xor_sync(FullMask, kmax, off)); + } + const float kscale = fmaxf(kmax / 6.0f, 0.001953125f); + if (lane < 4) { + const std::int64_t code = + gqa_kv_nvfp4_code_index(page, kv_head, group * 16, page_off); + cache_k[code + 2 * lane] = + static_cast(gqa_kv_nvfp4_e2m1_nibble(kx[0] / kscale) | + (gqa_kv_nvfp4_e2m1_nibble(kx[1] / kscale) << 4)); + cache_k[code + 2 * lane + 1] = + static_cast(gqa_kv_nvfp4_e2m1_nibble(kx[2] / kscale) | + (gqa_kv_nvfp4_e2m1_nibble(kx[3] / kscale) << 4)); + } + if (lane == 0) { + scale_k[gqa_kv_nvfp4_scale_index(page, kv_head, group, page_off)] = + gqa_kv_nvfp4_fp32_to_e4m3(kscale); + } + + // ---- K residual: second E2M1 stage over the first-stage error ---- + if (cache_k_residual != nullptr) { + float res[4] = {0.0f, 0.0f, 0.0f, 0.0f}; + if (lane < 4) { +#pragma unroll + for (int j = 0; j < 4; ++j) { + const std::uint8_t code_j = gqa_kv_nvfp4_e2m1_nibble(kx[j] / kscale); + res[j] = kx[j] - gqa_kv_nvfp4_e2m1_to_f32(code_j) * kscale; + } + } + float rmax = fmaxf(fmaxf(fabsf(res[0]), fabsf(res[1])), + fmaxf(fabsf(res[2]), fabsf(res[3]))); +#pragma unroll + for (int off = 1; off <= 2; off <<= 1) { + rmax = fmaxf(rmax, __shfl_xor_sync(FullMask, rmax, off)); + } + const float rscale = fmaxf(rmax / 6.0f, 0.001953125f); + if (lane < 4) { + const std::int64_t rcode = + gqa_kv_nvfp4_code_index(page, kv_head, group * 16, page_off); + cache_k_residual[rcode + 2 * lane] = + static_cast(gqa_kv_nvfp4_e2m1_nibble(res[0] / rscale) | + (gqa_kv_nvfp4_e2m1_nibble(res[1] / rscale) << 4)); + cache_k_residual[rcode + 2 * lane + 1] = + static_cast(gqa_kv_nvfp4_e2m1_nibble(res[2] / rscale) | + (gqa_kv_nvfp4_e2m1_nibble(res[3] / rscale) << 4)); + } + if (lane == 0) { + scale_k_residual[gqa_kv_nvfp4_scale_index(page, kv_head, group, page_off)] = + gqa_kv_nvfp4_fp32_to_e4m3(rscale); + } + } + + // ---- V: gain-only ISO3 pack ---- + const float v0 = lane < 16 ? __bfloat162float(v[gqa_kv_nvfp4_src_index( + kv_head, group * 16 + lane, token)]) + : 0.0f; + float vmax = fabsf(v0); +#pragma unroll + for (int off = 8; off > 0; off >>= 1) { + vmax = fmaxf(vmax, __shfl_xor_sync(FullMask, vmax, off)); + } + const float vscale = fmaxf(vmax / 7.0f, 0.001953125f); + if (lane < 8) { + const float ve = + __bfloat162float(v[gqa_kv_nvfp4_src_index(kv_head, group * 16 + lane * 2, + token)]); + const float vo = + __bfloat162float(v[gqa_kv_nvfp4_src_index(kv_head, group * 16 + lane * 2 + 1, + token)]); + const std::int64_t code = + gqa_kv_nvfp4_code_index(page, kv_head, group * 16, page_off); + cache_v[code + lane] = + static_cast(gqa_iso3_nibble(ve, vscale) | + (gqa_iso3_nibble(vo, vscale) << 4)); + } + if (lane == 0) { + scale_v[gqa_kv_nvfp4_scale_index(page, kv_head, group, page_off)] = + gqa_kv_nvfp4_fp32_to_e4m3(vscale); + } + + // ---- V residual: second ISO3 stage over the first-stage error ---- + if (cache_v_residual != nullptr) { + float res[2] = {0.0f, 0.0f}; + float rmax = 0.0f; + if (lane < 8) { + const float ve = + __bfloat162float(v[gqa_kv_nvfp4_src_index(kv_head, group * 16 + lane * 2, + token)]); + const float vo = __bfloat162float(v[gqa_kv_nvfp4_src_index( + kv_head, group * 16 + lane * 2 + 1, token)]); + const std::uint8_t ce = gqa_iso3_nibble(ve, vscale); + const std::uint8_t co = gqa_iso3_nibble(vo, vscale); + res[0] = ve - gqa_iso3_decode(ce) * vscale; + res[1] = vo - gqa_iso3_decode(co) * vscale; + rmax = fmaxf(fabsf(res[0]), fabsf(res[1])); + } else if (lane < 16) { + const float vd = + __bfloat162float(v[gqa_kv_nvfp4_src_index(kv_head, group * 16 + lane, + token)]); + const std::uint8_t code_d = gqa_iso3_nibble(vd, vscale); + res[0] = vd - gqa_iso3_decode(code_d) * vscale; + rmax = fabsf(res[0]); + } +#pragma unroll + for (int off = 8; off > 0; off >>= 1) { + rmax = fmaxf(rmax, __shfl_xor_sync(FullMask, rmax, off)); + } + const float rvscale = fmaxf(rmax / 7.0f, 0.001953125f); + if (lane < 8) { + const std::int64_t rcode = + gqa_kv_nvfp4_code_index(page, kv_head, group * 16, page_off); + cache_v_residual[rcode + lane] = + static_cast(gqa_iso3_nibble(res[0], rvscale) | + (gqa_iso3_nibble(res[1], rvscale) << 4)); + } + if (lane == 0) { + scale_v_residual[gqa_kv_nvfp4_scale_index(page, kv_head, group, page_off)] = + gqa_kv_nvfp4_fp32_to_e4m3(rvscale); + } + } +} + +template +__launch_bounds__(256) __global__ + void gqa_attention_prefill_fill_fp8_kernel(const __nv_bfloat16* __restrict__ k, + const __nv_bfloat16* __restrict__ v, + const std::int32_t* __restrict__ positions, + Metadata metadata, + std::uint8_t* __restrict__ cache_k, + std::uint8_t* __restrict__ cache_v, + std::uint8_t* __restrict__ scale_k, + std::uint8_t* __restrict__ scale_v, + std::int32_t width) { + constexpr int Warps = 8; + constexpr unsigned FullMask = 0xffffffffu; + const int tokens = metadata.valid_tokens(width); + const int warp = static_cast(threadIdx.x) >> 5; + const int lane = static_cast(threadIdx.x) & 31; + const int unit = static_cast(blockIdx.x) * Warps + warp; + const int units = tokens * Geometry::KVHeads * kGqaKvNvfp4Groups; + if (unit >= units) { return; } + + const int group = unit % kGqaKvNvfp4Groups; + const int tmp = unit / kGqaKvNvfp4Groups; + const int kv_head = tmp % Geometry::KVHeads; + const int token = tmp / Geometry::KVHeads; + const int position = positions[0] + token; + const std::int32_t* block_table = metadata.block_table(); + int page = lane == 0 ? paged_kv_physical_page(block_table, position) : 0; + page = __shfl_sync(FullMask, page, 0); + const int page_off = position & kPagedKVPageMask; + + // ---- K: rotate + FP8 pack ---- + float kx[4] = {0.0f, 0.0f, 0.0f, 0.0f}; + if (lane < 4) { + const int block = group * 4 + lane; + const std::int64_t src = + gqa_kv_nvfp4_src_index(kv_head, group * 16, token) + lane * 4; +#pragma unroll + for (int j = 0; j < 4; ++j) { kx[j] = __bfloat162float(k[src + j]); } + const float y0 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 0); + const float y1 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 1); + const float y2 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 2); + const float y3 = gqa_prefill_nvfp4_rot(kx[0], kx[1], kx[2], kx[3], block, 3); + kx[0] = y0; + kx[1] = y1; + kx[2] = y2; + kx[3] = y3; + } + float kmax = fmaxf(fmaxf(fabsf(kx[0]), fabsf(kx[1])), fmaxf(fabsf(kx[2]), fabsf(kx[3]))); +#pragma unroll + for (int off = 1; off <= 2; off <<= 1) { + kmax = fmaxf(kmax, __shfl_xor_sync(FullMask, kmax, off)); + } + const float kscale = fmaxf(kmax / 448.0f, 0.001953125f); + if (lane < 4) { + const std::int64_t base = paged_kv_element_offset( + page, kv_head, page_off, group * 16 + lane * 4); +#pragma unroll + for (int j = 0; j < 4; ++j) { + cache_k[base + j] = gqa_kv_nvfp4_fp32_to_e4m3(kx[j] / kscale); + } + } + if (lane == 0) { + scale_k[gqa_kv_nvfp4_scale_index(page, kv_head, group, page_off)] = + gqa_kv_nvfp4_fp32_to_e4m3(kscale); + } + + // ---- V: gain-only FP8 pack ---- + const float v0 = lane < 16 ? __bfloat162float(v[gqa_kv_nvfp4_src_index( + kv_head, group * 16 + lane, token)]) + : 0.0f; + float vmax = fabsf(v0); +#pragma unroll + for (int off = 8; off > 0; off >>= 1) { + vmax = fmaxf(vmax, __shfl_xor_sync(FullMask, vmax, off)); + } + const float vscale = fmaxf(vmax / 448.0f, 0.001953125f); + if (lane < 16) { + const std::int64_t base = paged_kv_element_offset( + page, kv_head, page_off, group * 16 + lane); + cache_v[base] = gqa_kv_nvfp4_fp32_to_e4m3(v0 / vscale); + } + if (lane == 0) { + scale_v[gqa_kv_nvfp4_scale_index(page, kv_head, group, page_off)] = + gqa_kv_nvfp4_fp32_to_e4m3(vscale); + } +} + +// Warp-specialized FlashAttention-2 forward over the packed cache. Producer +// warps dequantize; consumer warps run the exact BF16 tensor-core attention +// body with Bc = 32. +template +__launch_bounds__(kNvfp4PrefillThreads, 1) __global__ + void gqa_attention_prefill_nvfp4_kernel(const __nv_bfloat16* __restrict__ q, + const std::uint8_t* __restrict__ cache_k, + const std::uint8_t* __restrict__ cache_v, + const std::uint8_t* __restrict__ cache_k_scale, + const std::uint8_t* __restrict__ cache_v_scale, + const std::uint8_t* __restrict__ cache_k_residual, + const std::uint8_t* __restrict__ cache_k_residual_scale, + const std::uint8_t* __restrict__ cache_v_residual, + const std::uint8_t* __restrict__ cache_v_residual_scale, + const std::uint8_t* __restrict__ cold_k_slots, + const std::uint8_t* __restrict__ cold_v_slots, + const std::int32_t* __restrict__ cold_k_valid, + const std::int32_t* __restrict__ cold_v_valid, + int cold_slot_bytes, int sliding_window, int layer, + Metadata metadata, + const std::int32_t* __restrict__ positions, float scale, + __nv_bfloat16* __restrict__ out, std::int32_t width) { + constexpr int D = kGqaPrefillHeadDim; + constexpr int Br = kGqaPrefillBr; // 64 + constexpr int Bc = kNvfp4PrefillBc; // 32 + constexpr int Threads = kNvfp4PrefillThreads; // 256 + constexpr int ProducerThreads = 128; + constexpr int QKNt = Bc / 8; // 4 + constexpr int QKKs = D / 16; // 16 + constexpr int PVNt = D / 8; // 32 + constexpr int PVKs = Bc / 16; // 2 + constexpr float Log2E = 1.4426950408889634074f; + constexpr unsigned FullMask = 0xffffffffu; + + static_assert(Threads == 256); + static_assert(ProducerThreads == 128); + static_assert(QKNt == 4); + static_assert(PVKs == 2); + static_assert(KVDType == DType::NVFP4 || KVDType == DType::FP8_E4M3FN || + KVDType == DType::ISO3); + static_assert(VVDType == DType::NVFP4 || VVDType == DType::FP8_E4M3FN || + VVDType == DType::ISO3); + + extern __shared__ __align__(16) std::uint8_t nvfp4_smem[]; + constexpr bool Mxf4QK = KVDType == DType::NVFP4; + constexpr int Mxf4QKKs = D / 64; + static_assert(!Mxf4QK || Mxf4QKKs == 4); + + __nv_bfloat16* q_s = nullptr; + std::uint8_t* q_a = nullptr; + std::uint8_t* q_sf = nullptr; + std::uint8_t* k_pk0 = nullptr; + std::uint8_t* k_sf0 = nullptr; + std::uint8_t* k_rpk0 = nullptr; + std::uint8_t* k_rsf0 = nullptr; + std::uint8_t* k_pk1 = nullptr; + std::uint8_t* k_sf1 = nullptr; + std::uint8_t* k_rpk1 = nullptr; + std::uint8_t* k_rsf1 = nullptr; + __nv_bfloat16* k_s0 = nullptr; + __nv_bfloat16* k_s1 = nullptr; + __nv_bfloat16* v_s0 = nullptr; + __nv_bfloat16* v_s1 = nullptr; + volatile std::uint32_t* flags = nullptr; + if constexpr (Mxf4QK) { + // Q packed E2M1 + scales, two packed 32-key K main/residual tiles, + // then the BF16 V tiles consumed by the BF16 PV body. + std::uint8_t* smem8 = nvfp4_smem; + q_a = smem8; // [Br, 128] + q_sf = q_a + Br * 128; // [Br, 16] + k_pk0 = q_sf + Br * 16; // [Bc, 128] + k_rpk0 = k_pk0 + Bc * 128; + k_sf0 = k_rpk0 + Bc * 128; // [Bc, 16] + k_rsf0 = k_sf0 + Bc * 16; + k_pk1 = k_rsf0 + Bc * 16; + k_rpk1 = k_pk1 + Bc * 128; + k_sf1 = k_rpk1 + Bc * 128; + k_rsf1 = k_sf1 + Bc * 16; + v_s0 = reinterpret_cast<__nv_bfloat16*>(k_rsf1 + Bc * 16); + v_s1 = v_s0 + Bc * D; + flags = reinterpret_cast(v_s1 + Bc * D); + } else { + q_s = reinterpret_cast<__nv_bfloat16*>(nvfp4_smem); // [Br, D] + k_s0 = q_s + Br * D; + k_s1 = k_s0 + Bc * D; + v_s0 = k_s1 + Bc * D; + v_s1 = v_s0 + Bc * D; + flags = reinterpret_cast(v_s1 + Bc * D); + } + + const int q_block = static_cast(blockIdx.x); + const int q_head = static_cast(blockIdx.y); + const int tid = static_cast(threadIdx.x); + const int warp = tid >> 5; + const int lane = tid & 31; + const int q0 = q_block * Br; + const int kv_head = q_head / Geometry::GroupSize; + const int tokens = metadata.valid_tokens(width); + + if (q_head >= Geometry::QHeads || q0 >= width) { return; } + if (q0 >= tokens) { + gqa_prefill_zero_output_rows(out, q_head, q0, min(q0 + Br, width), tid, Threads); + return; + } + const int base_pos = positions[0]; + const std::int32_t* block_table = metadata.block_table(); + + // ---- stage Q into smem once (all threads) ---- + if constexpr (Mxf4QK) { + // On-chip Q quantization: rotate each 4-channel block with the baked + // IsoQuant matrix, then pack per-16-group E2M1 with E4M3 scales. + for (int i = tid; i < Br * 128; i += Threads) { q_a[i] = 0; } + for (int i = tid; i < Br * 16; i += Threads) { q_sf[i] = kGqaPrefillMxf4E4M3One; } + __syncthreads(); + constexpr int Groups = kGqaKvNvfp4Groups; + const int q_rows = min(Br, tokens - q0); + for (int unit = warp; unit < q_rows * Groups; unit += 8) { + const int row = unit / Groups; + const int grp = unit - row * Groups; + const __nv_bfloat16* src = + q + gqa_prefill_q_index(q_head, grp * 16, q0 + row); + float qx[4]; + gqa_prefill_mxf4_rotate_4(qx, src, grp, lane); +#pragma unroll + for (int j = 0; j < 4; ++j) { + qx[j] *= gqa_kv_row_scale_inv(layer, kv_head, grp * 16 + lane * 4 + j); + } + float qmax = fmaxf(fmaxf(fabsf(qx[0]), fabsf(qx[1])), + fmaxf(fabsf(qx[2]), fabsf(qx[3]))); + qmax = gqa_prefill_mxf4_group_max4(qmax, FullMask); + const float qscale = fmaxf(qmax / 6.0f, kGqaPrefillMxf4MinScale); + if (lane < 4) { + q_a[row * 128 + grp * 8 + 2 * lane] = + static_cast(gqa_kv_nvfp4_e2m1_nibble(qx[0] / qscale) | + (gqa_kv_nvfp4_e2m1_nibble(qx[1] / qscale) << 4)); + q_a[row * 128 + grp * 8 + 2 * lane + 1] = + static_cast(gqa_kv_nvfp4_e2m1_nibble(qx[2] / qscale) | + (gqa_kv_nvfp4_e2m1_nibble(qx[3] / qscale) << 4)); + } + if (lane == 0) { + q_sf[row * 16 + grp] = gqa_kv_nvfp4_fp32_to_e4m3(qscale); + } + } + } else { + constexpr int VecPerRow = D / 8; + constexpr int QRowStride = D * Geometry::QHeads; + const __nv_bfloat16* q_block = q + gqa_prefill_q_index(q_head, 0, q0); + for (int chunk = tid; chunk < Br * VecPerRow; chunk += Threads) { + const int row = chunk / VecPerRow; + const int d = (chunk - row * VecPerRow) << 3; + __nv_bfloat16* p = &q_s[row * D + gqa_prefill_swz(row, d)]; + if (q0 + row < tokens) { + float x[8]; +#pragma unroll + for (int j = 0; j < 8; ++j) { + x[j] = __bfloat162float(q_block[row * QRowStride + d + j]); + } + gqa_prefill_nvfp4_rotate_8(x, d); + unsigned packed[4]; +#pragma unroll + for (int i = 0; i < 4; ++i) { + packed[i] = pack_bf16x2(x[2 * i], x[2 * i + 1]); + } + store_vec(p, make_int4(static_cast(packed[0]), static_cast(packed[1]), + static_cast(packed[2]), static_cast(packed[3]))); + } else { + store_vec(p, make_int4(0, 0, 0, 0)); + } + } + } + + for (int i = tid; i < 8; i += Threads) { flags[i] = 0; } + if (tid == 1) { flags[1] = 1; } // K slot 0 free + if (tid == 3) { flags[3] = 1; } // K slot 1 free + if (tid == 5) { flags[5] = 1; } // V slot 0 free + if (tid == 7) { flags[7] = 1; } // V slot 1 free + __syncthreads(); + + const int tile_rows = min(Br, tokens - q0); + const int max_query_abs = base_pos + q0 + tile_rows - 1; + const int window = (sliding_window > 0 && KVDType == DType::NVFP4) ? sliding_window : 0; + const int visible_start = window > 0 ? max(0, base_pos + q0 - window + 1) : 0; + const int kb_start = visible_start / (2 * Bc); + const int n_block64 = (max_query_abs / (2 * Bc)) + 1 - kb_start; + const float scale_l2 = scale * Log2E; + + if (warp >= 4) { + // ---- producer: stage packed K and dequantized V sub-tiles into + // ping-pong smem buffers. Named barrier 0 is the full-CTA handshake; + // producer threads first decode any cold slot half-page, synchronized + // by producer-only named barrier 1. ---- + const int ptid = tid - ProducerThreads; + const auto stage_v = [&](__nv_bfloat16* v_s, int k0i, int page) { + if constexpr (VVDType == DType::FP8_E4M3FN) { + gqa_prefill_fp8_stage_kv( + v_s, cache_v, cache_v_scale, kv_head, k0i, max_query_abs, page, ptid); + } else if constexpr (VVDType == DType::ISO3) { + gqa_prefill_iso3_stage_kv( + v_s, cache_v, cache_v_scale, kv_head, k0i, max_query_abs, page, ptid); + if (cache_v_residual != nullptr) { + gqa_prefill_iso3_stage_v_residual( + v_s, cache_v_residual, cache_v_residual_scale, kv_head, k0i, + max_query_abs, page, ptid); + } + } else { + gqa_prefill_nvfp4_stage_kv( + v_s, cache_v, cache_v_scale, kv_head, k0i, visible_start, max_query_abs, page, + ptid); + } + }; + const auto stage_k_bf16 = [&](__nv_bfloat16* k_s, int k0i, int page) { + if constexpr (KVDType == DType::FP8_E4M3FN) { + gqa_prefill_fp8_stage_kv( + k_s, cache_k, cache_k_scale, kv_head, k0i, max_query_abs, page, ptid); + } else if constexpr (KVDType == DType::ISO3) { + gqa_prefill_iso3_stage_kv( + k_s, cache_k, cache_k_scale, kv_head, k0i, max_query_abs, page, ptid); + } else { + gqa_prefill_nvfp4_stage_kv( + k_s, cache_k, cache_k_scale, kv_head, k0i, visible_start, max_query_abs, page, + ptid); + } + }; + const auto stage_k_cold_bf16 = [&](__nv_bfloat16* k_s, const std::uint8_t* k_slot, + int half, int k0i) { + if (ptid < kEntropyNvfp4SlotStreamsPerHalf) { + gqa_prefill_nvfp4_cold_decode_kv( + k_s, k_slot, entropy_nvfp4_slot_scales(k_slot, cold_slot_bytes), half, k0i, + visible_start, max_query_abs, ptid); + } + }; + for (int kb = 0; kb < n_block64; ++kb) { + const int kb64 = kb_start + kb; + const int k0 = kb64 * 2 * Bc; + const int table_entry = block_table[kb64]; + const bool cold_available = table_entry <= -2 && cold_k_slots != nullptr && + cold_v_slots != nullptr && cold_k_valid != nullptr && + cold_v_valid != nullptr && cold_slot_bytes >= 1024 + 320; + const int slot_base = cold_available ? -table_entry - 2 : 0; + // Region-relative flat slot index: slot * 2*KVHeads + head; the V + // plane's valid entries sit one KVHeads block later. + const int cold_slot_id = slot_base * (2 * Geometry::KVHeads) + kv_head; + const bool cold = cold_available && cold_k_valid[cold_slot_id] != 0 && + cold_v_valid[cold_slot_id + Geometry::KVHeads] != 0; + const int physical_page = cold ? 0 : table_entry; + const std::uint8_t* k_slot = + cold ? cold_k_slots + static_cast(cold_slot_id) * cold_slot_bytes + : nullptr; + const std::uint8_t* v_slot = + cold ? cold_v_slots + static_cast(cold_slot_id) * cold_slot_bytes + : nullptr; + + // ---- half 0 (slot 0) ---- + if constexpr (Mxf4QK) { + if (cold) { + gqa_prefill_mxf4_stage_k_cold( + k_pk0, k_sf0, k_slot, cold_slot_bytes, 0, k0, visible_start, + max_query_abs, ptid); + for (int chunk = ptid; chunk < Bc * 8; chunk += ProducerThreads) { + const int key_l = chunk >> 3; + const int j = chunk & 7; + store_vec(&k_rpk0[key_l * 128 + j * 16], make_int4(0, 0, 0, 0)); + } + for (int row = ptid; row < Bc; row += ProducerThreads) { + store_vec(&k_rsf0[row * 16], make_int4(0, 0, 0, 0)); + } + } else { + gqa_prefill_mxf4_stage_k_packed( + k_pk0, k_sf0, cache_k, cache_k_scale, kv_head, k0, visible_start, + max_query_abs, physical_page, ptid); + if (cache_k_residual != nullptr) { + gqa_prefill_mxf4_stage_k_packed( + k_rpk0, k_rsf0, cache_k_residual, cache_k_residual_scale, kv_head, k0, + visible_start, max_query_abs, physical_page, ptid); + } else { + for (int chunk = ptid; chunk < Bc * 8; chunk += ProducerThreads) { + const int key_l = chunk >> 3; + const int j = chunk & 7; + store_vec(&k_rpk0[key_l * 128 + j * 16], make_int4(0, 0, 0, 0)); + } + for (int row = ptid; row < Bc; row += ProducerThreads) { + store_vec(&k_rsf0[row * 16], make_int4(0, 0, 0, 0)); + } + } + } + } else { + if (cold) { + stage_k_cold_bf16(k_s0, k_slot, 0, k0); + } else { + stage_k_bf16(k_s0, k0, physical_page); + } + } + if (cold) { + if (ptid >= kEntropyNvfp4SlotStreamsPerHalf && + ptid < 2 * kEntropyNvfp4SlotStreamsPerHalf) { + gqa_prefill_nvfp4_cold_decode_kv( + v_s0, v_slot, entropy_nvfp4_slot_scales(v_slot, cold_slot_bytes), 0, k0, + visible_start, max_query_abs, ptid - kEntropyNvfp4SlotStreamsPerHalf); + } + gqa_prefill_bar_sync(1, ProducerThreads); + } else { + stage_v(v_s0, k0, physical_page); + } + gqa_prefill_bar_sync(0, Threads); + + // ---- half 1 (slot 1) ---- + if constexpr (Mxf4QK) { + if (cold) { + gqa_prefill_mxf4_stage_k_cold( + k_pk1, k_sf1, k_slot, cold_slot_bytes, 1, k0 + Bc, visible_start, + max_query_abs, ptid); + for (int chunk = ptid; chunk < Bc * 8; chunk += ProducerThreads) { + const int key_l = chunk >> 3; + const int j = chunk & 7; + store_vec(&k_rpk1[key_l * 128 + j * 16], make_int4(0, 0, 0, 0)); + } + for (int row = ptid; row < Bc; row += ProducerThreads) { + store_vec(&k_rsf1[row * 16], make_int4(0, 0, 0, 0)); + } + } else { + gqa_prefill_mxf4_stage_k_packed( + k_pk1, k_sf1, cache_k, cache_k_scale, kv_head, k0 + Bc, visible_start, + max_query_abs, physical_page, ptid); + if (cache_k_residual != nullptr) { + gqa_prefill_mxf4_stage_k_packed( + k_rpk1, k_rsf1, cache_k_residual, cache_k_residual_scale, kv_head, + k0 + Bc, visible_start, max_query_abs, physical_page, ptid); + } else { + for (int chunk = ptid; chunk < Bc * 8; chunk += ProducerThreads) { + const int key_l = chunk >> 3; + const int j = chunk & 7; + store_vec(&k_rpk1[key_l * 128 + j * 16], make_int4(0, 0, 0, 0)); + } + for (int row = ptid; row < Bc; row += ProducerThreads) { + store_vec(&k_rsf1[row * 16], make_int4(0, 0, 0, 0)); + } + } + } + } else { + if (cold) { + stage_k_cold_bf16(k_s1, k_slot, 1, k0 + Bc); + } else { + stage_k_bf16(k_s1, k0 + Bc, physical_page); + } + } + if (cold) { + if (ptid >= kEntropyNvfp4SlotStreamsPerHalf && + ptid < 2 * kEntropyNvfp4SlotStreamsPerHalf) { + gqa_prefill_nvfp4_cold_decode_kv( + v_s1, v_slot, entropy_nvfp4_slot_scales(v_slot, cold_slot_bytes), 1, + k0 + Bc, visible_start, max_query_abs, + ptid - kEntropyNvfp4SlotStreamsPerHalf); + } + gqa_prefill_bar_sync(1, ProducerThreads); + } else { + stage_v(v_s1, k0 + Bc, physical_page); + } + gqa_prefill_bar_sync(0, Threads); + + gqa_prefill_bar_sync(0, Threads); + } + return; + } + + // ---- consumer: exact BF16 FlashAttention body over the dequantized tiles ---- + const int gid = lane >> 2; + const int lid = lane & 3; + + const int b_rin = lane & 7; + const int warp_row0 = warp * 16; + + const unsigned v_as = static_cast((lane >> 4) << 4); + const unsigned v_r = static_cast(b_rin << 4); + + float acc[PVNt][4]; +#pragma unroll + for (int n = 0; n < PVNt; ++n) { +#pragma unroll + for (int i = 0; i < 4; ++i) { acc[n][i] = 0.0f; } + } + float m0 = -CUDART_INF_F, m1 = -CUDART_INF_F, l0 = 0.0f, l1 = 0.0f; + + constexpr int QKNt64 = 8; // 64-key score n-tiles + constexpr int PVKs64 = 4; // 64-key PV contraction groups + + const auto qk_half_mxf4 = [&](const std::uint8_t* k_pk, const std::uint8_t* k_sf, + const std::uint8_t* k_rpk, const std::uint8_t* k_rsf, + float (&score)[QKNt][4]) { +#pragma unroll + for (int nt = 0; nt < QKNt; ++nt) { + score[nt][0] = score[nt][1] = score[nt][2] = score[nt][3] = 0.0f; + } +#pragma unroll + for (int k = 0; k < Mxf4QKKs; ++k) { + unsigned af[4]; + gqa_prefill_mxf4_load_a_frag(af, q_a + warp_row0 * 128, lane, k); + const unsigned sfa = load_vec( + q_sf + warp_row0 * 16 + (gid + (lid & 1) * 8) * 16 + k * 4); +#pragma unroll + for (int nt = 0; nt < QKNt; ++nt) { + unsigned bf[2]; + gqa_prefill_mxf4_load_b_frag(bf, k_pk, lane, nt, k); + const unsigned sfb = load_vec(k_sf + (gid + nt * 8) * 16 + k * 4); + mma_nvfp4_e4m3(score[nt][0], score[nt][1], score[nt][2], score[nt][3], + af[0], af[1], af[2], af[3], bf[0], bf[1], sfa, sfb); + } + } + // Second pass accumulates the E2M1 residual K plane. +#pragma unroll + for (int k = 0; k < Mxf4QKKs; ++k) { + unsigned af[4]; + gqa_prefill_mxf4_load_a_frag(af, q_a + warp_row0 * 128, lane, k); + const unsigned sfa = load_vec( + q_sf + warp_row0 * 16 + (gid + (lid & 1) * 8) * 16 + k * 4); +#pragma unroll + for (int nt = 0; nt < QKNt; ++nt) { + unsigned bf[2]; + gqa_prefill_mxf4_load_b_frag(bf, k_rpk, lane, nt, k); + const unsigned sfb = load_vec(k_rsf + (gid + nt * 8) * 16 + k * 4); + mma_nvfp4_e4m3(score[nt][0], score[nt][1], score[nt][2], score[nt][3], + af[0], af[1], af[2], af[3], bf[0], bf[1], sfa, sfb); + } + } + }; + + const auto qk_half_bf16 = [&](const __nv_bfloat16* k_s, float (&score)[QKNt][4]) { + const int a_mat = lane >> 3; + const int a_rin = lane & 7; + const int a_rowoff = a_rin + ((a_mat & 1) << 3); + const int b_koff = ((lane >> 3) & 1) << 3; + const unsigned q_sbase = smem_addr(q_s); + const unsigned q_lane_base = + q_sbase + static_cast((warp_row0 + a_rowoff) * 512); + const unsigned q_as = static_cast((a_mat >> 1) << 4); + const unsigned q_r = static_cast(a_rin << 4); + const unsigned k_as = static_cast((b_koff >> 3) << 4); + const unsigned k_r = static_cast(b_rin << 4); + const unsigned k_sbase = smem_addr(k_s); + const unsigned k_lane_base = + k_sbase + static_cast(b_rin * 512) + + (static_cast(lane >> 4) << 12); +#pragma unroll + for (int nt = 0; nt < QKNt; ++nt) { + score[nt][0] = score[nt][1] = score[nt][2] = score[nt][3] = 0.0f; + } + unsigned af[2][4]; + unsigned bf[2][QKNt][2]; + { + ldmatrix_x4(af[0][0], af[0][1], af[0][2], af[0][3], + gqa_prefill_swz_addr(q_lane_base, 0u, q_as, q_r)); +#pragma unroll + for (int nt2 = 0; nt2 < QKNt; nt2 += 2) { + ldmatrix_x4(bf[0][nt2][0], bf[0][nt2][1], bf[0][nt2 + 1][0], bf[0][nt2 + 1][1], + gqa_prefill_swz_addr( + k_lane_base + static_cast(nt2 * 4096), 0u, k_as, k_r)); + } + } +#pragma unroll + for (int k = 0; k < QKKs; ++k) { + const int cur = k & 1; + const int nxt = cur ^ 1; + if (k + 1 < QKKs) { + const unsigned ck = static_cast((k + 1) << 5); + ldmatrix_x4(af[nxt][0], af[nxt][1], af[nxt][2], af[nxt][3], + gqa_prefill_swz_addr(q_lane_base, ck, q_as, q_r)); +#pragma unroll + for (int nt2 = 0; nt2 < QKNt; nt2 += 2) { + ldmatrix_x4( + bf[nxt][nt2][0], bf[nxt][nt2][1], bf[nxt][nt2 + 1][0], + bf[nxt][nt2 + 1][1], + gqa_prefill_swz_addr( + k_lane_base + static_cast(nt2 * 4096), ck, k_as, k_r)); + } + } +#pragma unroll + for (int nt = 0; nt < QKNt; ++nt) { + mma_bf16(score[nt][0], score[nt][1], score[nt][2], score[nt][3], af[cur][0], + af[cur][1], af[cur][2], af[cur][3], bf[cur][nt][0], bf[cur][nt][1]); + } + } + }; + + for (int kb = 0; kb < n_block64; ++kb) { + const int k0 = (kb_start + kb) * 2 * Bc; + + // ---- QK^T over the two 32-key halves, then one 64-key softmax ---- + gqa_prefill_bar_sync(0, Threads); // slot 0 staged by producers + float score_a[QKNt][4]; + if constexpr (Mxf4QK) { + qk_half_mxf4(k_pk0, k_sf0, k_rpk0, k_rsf0, score_a); + } else { + qk_half_bf16(k_s0, score_a); + } + + gqa_prefill_bar_sync(0, Threads); // slot 1 staged; slot 0 read done + float score_b[QKNt][4]; + if constexpr (Mxf4QK) { + qk_half_mxf4(k_pk1, k_sf1, k_rpk1, k_rsf1, score_b); + } else { + qk_half_bf16(k_s1, score_b); + } + + float score[QKNt64][4]; +#pragma unroll + for (int nt = 0; nt < QKNt; ++nt) { + score[nt][0] = score_a[nt][0]; + score[nt][1] = score_a[nt][1]; + score[nt][2] = score_a[nt][2]; + score[nt][3] = score_a[nt][3]; + score[QKNt + nt][0] = score_b[nt][0]; + score[QKNt + nt][1] = score_b[nt][1]; + score[QKNt + nt][2] = score_b[nt][2]; + score[QKNt + nt][3] = score_b[nt][3]; + } + + const int row0 = warp_row0 + gid; + const int row1 = warp_row0 + gid + 8; + const int qrow0 = q0 + row0; + const int qrow1 = q0 + row1; + const int qabs0 = (qrow0 < tokens) ? base_pos + qrow0 : -1; + const int qabs1 = (qrow1 < tokens) ? base_pos + qrow1 : -1; + const bool full_score_tile = + (q0 + Br <= tokens) && ((k0 + 2 * Bc - 1) <= (base_pos + q0)) && + (window == 0 || k0 >= max(0, max_query_abs - window + 1)); + + float bm0 = -CUDART_INF_F, bm1 = -CUDART_INF_F; + if (full_score_tile) { +#pragma unroll + for (int nt = 0; nt < QKNt64; ++nt) { + bm0 = fmaxf(bm0, fmaxf(score[nt][0], score[nt][1])); + bm1 = fmaxf(bm1, fmaxf(score[nt][2], score[nt][3])); + } + } else { +#pragma unroll + for (int nt = 0; nt < QKNt64; ++nt) { + const int key0 = k0 + nt * 8 + 2 * lid; + const int key1 = key0 + 1; + const int row0_start = (window > 0 && qabs0 >= 0) ? max(0, qabs0 - window + 1) : 0; + const int row1_start = (window > 0 && qabs1 >= 0) ? max(0, qabs1 - window + 1) : 0; + score[nt][0] = (qrow0 < tokens && key0 <= qabs0 && key0 >= row0_start) + ? score[nt][0] + : -CUDART_INF_F; + score[nt][1] = (qrow0 < tokens && key1 <= qabs0 && key1 >= row0_start) + ? score[nt][1] + : -CUDART_INF_F; + score[nt][2] = (qrow1 < tokens && key0 <= qabs1 && key0 >= row1_start) + ? score[nt][2] + : -CUDART_INF_F; + score[nt][3] = (qrow1 < tokens && key1 <= qabs1 && key1 >= row1_start) + ? score[nt][3] + : -CUDART_INF_F; + bm0 = fmaxf(bm0, fmaxf(score[nt][0], score[nt][1])); + bm1 = fmaxf(bm1, fmaxf(score[nt][2], score[nt][3])); + } + } + bm0 = warp_max<4>(bm0, FullMask); + bm1 = warp_max<4>(bm1, FullMask); + + const float nm0 = fmaxf(m0, bm0); + const float nm1 = fmaxf(m1, bm1); + const float nm0_scaled = nm0 * scale_l2; + const float nm1_scaled = nm1 * scale_l2; + const float alpha0 = exp2_approx(__fmaf_rn(m0, scale_l2, -nm0_scaled)); + const float alpha1 = exp2_approx(__fmaf_rn(m1, scale_l2, -nm1_scaled)); + + float bl0 = 0.0f, bl1 = 0.0f; + unsigned p_frag[PVKs64][4]; + if (full_score_tile) { +#pragma unroll + for (int nt = 0; nt < QKNt64; ++nt) { + const float p00 = exp2_approx(__fmaf_rn(score[nt][0], scale_l2, -nm0_scaled)); + const float p01 = exp2_approx(__fmaf_rn(score[nt][1], scale_l2, -nm0_scaled)); + const float p10 = exp2_approx(__fmaf_rn(score[nt][2], scale_l2, -nm1_scaled)); + const float p11 = exp2_approx(__fmaf_rn(score[nt][3], scale_l2, -nm1_scaled)); + bl0 += p00 + p01; + bl1 += p10 + p11; + const int pk = nt >> 1; + if ((nt & 1) == 0) { + p_frag[pk][0] = pack_bf16x2(p00, p01); + p_frag[pk][1] = pack_bf16x2(p10, p11); + } else { + p_frag[pk][2] = pack_bf16x2(p00, p01); + p_frag[pk][3] = pack_bf16x2(p10, p11); + } + } + } else { +#pragma unroll + for (int nt = 0; nt < QKNt64; ++nt) { + const float p00 = (score[nt][0] > -CUDART_INF_F) + ? exp2_approx(__fmaf_rn(score[nt][0], scale_l2, -nm0_scaled)) + : 0.0f; + const float p01 = (score[nt][1] > -CUDART_INF_F) + ? exp2_approx(__fmaf_rn(score[nt][1], scale_l2, -nm0_scaled)) + : 0.0f; + const float p10 = (score[nt][2] > -CUDART_INF_F) + ? exp2_approx(__fmaf_rn(score[nt][2], scale_l2, -nm1_scaled)) + : 0.0f; + const float p11 = (score[nt][3] > -CUDART_INF_F) + ? exp2_approx(__fmaf_rn(score[nt][3], scale_l2, -nm1_scaled)) + : 0.0f; + bl0 += p00 + p01; + bl1 += p10 + p11; + const int pk = nt >> 1; + if ((nt & 1) == 0) { + p_frag[pk][0] = pack_bf16x2(p00, p01); + p_frag[pk][1] = pack_bf16x2(p10, p11); + } else { + p_frag[pk][2] = pack_bf16x2(p00, p01); + p_frag[pk][3] = pack_bf16x2(p10, p11); + } + } + } + + l0 = __fmaf_rn(l0, alpha0, bl0); + l1 = __fmaf_rn(l1, alpha1, bl1); + m0 = nm0; + m1 = nm1; +#pragma unroll + for (int n = 0; n < PVNt; ++n) { + acc[n][0] *= alpha0; + acc[n][1] *= alpha0; + acc[n][2] *= alpha1; + acc[n][3] *= alpha1; + } + + // ---- O += P V over the two 32-key V halves ---- + constexpr int PVHalf = PVNt / 2; + constexpr int PVLoads = PVKs * PVHalf; +#pragma unroll + for (int half = 0; half < 2; ++half) { + const __nv_bfloat16* v_s = half == 0 ? v_s0 : v_s1; + const unsigned v_sbase = smem_addr(v_s); + const unsigned v_lane_base = + v_sbase + static_cast(((lane >> 3) & 1) * 4096) + + static_cast(b_rin * 512); + unsigned vf[2][4]; + { + ldmatrix_x4_t(vf[0][0], vf[0][1], vf[0][2], vf[0][3], + gqa_prefill_swz_addr(v_lane_base, 0u, v_as, v_r)); + } +#pragma unroll + for (int li = 0; li < PVLoads; ++li) { + const int k = li / PVHalf; + const int n2 = (li % PVHalf) * 2; + const int cur = li & 1; + const int nxt = cur ^ 1; + if (li + 1 < PVLoads) { + const int k2 = (li + 1) / PVHalf; + const int n2b = ((li + 1) % PVHalf) * 2; + const unsigned ckv = static_cast(n2b << 4); + ldmatrix_x4_t(vf[nxt][0], vf[nxt][1], vf[nxt][2], vf[nxt][3], + gqa_prefill_swz_addr( + v_lane_base + static_cast(k2 * 8192), ckv, v_as, + v_r)); + } + const int pk = half * PVKs + k; + mma_bf16(acc[n2][0], acc[n2][1], acc[n2][2], acc[n2][3], p_frag[pk][0], + p_frag[pk][1], p_frag[pk][2], p_frag[pk][3], vf[cur][0], vf[cur][1]); + mma_bf16(acc[n2 + 1][0], acc[n2 + 1][1], acc[n2 + 1][2], acc[n2 + 1][3], + p_frag[pk][0], p_frag[pk][1], p_frag[pk][2], p_frag[pk][3], vf[cur][2], + vf[cur][3]); + } + } + gqa_prefill_bar_sync(0, Threads); // both halves consumed; buffers reusable + } + + l0 = warp_sum<4>(l0, FullMask); + l1 = warp_sum<4>(l1, FullMask); + + const float inv_l0 = (l0 > 0.0f) ? __frcp_rn(l0) : 0.0f; + const float inv_l1 = (l1 > 0.0f) ? __frcp_rn(l1) : 0.0f; +#pragma unroll + for (int n = 0; n < PVNt; ++n) { + const int d0 = n * 8 + 2 * lid; + const int qrow0 = q0 + warp_row0 + gid; + const int qrow1 = q0 + warp_row0 + gid + 8; + if (qrow0 < tokens) { + *reinterpret_cast(&out[gqa_prefill_q_index(q_head, d0, qrow0)]) = + pack_bf16x2(acc[n][0] * inv_l0, acc[n][1] * inv_l0); + } + if (qrow1 < tokens) { + *reinterpret_cast(&out[gqa_prefill_q_index(q_head, d0, qrow1)]) = + pack_bf16x2(acc[n][2] * inv_l1, acc[n][3] * inv_l1); + } + } + gqa_prefill_zero_output_rows(out, q_head, tokens, min(q0 + Br, width), tid, + ProducerThreads); +} + +} // namespace ninfer::ops diff --git a/src/ops/kernel/gqa_isoquant_rot.cuh b/src/ops/kernel/gqa_isoquant_rot.cuh new file mode 100644 index 0000000000..2ffd3533ee --- /dev/null +++ b/src/ops/kernel/gqa_isoquant_rot.cuh @@ -0,0 +1,19 @@ +#pragma once + +// Baked IsoQuant per-4-channel SO(4) rotations, [64][4][4] fp32, +// imported from the nvfp4rtx offline calibration (isoquant_rot.npy). +// The table lives in constant memory (gqa_isoquant_rot.cu) so kernels index +// it with LDC; a function-local constexpr copy previously made nvcc expand +// the whole table into the hot Q/K quantization loops and spill it through +// the per-thread stack. Applied to K on cache write and to Q before NVFP4 +// quantization so QK^T is preserved in the rotated domain. + +extern __constant__ float kGqaIsoquantRotDev[64][4][4]; + +namespace ninfer::ops { + +__device__ __forceinline__ float gqa_isoquant_rot_value(int block, int row, int col) { + return ::kGqaIsoquantRotDev[block][row][col]; +} + +} // namespace ninfer::ops diff --git a/src/ops/kernel/gqa_isoquant_row_scale.cuh b/src/ops/kernel/gqa_isoquant_row_scale.cuh new file mode 100644 index 0000000000..d2c9387360 --- /dev/null +++ b/src/ops/kernel/gqa_isoquant_row_scale.cuh @@ -0,0 +1,31 @@ +#pragma once + +// Sinkhorn-constrained row scales for the rotated NVFP4 K domain. +// +// For every full-attention (layer, kv_head), a token-independent per-channel +// scale s_d in [0.5, 2.0] balances rotated K row RMS before E4M3/E2M1 +// quantization. K is multiplied by s_d on cache write; Q is multiplied by +// 1/s_d before QK, so QK^T is preserved. This mainly protects low-energy +// channels whose E4M3 group scale would otherwise collapse to denormals. +// +// The table is baked from kvcalib-a and stored as BF16 words in constant +// memory (16 * 4 * 256 * 2 = 32 KiB, together with the SO(4) rotation table). + +#include + +#include + +extern __constant__ unsigned short kGqaKvRowScaleDev[16][4][256]; + +namespace ninfer::ops { + +__device__ __forceinline__ float gqa_kv_row_scale(int layer, int kv_head, int d) { + const unsigned short raw = ::kGqaKvRowScaleDev[layer][kv_head][d]; + return __bfloat162float(*reinterpret_cast(&raw)); +} + +__device__ __forceinline__ float gqa_kv_row_scale_inv(int layer, int kv_head, int d) { + return 1.0f / gqa_kv_row_scale(layer, kv_head, d); +} + +} // namespace ninfer::ops From 562b5dc8ac3092f940079e5ec9ebef82f4ddaa88 Mon Sep 17 00:00:00 2001 From: NInfer Agent Date: Sun, 30 Aug 2026 22:11:09 +0800 Subject: [PATCH 08/11] feat(runtime): cold-compress pass at the decode boundary + build wiring --- src/ops/kernel/entropy_cold_requant_kernels.cuh | 15 ++++++++++++++- 1 file changed, 14 insertions(+), 1 deletion(-) diff --git a/src/ops/kernel/entropy_cold_requant_kernels.cuh b/src/ops/kernel/entropy_cold_requant_kernels.cuh index 8269f69028..6c8d52b273 100644 --- a/src/ops/kernel/entropy_cold_requant_kernels.cuh +++ b/src/ops/kernel/entropy_cold_requant_kernels.cuh @@ -20,7 +20,20 @@ #include "ops/kernel/gqa_attention_kv_nvfp4.cuh" #include "ops/kernel/gqa_attention_kv_quant.cuh" -#include "ops/kernel/gqa_attention_prefill_nvfp4.cuh" // gqa_iso3_nibble / gqa_iso3_decode +// ISO3 = sign-magnitude INT3: low 3 bits magnitude 0..7, bit 3 sign. +// Inlined here to keep the op header dependency-free. +__device__ __forceinline__ std::uint8_t gqa_iso3_nibble(float value, float scale) { + float mag = roundf(fabsf(value) / scale); + if (mag > 7.0f) { mag = 7.0f; } + if (mag < 0.0f) { mag = 0.0f; } + std::uint8_t code = static_cast(mag); + if (value < 0.0f && code != 0) { code |= 0x08u; } + return code; +} +__device__ __forceinline__ float gqa_iso3_decode(std::uint8_t code) { + const float mag = static_cast(code & 0x07u); + return (code & 0x08u) != 0 ? -mag : mag; +} #include "ops/launcher/entropy_cold_requant.h" #include From 17f56aff12fb655f893c3e6adaec281604f67d14 Mon Sep 17 00:00:00 2001 From: NInfer Agent Date: Sun, 30 Aug 2026 22:23:28 +0800 Subject: [PATCH 09/11] feat(kv): complete cold-pool mechanism on the paged KV store Cold slots are allocated as per-layer regions (9232 B raw slots + I32 validity) by the decoder state; the pool exposes allocate/ release with a used bitmap. A decode-boundary pass packs the retired tail (valid - cold_keep_tokens) of a sequence's text KV into raw entropy slots, shrinks the address space entitlement, returns the physical pages to the pool, and publishes block-table sentinels (entry <= -2, slot base = -2 - entry) via publish_indices. Attention producers decode sentinel pages inline from the slots in the cold staging branches (INT8 adapter preserves int8 QK cores; NVFP4 tier keeps its native E2M1/ISO3 nibble semantics). --- src/core/paged_kv_cache.h | 18 ++++ .../ninfer/targets/qwen3_6/decoder_state.h | 19 ++++ src/targets/qwen3_6/impl/runtime/layouts.h | 3 + .../qwen3_6/impl/runtime/layouts_impl.h | 9 ++ src/targets/qwen3_6/impl/runtime/program.h | 9 ++ .../qwen3_6/impl/runtime/program_impl.h | 100 ++++++++++++++++++ .../qwen3_6/impl/state/decoder_state.cpp | 46 +++++++- 7 files changed, 203 insertions(+), 1 deletion(-) diff --git a/src/core/paged_kv_cache.h b/src/core/paged_kv_cache.h index 9a4e079148..f59f3242eb 100644 --- a/src/core/paged_kv_cache.h +++ b/src/core/paged_kv_cache.h @@ -16,6 +16,21 @@ namespace ninfer { inline constexpr std::int32_t kPagedKVPageSize = 64; +// Cold-slot sentinel encoding in block tables: entries <= -2 address the +// entropy pool (slot base = -2 - entry). The pool itself is per-layer and +// owned by the decoder state; the table row only carries the slot base. +inline constexpr std::int32_t kPagedKVColdSentinelBase = -2; + +[[nodiscard]] inline bool paged_kv_is_cold(std::int32_t entry) noexcept { + return entry <= kPagedKVColdSentinelBase; +} +[[nodiscard]] inline std::int32_t paged_kv_cold_slot_base(std::int32_t entry) noexcept { + return kPagedKVColdSentinelBase - entry; +} +[[nodiscard]] inline std::int32_t paged_kv_cold_entry(std::int32_t slot_base) noexcept { + return kPagedKVColdSentinelBase - slot_base; +} + /** Non-owning, single-sequence view consumed by growing-cache Ops. */ struct PagedKVLayerView { Tensor k_pages; @@ -36,6 +51,9 @@ struct PagedKVBatchLayerView { Tensor k_scale_pages; Tensor v_scale_pages; Tensor block_tables; + // Entropy-coded cold pool: fixed raw slots + validity flags per layer. + Tensor cold_slots; + Tensor cold_slot_valid; std::int32_t head_dim = 0; std::int32_t num_kv_heads = 0; DType dtype = DType::BF16; diff --git a/src/targets/qwen3_6/export/ninfer/targets/qwen3_6/decoder_state.h b/src/targets/qwen3_6/export/ninfer/targets/qwen3_6/decoder_state.h index f1193f588a..7fded328b3 100644 --- a/src/targets/qwen3_6/export/ninfer/targets/qwen3_6/decoder_state.h +++ b/src/targets/qwen3_6/export/ninfer/targets/qwen3_6/decoder_state.h @@ -24,6 +24,8 @@ struct DecoderStateSpec { std::int32_t kv_table_rows = 1; std::uint32_t text_physical_page_groups = 0; std::uint32_t mtp_physical_page_groups = 0; + // Entropy-coded cold pool capacity in pages; 0 disables the pool. + std::uint32_t max_cold_pages = 0; }; struct PagedKVCacheLayout { @@ -35,6 +37,12 @@ struct PagedKVCacheLayout { std::int32_t head_dim = 0; DType dtype = DType::BF16; std::int32_t quant_group = 0; + // Cold slots per layer: [slot_bytes, kv_heads, 2, max_cold_pages] + // plus an I32 validity plane of [kv_heads, 2, max_cold_pages]. + std::array cold_slots; + std::array cold_slot_valid; + std::int32_t cold_slot_bytes = 0; + std::uint32_t max_cold_pages = 0; [[nodiscard]] std::size_t payload_bytes() const noexcept { return pages.payload_bytes(); } }; @@ -48,6 +56,12 @@ class PagedKVCacheView { [[nodiscard]] bool valid() const noexcept { return cache_ != nullptr; } [[nodiscard]] std::uint32_t max_context() const noexcept; + // Cold-slot pool: fixed raw slots per (layer, kv_head, plane). + [[nodiscard]] std::int32_t cold_slot_bytes() const noexcept { return cold_slot_bytes_; } + [[nodiscard]] std::uint32_t max_cold_pages() const noexcept { return max_cold_pages_; } + std::int32_t allocate_cold_slot() noexcept; + void release_cold_slot(std::int32_t slot) noexcept; + [[nodiscard]] PagedKVLayerView layer_view(std::uint32_t layer) const; private: @@ -96,6 +110,11 @@ class PagedKVCache { std::int32_t kv_heads_ = 0; std::int32_t head_dim_ = 0; DType dtype_ = DType::BF16; + std::array cold_slots_; + std::array cold_slot_valid_; + std::int32_t cold_slot_bytes_ = 0; + std::uint32_t max_cold_pages_ = 0; + std::vector cold_slot_used_; std::int32_t quant_group_ = 0; }; diff --git a/src/targets/qwen3_6/impl/runtime/layouts.h b/src/targets/qwen3_6/impl/runtime/layouts.h index 859f4f5d7e..460782a60f 100644 --- a/src/targets/qwen3_6/impl/runtime/layouts.h +++ b/src/targets/qwen3_6/impl/runtime/layouts.h @@ -80,6 +80,9 @@ struct SequencePlanningInputs { StartupFeatures features; bool use_cuda_graph = true; bool causal_scoring = false; + ColdPolicy cold_policy = ColdPolicy::None; + std::uint32_t cold_keep_tokens = 128; + std::uint64_t cold_host_bytes = 4ULL << 30; int device = 0; ContextCacheOptions context_cache; }; diff --git a/src/targets/qwen3_6/impl/runtime/layouts_impl.h b/src/targets/qwen3_6/impl/runtime/layouts_impl.h index b21165ac86..3fb25fb154 100644 --- a/src/targets/qwen3_6/impl/runtime/layouts_impl.h +++ b/src/targets/qwen3_6/impl/runtime/layouts_impl.h @@ -137,6 +137,9 @@ PersistentLayout persistent_layout(const SequencePlanImpl& plan) { .kv_dtype = plan.kv_dtype, .kv_quant_group = plan.kv_quant_group, .enable_mtp = plan.features.mtp(), + .max_cold_pages = plan.cold_policy == ColdPolicy::Window + ? plan.cold_keep_tokens / kPagedKVPageSize + 16 + : 0, .kv_table_rows = static_cast(plan.max_concurrency), .text_physical_page_groups = physical_pages, .mtp_physical_page_groups = mtp_physical_pages, @@ -670,6 +673,9 @@ std::unique_ptr build_sequence_candidate(const SequencePlannin impl->proposal_head = inputs.proposal_head; impl->features = inputs.features; impl->use_cuda_graph = inputs.use_cuda_graph; + impl->cold_policy = inputs.cold_policy; + impl->cold_keep_tokens = inputs.cold_keep_tokens; + impl->cold_host_bytes = inputs.cold_host_bytes; impl->causal_scoring = inputs.causal_scoring; impl->device = inputs.device; impl->context_cache = inputs.context_cache; @@ -745,6 +751,9 @@ make_sequence_planner_impl(DeviceContext& device, const EngineOptions& options, .proposal_head = options.speculative.proposal_head, .features = qwen3_6::startup_features(options), .use_cuda_graph = options.use_cuda_graph, + .cold_policy = options.cold_policy, + .cold_keep_tokens = options.cold_keep_tokens, + .cold_host_bytes = options.cold_host_bytes, .causal_scoring = options.purpose == EnginePurpose::CausalScoring, .device = options.device, .context_cache = options.context_cache, diff --git a/src/targets/qwen3_6/impl/runtime/program.h b/src/targets/qwen3_6/impl/runtime/program.h index 61e566f998..4f435f77b6 100644 --- a/src/targets/qwen3_6/impl/runtime/program.h +++ b/src/targets/qwen3_6/impl/runtime/program.h @@ -688,6 +688,15 @@ class ProgramImplCore { std::size_t workspace_logical_peak_bytes = 0; + // Cold-pool maintenance (rev 2b): staging + per-step compress pass. + ColdPolicy cold_policy = ColdPolicy::None; + std::uint32_t cold_keep_tokens = 128; + std::uint64_t cold_host_bytes = 4ULL << 30; + void* cold_requant_codes = nullptr; + void* cold_requant_scales = nullptr; + std::uint32_t cold_requant_heads = 0; + void enqueue_cold_compressions(SequenceState& sequence); + std::size_t vision_handoff_peak_bytes = 0; private: diff --git a/src/targets/qwen3_6/impl/runtime/program_impl.h b/src/targets/qwen3_6/impl/runtime/program_impl.h index 751117abaf..2432be9280 100644 --- a/src/targets/qwen3_6/impl/runtime/program_impl.h +++ b/src/targets/qwen3_6/impl/runtime/program_impl.h @@ -729,6 +729,8 @@ ProgramImplCore::ProgramImplCore(const LoadedModelData& model_in, const Sequence speculative_backend(plan.speculative_backend), kv_dtype(plan.kv_dtype), kv_quant_group(plan.kv_quant_group), proposal_head(plan.proposal_head), vision_enabled(plan.features.vision), use_cuda_graph(plan.use_cuda_graph), + cold_policy(plan.cold_policy), cold_keep_tokens(plan.cold_keep_tokens), + cold_host_bytes(plan.cold_host_bytes), causal_scoring(plan.causal_scoring), kv_payload_bytes(plan.persistent.kv_payload_bytes), graph_allowance_bytes(plan.graph_allowance_bytes), workspace_plan(plan.workspace), persistent(plan.persistent.bytes), workspace_storage(plan.workspace.capacity), @@ -10144,6 +10146,99 @@ void ProgramImplCore::ordered_reset(SequenceState& sequence) { sequence.dflash_context_frontier = 0; } + +// Cold-pool maintenance: pack the retired tail of a sequence's text KV into +// raw entropy slots and detach those pages (sentinel entries in the block +// table; physical pages return to the pool). Runs when the window policy is +// active and at least cold_keep_tokens lie behind the decode frontier. +void ProgramImplCore::enqueue_cold_compressions(SequenceState& sequence) { + if (cold_policy != ColdPolicy::Window || !sequence.kv || decoder == nullptr || + !sequence.kv->text.valid()) { + return; + } + KVAddressSpaceStore& store = *text_kv_addresses; + KVAddressSpaceHandle text = sequence.kv->text; + const std::uint32_t total_pages = store.entitlement(text); + if (total_pages == 0) { return; } + + const std::uint32_t cold_pages = + sequence.text_kv_valid > cold_keep_tokens + ? (sequence.text_kv_valid - cold_keep_tokens) / kPagedKVPageSize + : 0; + if (cold_pages == 0 || cold_pages >= total_pages) { return; } + + // Allocate cold slots for the tail (one slot per page; shared across heads + // through the flat slot addressing in the kernels). + const std::int32_t slot = decoder->text_kv.allocate_cold_slot(); + if (slot < 0) { return; } + + const std::uint32_t keep_pages = total_pages - cold_pages; + const std::int32_t kv_heads = + decoder->text_kv.batch_layer_view(0).num_kv_heads; + + // Pack every layer's cold tail pages into the raw slots. + for (std::uint32_t layer = 0; layer < decoder->text_kv.layers(); ++layer) { + const PagedKVBatchLayerView view = decoder->text_kv.batch_layer_view(layer); + const Tensor cold_slots = view.cold_slots; + if (cold_slots.data == nullptr) { continue; } + const bool int8_layer = view.dtype == DType::I8; + const auto k_mode = int8_layer ? ops::EntropyColdRequantMode::Int8G64 + : ops::EntropyColdRequantMode::Nvfp4G16; + const auto v_mode = int8_layer ? ops::EntropyColdRequantMode::Int8G64 + : ops::EntropyColdRequantMode::Iso3VG16; + for (std::uint32_t p = 0; p < cold_pages; ++p) { + const DeviceKVPageHandle ph = store.physical_page(text, keep_pages + p); + if (!ph.valid()) { continue; } + const std::int32_t physical = ph.index(); + auto* k_codes = static_cast(view.k_pages.data) + + physical * view.k_pages.nb[3]; + auto* v_codes = static_cast(view.v_pages.data) + + physical * view.v_pages.nb[3]; + auto* k_scales = static_cast(view.k_scale_pages.data) + + physical * view.k_scale_pages.nb[3]; + auto* v_scales = static_cast(view.v_scale_pages.data) + + physical * view.v_scale_pages.nb[3]; + const std::int64_t slot_off = + (static_cast(slot) * 2 * kv_heads) * cold_slots.nb[0]; + auto* k_slot = static_cast(cold_slots.data) + slot_off; + auto* v_slot = k_slot + cold_slots.nb[2]; + auto* k_valid = static_cast(view.cold_slot_valid.data) + + static_cast(slot) * view.cold_slot_valid.nb[2]; + auto* v_valid = reinterpret_cast( + reinterpret_cast(k_valid) + view.cold_slot_valid.nb[1]); + ops::entropy_cold_requant_raw( + k_codes, k_scales, k_mode, kv_heads, 1, + static_cast(cold_requant_codes), + static_cast(cold_requant_scales), device.stream); + ops::cold_i8_slot_pack_raw( + static_cast(cold_requant_codes), + static_cast(cold_requant_scales), kv_heads, 1, k_slot, + k_valid, device.stream); + ops::entropy_cold_requant_raw( + v_codes, v_scales, v_mode, kv_heads, 1, + static_cast(cold_requant_codes), + static_cast(cold_requant_scales), device.stream); + ops::cold_i8_slot_pack_raw( + static_cast(cold_requant_codes), + static_cast(cold_requant_scales), kv_heads, 1, v_slot, + v_valid, device.stream); + } + } + device.synchronize(); + + // Detach the tail: shrink the address space, return the physical pages, + // and publish sentinel entries in the execution row. + const KVExecutionRowLease& row = store.execution_row(text); + std::vector sentinel(cold_pages, paged_kv_cold_entry(slot)); + decoder->text_kv.execution_tables().publish_indices( + row.handle(), keep_pages, sentinel, device.stream); + store.deactivate(text); + (void)store.activate(text, keep_pages, static_cast(row.index())); + device.synchronize(); + std::fprintf(stderr, "[cold] compressed tail %u pages -> slot %d (kept %u)\n", + cold_pages, slot, keep_pages); +} + void ProgramImplCore::prepare_graphs() { if (!use_cuda_graph) { return; } nvtx::ScopedRange prepare_range(nvtx::Name::CudaGraphPrepare, nvtx::Category::Graph); @@ -11424,6 +11519,11 @@ runtime::BatchedGeneratedRound ProgramImplCore::decode_raw(std::span lanes, std::span budgets, runtime::ExecutionTiming* failed_timing) { + // Cold-pool maintenance at the round boundary (window policy only). + if (cold_policy == ColdPolicy::Window && lanes.size() == 1 && + sequences[lanes[0]].kv) { + enqueue_cold_compressions(sequences[lanes[0]]); + } if (speculative_backend == SpeculativeBackend::None) { return decode_ordinary_batch(lanes, budgets, failed_timing); } diff --git a/src/targets/qwen3_6/impl/state/decoder_state.cpp b/src/targets/qwen3_6/impl/state/decoder_state.cpp index 5e7372df45..ddf0a03e23 100644 --- a/src/targets/qwen3_6/impl/state/decoder_state.cpp +++ b/src/targets/qwen3_6/impl/state/decoder_state.cpp @@ -1,4 +1,5 @@ #include +#include "ninfer/ops/cold_i8.h" #include #include @@ -74,17 +75,58 @@ DecoderStateLayout plan_decoder_state(LayoutBuilder& builder, const DecoderState spec.attention_head_dim, spec.kv_dtype, spec.kv_quant_group, spec.kv_table_rows, spec.mtp_physical_page_groups); } + // Entropy-coded cold pool: fixed raw slots (9232 B) plus an I32 validity + // plane, per full-attention layer. Only active when the spec opts in. + const std::int32_t cold_slot_bytes = ops::kColdI8SlotBytes; + if (spec.max_cold_pages != 0) { + const std::uint32_t cold_pages = spec.max_cold_pages; + for (std::uint32_t layer = 0; layer < spec.full_attention_layers; ++layer) { + layout.cold_slots[layer] = builder.add_tensor( + DType::U8, {cold_slot_bytes, static_cast(spec.kv_heads), + 2, cold_pages}, + 256, "cold slots L" + std::to_string(layer)); + layout.cold_slot_valid[layer] = builder.add_tensor( + DType::I32, {static_cast(spec.kv_heads), 2, cold_pages}, 256, + "cold slot valid L" + std::to_string(layer)); + } + } return layout; } PagedKVCache::PagedKVCache(DeviceSpan backing, const PagedKVCacheLayout& layout) : pages_(backing, layout.pages), execution_tables_(backing, layout.execution_tables, pages_), layers_(layout.layers), max_context_(layout.max_context), kv_heads_(layout.kv_heads), - head_dim_(layout.head_dim), dtype_(layout.dtype), quant_group_(layout.quant_group) {} + head_dim_(layout.head_dim), dtype_(layout.dtype), quant_group_(layout.quant_group), + cold_slot_bytes_(layout.cold_slot_bytes), max_cold_pages_(layout.max_cold_pages) { + cold_slot_used_.assign(max_cold_pages_, 0); + for (std::uint32_t layer = 0; layer < layers_; ++layer) { + if (layout.cold_slots[layer].region.bytes != 0) { + cold_slots_[layer] = layout.cold_slots[layer].bind(backing); + cold_slot_valid_[layer] = layout.cold_slot_valid[layer].bind(backing); + } + } +} PagedKVCacheView::PagedKVCacheView(const PagedKVCache& cache, Tensor block_table) noexcept : cache_(&cache), block_table_(block_table) {} +std::int32_t PagedKVCache::allocate_cold_slot() noexcept { + if (max_cold_pages_ == 0) { return -1; } + for (std::uint32_t slot = 0; slot < max_cold_pages_; ++slot) { + if (!cold_slot_used_[slot]) { + cold_slot_used_[slot] = true; + return static_cast(slot); + } + } + return -1; +} + +void PagedKVCache::release_cold_slot(std::int32_t slot) noexcept { + if (slot >= 0 && static_cast(slot) < max_cold_pages_) { + cold_slot_used_[slot] = false; + } +} + std::uint32_t PagedKVCacheView::max_context() const noexcept { return cache_ == nullptr ? 0 : cache_->max_context(); } @@ -130,6 +172,8 @@ PagedKVBatchLayerView PagedKVCache::batch_layer_view(std::uint32_t layer) const .k_scale_pages = scaled ? pages_.plane(base + 2) : Tensor(), .v_scale_pages = scaled ? pages_.plane(base + 3) : Tensor(), .block_tables = execution_tables_.matrix(), + .cold_slots = cold_slots_[layer], + .cold_slot_valid = cold_slot_valid_[layer], .head_dim = head_dim_, .num_kv_heads = kv_heads_, .dtype = dtype_, From d6fad5694a4373492e153873530299121c727ae2 Mon Sep 17 00:00:00 2001 From: NInfer Agent Date: Sun, 30 Aug 2026 23:10:52 +0800 Subject: [PATCH 10/11] feat(kv): full cold-pool window mechanism on master's paged KV store --- src/core/paged_kv_cache.h | 15 +- .../dense/causal_cache/small_t.cu | 14 +- .../dense/causal_cache/small_t_bf16.cuh | 81 +++++++ .../dense/causal_cache/small_t_i8.cuh | 95 ++++++++ .../ninfer/targets/qwen3_6/decoder_state.h | 18 +- src/targets/qwen3_6/impl/runtime/layouts.h | 3 + .../qwen3_6/impl/runtime/layouts_impl.h | 8 +- .../qwen3_6/impl/runtime/logical_kv_store.h | 84 +++++++ src/targets/qwen3_6/impl/runtime/program.h | 13 + .../qwen3_6/impl/runtime/program_impl.h | 226 ++++++++++++++---- .../qwen3_6/impl/state/decoder_state.cpp | 10 +- 11 files changed, 502 insertions(+), 65 deletions(-) diff --git a/src/core/paged_kv_cache.h b/src/core/paged_kv_cache.h index f59f3242eb..c4c9c6089c 100644 --- a/src/core/paged_kv_cache.h +++ b/src/core/paged_kv_cache.h @@ -38,6 +38,11 @@ struct PagedKVLayerView { Tensor k_scale_pages; Tensor v_scale_pages; Tensor block_table; + // Entropy-coded cold pool (single-sequence window view): fixed raw slots + + // validity flags. Empty when the cache has no cold pool. + Tensor cold_slots; + Tensor cold_slot_valid; + std::int32_t cold_slot_bytes = 0; std::int32_t head_dim = 0; std::int32_t num_kv_heads = 0; DType dtype = DType::BF16; @@ -54,6 +59,7 @@ struct PagedKVBatchLayerView { // Entropy-coded cold pool: fixed raw slots + validity flags per layer. Tensor cold_slots; Tensor cold_slot_valid; + std::int32_t cold_slot_bytes = 0; std::int32_t head_dim = 0; std::int32_t num_kv_heads = 0; DType dtype = DType::BF16; @@ -130,6 +136,10 @@ class DeviceKVPageHandle { [[nodiscard]] bool valid() const noexcept { return owner_ != nullptr; } + // Physical page index within the owning pool (public read access for the + // cold-compression pass). + [[nodiscard]] std::int32_t index() const noexcept { return index_; } + private: friend class DeviceKVPagePool; friend class DeviceKVPageLease; @@ -369,13 +379,14 @@ class KVExecutionTablePool { [[nodiscard]] const Tensor& matrix() const noexcept { return block_tables_; } + void publish_indices(KVExecutionRowHandle row, std::uint32_t logical_begin, + std::span indices, cudaStream_t stream); private: friend class KVExecutionRowLease; [[nodiscard]] bool valid_handle(KVExecutionRowHandle handle) const noexcept; bool release_row(std::int32_t row, std::uint32_t generation) noexcept; - void publish_indices(KVExecutionRowHandle row, std::uint32_t logical_begin, - std::span indices, cudaStream_t stream); + KVExecutionTableSpec spec_; const DeviceKVPagePool* pages_ = nullptr; diff --git a/src/ops/softmax_attention/dense/causal_cache/small_t.cu b/src/ops/softmax_attention/dense/causal_cache/small_t.cu index 6caa50ac56..383aeabf11 100644 --- a/src/ops/softmax_attention/dense/causal_cache/small_t.cu +++ b/src/ops/softmax_attention/dense/causal_cache/small_t.cu @@ -115,8 +115,10 @@ void launch_tc_partial_bf16(const Tensor& q, CacheInput input, const Tensor& pos : static_cast(invocation.table_rows->data), cache.block_tables.ne[0], invocation.width, invocation.full_width, invocation.column_begin, logical_capacity, scale, - static_cast<__nv_bfloat16*>(partial_acc.data), static_cast(partial_m.data), - static_cast(partial_l.data)); + static_cast(cache.cold_slots.data), + static_cast(cache.cold_slot_valid.data), + cache.cold_slot_bytes, static_cast<__nv_bfloat16*>(partial_acc.data), + static_cast(partial_m.data), static_cast(partial_l.data)); CUDA_CHECK(cudaGetLastError()); } @@ -158,7 +160,10 @@ void launch_tc_partial_i8(const Tensor& q, CacheInput input, const Tensor& pos, ? nullptr : static_cast(invocation.table_rows->data), cache.block_tables.ne[0], invocation.full_width, invocation.column_begin, - logical_capacity, scale, static_cast<__nv_bfloat16*>(partial_acc.data), + logical_capacity, scale, + static_cast(cache.cold_slots.data), + static_cast(cache.cold_slot_valid.data), + cache.cold_slot_bytes, static_cast<__nv_bfloat16*>(partial_acc.data), static_cast(partial_m.data), static_cast(partial_l.data)); }; if constexpr (TokenTile == 6) { @@ -216,6 +221,9 @@ PagedKVBatchLayerView single_row_batch_view(const PagedKVLayerView& cache) { .k_scale_pages = cache.k_scale_pages, .v_scale_pages = cache.v_scale_pages, .block_tables = cache.block_table.view({cache.block_table.ne[0], 1}), + .cold_slots = cache.cold_slots, + .cold_slot_valid = cache.cold_slot_valid, + .cold_slot_bytes = cache.cold_slot_bytes, .head_dim = cache.head_dim, .num_kv_heads = cache.num_kv_heads, .dtype = cache.dtype, diff --git a/src/ops/softmax_attention/dense/causal_cache/small_t_bf16.cuh b/src/ops/softmax_attention/dense/causal_cache/small_t_bf16.cuh index 2b8a53b11d..74c9a9347b 100644 --- a/src/ops/softmax_attention/dense/causal_cache/small_t_bf16.cuh +++ b/src/ops/softmax_attention/dense/causal_cache/small_t_bf16.cuh @@ -9,6 +9,7 @@ #include #include +#include "ops/kernel/cold_i8_kernels.cuh" #include "ops/softmax_attention/dense/causal_cache/small_t.cuh" #include @@ -22,6 +23,7 @@ __launch_bounds__(128, 2) __global__ void causal_attention_small_t_tc_partial_bf __nv_bfloat16* cache_v, const std::int32_t* block_tables, const std::int32_t* valid_columns, const std::int32_t* table_rows, std::int32_t table_stride, std::int32_t tokens, std::int32_t full_width, std::int32_t column_begin, std::int32_t logical_capacity, float scale, + const std::uint8_t* cold_slots, const std::int32_t* cold_valid, std::int32_t cold_slot_bytes, __nv_bfloat16* partial_acc, float* partial_m, float* partial_l) { static_assert(TokenTile >= 1 && TokenTile <= 6); static_assert(WarpsPerCta >= 1 && WarpsPerCta <= 4); @@ -217,6 +219,15 @@ __launch_bounds__(128, 2) __global__ void causal_attention_small_t_tc_partial_bf if (kb != 0 && (k0 & kPagedKVPageMask) == 0) { physical_page = physical_pages_s[(k0 >> kPagedKVPageShift) - first_page]; } + // A Bc=32 tile never crosses a 64-token page boundary, so the entry + // cached for this tile decides the whole tile's load path. Cold pages + // carry a sentinel (<= -2): decode E2M1 nibbles + E4M3 g16 scales + // straight from the raw slot into the bf16 tile. + const int entry = physical_pages_s[(k0 >> kPagedKVPageShift) - first_page]; + const bool cold = entry <= -2 && cold_slots != nullptr && cold_slot_bytes >= 1024 + 320 && + cold_valid[(-entry - 2) * (2 * Geometry::KVHeads) + kv_head] != 0 && + cold_valid[(-entry - 2) * (2 * Geometry::KVHeads) + Geometry::KVHeads + + kv_head] != 0; // Stage the bf16 K/V key tile with one cp.async wave (16B/thread, high MLP). // Current-step tokens come from k_new/v_new; tail slots are zeroed. #pragma unroll 1 @@ -236,12 +247,82 @@ __launch_bounds__(128, 2) __global__ void causal_attention_small_t_tc_partial_bf kv_cache_int8_new_index(kv_head, d, new_token); ninfer::ops::cp_async<16>(k_dst, &input.k[off]); ninfer::ops::cp_async<16>(v_dst, &input.v[off]); + } else if (cold) { + const int slot_base = -entry - 2; + const std::int64_t k_off = + static_cast(slot_base * (2 * Geometry::KVHeads) + + kv_head) * + cold_slot_bytes; + const std::int64_t v_off = k_off + static_cast( + Geometry::KVHeads) * + cold_slot_bytes; + const std::uint8_t* k_row = detail::cold_i8_slot_codes(cold_slots + k_off) + + (key & kPagedKVPageMask) * 128; + const std::uint8_t* k_row_s = + detail::cold_i8_slot_scales(cold_slots + k_off) + + (key & kPagedKVPageMask) * 16; + const std::uint8_t* v_row = detail::cold_i8_slot_codes(cold_slots + v_off) + + (key & kPagedKVPageMask) * 128; + const std::uint8_t* v_row_s = + detail::cold_i8_slot_scales(cold_slots + v_off) + + (key & kPagedKVPageMask) * 16; +#pragma unroll + for (int i = 0; i < 8; ++i) { + const int chan = d + i; + const std::uint8_t kb = k_row[chan >> 1]; + const std::uint8_t vb = v_row[chan >> 1]; + const float k_code = + gqa_kv_nvfp4_e2m1_to_f32((chan & 1) ? (kb >> 4) : (kb & 0x0F)); + const float v_code = + gqa_kv_nvfp4_e2m1_to_f32((chan & 1) ? (vb >> 4) : (vb & 0x0F)); + const float k_scale = + gqa_kv_nvfp4_e4m3_to_f32(k_row_s[chan >> 4]); + const float v_scale = + gqa_kv_nvfp4_e4m3_to_f32(v_row_s[chan >> 4]); + k_dst[i] = __float2bfloat16(k_code * k_scale); + v_dst[i] = __float2bfloat16(v_code * v_scale); + } } else { const std::int64_t off = causal_cache_index( physical_page, kv_head, d, key & kPagedKVPageMask); ninfer::ops::cp_async<16>(k_dst, &cache_k[off]); ninfer::ops::cp_async<16>(v_dst, &cache_v[off]); } + } else if (cold) { + const int slot_base = -entry - 2; + const std::int64_t k_off = + static_cast(slot_base * (2 * Geometry::KVHeads) + + kv_head) * + cold_slot_bytes; + const std::int64_t v_off = k_off + static_cast( + Geometry::KVHeads) * + cold_slot_bytes; + const std::uint8_t* k_row = detail::cold_i8_slot_codes(cold_slots + k_off) + + (key & kPagedKVPageMask) * 128; + const std::uint8_t* k_row_s = + detail::cold_i8_slot_scales(cold_slots + k_off) + + (key & kPagedKVPageMask) * 16; + const std::uint8_t* v_row = detail::cold_i8_slot_codes(cold_slots + v_off) + + (key & kPagedKVPageMask) * 128; + const std::uint8_t* v_row_s = + detail::cold_i8_slot_scales(cold_slots + v_off) + + (key & kPagedKVPageMask) * 16; +#pragma unroll + for (int i = 0; i < 8; ++i) { + const int chan = d + i; + const std::uint8_t kb = k_row[chan >> 1]; + const std::uint8_t vb = v_row[chan >> 1]; + const float k_code = + gqa_kv_nvfp4_e2m1_to_f32((chan & 1) ? (kb >> 4) : (kb & 0x0F)); + const float v_code = + gqa_kv_nvfp4_e2m1_to_f32((chan & 1) ? (vb >> 4) : (vb & 0x0F)); + const float k_scale = + gqa_kv_nvfp4_e4m3_to_f32(k_row_s[chan >> 4]); + const float v_scale = + gqa_kv_nvfp4_e4m3_to_f32(v_row_s[chan >> 4]); + k_dst[i] = __float2bfloat16(k_code * k_scale); + v_dst[i] = __float2bfloat16(v_code * v_scale); + } } else { const std::int64_t off = causal_cache_index(physical_page, kv_head, d, key & kPagedKVPageMask); diff --git a/src/ops/softmax_attention/dense/causal_cache/small_t_i8.cuh b/src/ops/softmax_attention/dense/causal_cache/small_t_i8.cuh index 2145eeb585..be72570579 100644 --- a/src/ops/softmax_attention/dense/causal_cache/small_t_i8.cuh +++ b/src/ops/softmax_attention/dense/causal_cache/small_t_i8.cuh @@ -22,6 +22,7 @@ #include #include +#include "ops/kernel/cold_i8_kernels.cuh" #include "ops/softmax_attention/dense/causal_cache/small_t.cuh" #include "ops/kv_cache/int8_g64_codec.cuh" @@ -62,6 +63,7 @@ __launch_bounds__(WarpsPerCta * 32, MinBlocksPerSm) __global__ const std::int32_t* block_tables, const std::int32_t* valid_columns, const std::int32_t* table_rows, std::int32_t table_stride, std::int32_t full_width, std::int32_t column_begin, std::int32_t logical_capacity, float scale, + const std::uint8_t* cold_slots, const std::int32_t* cold_valid, std::int32_t cold_slot_bytes, __nv_bfloat16* partial_acc, float* partial_m, float* partial_l) { constexpr int Wc = WarpsPerCta; constexpr int RowCount = TokenTile * Geometry::GroupSize; @@ -357,6 +359,99 @@ __launch_bounds__(WarpsPerCta * 32, MinBlocksPerSm) __global__ float l0 = 0.0f, l1 = 0.0f; auto issue_kv_tile = [&](int tile_k0, int physical_page) { + // A Bc=32 tile never crosses a 64-token page boundary. Cold pages + // carry a sentinel (<= -2): decode E2M1 nibbles + E4M3 g16 scales + // from the raw slot back into int8 codes + fp16 g64 scales, matching + // the native planes the hot path stages with cp.async. + const int entry = physical_pages_s[(tile_k0 >> kPagedKVPageShift) - first_page]; + const bool cold = entry <= -2 && cold_slots != nullptr && cold_slot_bytes >= 1024 + 320 && + cold_valid[(-entry - 2) * (2 * Geometry::KVHeads) + kv_head] != 0 && + cold_valid[(-entry - 2) * (2 * Geometry::KVHeads) + Geometry::KVHeads + + kv_head] != 0; + if (cold) { + const int slot_base = -entry - 2; + const std::int64_t k_off = + static_cast(slot_base * (2 * Geometry::KVHeads) + kv_head) * + cold_slot_bytes; + const std::int64_t v_off = + k_off + static_cast(Geometry::KVHeads) * cold_slot_bytes; + for (int key_l = tid; key_l < Bc; key_l += Threads) { + const int key = tile_k0 + key_l; + if (key >= split_start && key < split_end) { + const int row = key & kPagedKVPageMask; + const std::uint8_t* k_rs = + detail::cold_i8_slot_scales(cold_slots + k_off) + row * 16; + const std::uint8_t* v_rs = + detail::cold_i8_slot_scales(cold_slots + v_off) + row * 16; +#pragma unroll + for (int g = 0; g < Groups; ++g) { + float mk = 0.0f, mv = 0.0f; +#pragma unroll + for (int s = 0; s < 4; ++s) { + mk = fmaxf(mk, gqa_kv_nvfp4_e4m3_to_f32(k_rs[g * 4 + s])); + mv = fmaxf(mv, gqa_kv_nvfp4_e4m3_to_f32(v_rs[g * 4 + s])); + } + k_scale_s[key_l * Groups + g] = __float2half(mk * 6.0f / 127.0f); + v_scale_s[key_l * Groups + g] = __float2half(mv * 6.0f / 127.0f); + } + } else { + store_vec(&k_scale_s[key_l * Groups], make_int2(0, 0)); + store_vec(&v_scale_s[key_l * Groups], make_int2(0, 0)); + } + } +#pragma unroll 1 + for (int chunk = tid; chunk < Bc * (D / 16); chunk += Threads) { + const int key_l = chunk / (D / 16); + const int dc = chunk - key_l * (D / 16); + const int d = dc * 16; + const int key = tile_k0 + key_l; + std::int8_t* dst = &k_i8[key_l * D + causal_small_t_tc_swz(key_l, dc * 8) * 2]; + if (key >= split_start && key < split_end) { + const int row = key & kPagedKVPageMask; + const std::uint8_t* k_rc = + detail::cold_i8_slot_codes(cold_slots + k_off) + row * 128; + const std::uint8_t* k_rs = + detail::cold_i8_slot_scales(cold_slots + k_off) + row * 16; + const std::uint8_t* v_rc = + detail::cold_i8_slot_codes(cold_slots + v_off) + row * 128; + const std::uint8_t* v_rs = + detail::cold_i8_slot_scales(cold_slots + v_off) + row * 16; + // 16 channels never straddle a 64-channel group, so each + // chunk recomputes its own upper-bound group scale. + const int g = d >> 6; + float mk = 0.0f, mv = 0.0f; +#pragma unroll + for (int s = 0; s < 4; ++s) { + mk = fmaxf(mk, gqa_kv_nvfp4_e4m3_to_f32(k_rs[g * 4 + s])); + mv = fmaxf(mv, gqa_kv_nvfp4_e4m3_to_f32(v_rs[g * 4 + s])); + } + const float k_inv = mk > 0.0f ? 127.0f / (mk * 6.0f) : 0.0f; + const float v_inv = mv > 0.0f ? 127.0f / (mv * 6.0f) : 0.0f; +#pragma unroll + for (int i = 0; i < 16; ++i) { + const int chan = d + i; + const std::uint8_t kb = k_rc[chan >> 1]; + const std::uint8_t vb = v_rc[chan >> 1]; + const float k_code = + gqa_kv_nvfp4_e2m1_to_f32((chan & 1) ? (kb >> 4) : (kb & 0x0F)); + const float v_code = + gqa_kv_nvfp4_e2m1_to_f32((chan & 1) ? (vb >> 4) : (vb & 0x0F)); + const float k_scale = gqa_kv_nvfp4_e4m3_to_f32(k_rs[chan >> 4]); + const float v_scale = gqa_kv_nvfp4_e4m3_to_f32(v_rs[chan >> 4]); + int kc = __float2int_rn(k_code * k_scale * k_inv); + int vc = __float2int_rn(v_code * v_scale * v_inv); + kc = max(-127, min(127, kc)); + vc = max(-127, min(127, vc)); + dst[i] = static_cast(kc); + v_i8[key_l * D + chan] = static_cast(vc); + } + } else { + store_vec(dst, make_int4(0, 0, 0, 0)); + store_vec(&v_i8[key_l * D + d], make_int4(0, 0, 0, 0)); + } + } + return; + } for (int key_l = tid; key_l < Bc; key_l += Threads) { const int key = tile_k0 + key_l; if (key >= split_start && key < split_end) { diff --git a/src/targets/qwen3_6/export/ninfer/targets/qwen3_6/decoder_state.h b/src/targets/qwen3_6/export/ninfer/targets/qwen3_6/decoder_state.h index 7fded328b3..5233cc8f55 100644 --- a/src/targets/qwen3_6/export/ninfer/targets/qwen3_6/decoder_state.h +++ b/src/targets/qwen3_6/export/ninfer/targets/qwen3_6/decoder_state.h @@ -56,11 +56,6 @@ class PagedKVCacheView { [[nodiscard]] bool valid() const noexcept { return cache_ != nullptr; } [[nodiscard]] std::uint32_t max_context() const noexcept; - // Cold-slot pool: fixed raw slots per (layer, kv_head, plane). - [[nodiscard]] std::int32_t cold_slot_bytes() const noexcept { return cold_slot_bytes_; } - [[nodiscard]] std::uint32_t max_cold_pages() const noexcept { return max_cold_pages_; } - std::int32_t allocate_cold_slot() noexcept; - void release_cold_slot(std::int32_t slot) noexcept; [[nodiscard]] PagedKVLayerView layer_view(std::uint32_t layer) const; @@ -81,7 +76,13 @@ class PagedKVCache { PagedKVCache(PagedKVCache&&) = delete; PagedKVCache& operator=(PagedKVCache&&) = delete; - [[nodiscard]] std::uint32_t max_context() const noexcept { return max_context_; } + // Cold-slot pool: fixed raw slots per (layer, kv_head, plane). + [[nodiscard]] std::int32_t cold_slot_bytes() const noexcept { return cold_slot_bytes_; } + [[nodiscard]] std::uint32_t max_cold_pages() const noexcept { return max_cold_pages_; } + std::int32_t allocate_cold_slot() noexcept; + void release_cold_slot(std::int32_t slot) noexcept; + +[[nodiscard]] std::uint32_t max_context() const noexcept { return max_context_; } [[nodiscard]] std::uint32_t layers() const noexcept { return layers_; } @@ -99,6 +100,10 @@ class PagedKVCache { [[nodiscard]] PagedKVBatchLayerView batch_layer_view(std::uint32_t layer) const; + [[nodiscard]] Tensor cold_slot_valid(std::uint32_t layer) const noexcept { + return layer < layers_ ? cold_slot_valid_[layer] : Tensor{}; + } + private: friend class PagedKVCacheView; [[nodiscard]] PagedKVLayerView layer_view(std::uint32_t layer, Tensor block_table) const; @@ -108,6 +113,7 @@ class PagedKVCache { std::uint32_t layers_ = 0; std::uint32_t max_context_ = 0; std::int32_t kv_heads_ = 0; + std::int32_t head_dim_ = 0; DType dtype_ = DType::BF16; std::array cold_slots_; diff --git a/src/targets/qwen3_6/impl/runtime/layouts.h b/src/targets/qwen3_6/impl/runtime/layouts.h index 460782a60f..8ba561c31e 100644 --- a/src/targets/qwen3_6/impl/runtime/layouts.h +++ b/src/targets/qwen3_6/impl/runtime/layouts.h @@ -106,6 +106,9 @@ struct SequencePlanImpl { ProposalHead proposal_head = ProposalHead::Full; StartupFeatures features; bool use_cuda_graph = true; + ColdPolicy cold_policy = ColdPolicy::None; + std::uint32_t cold_keep_tokens = 128; + std::uint64_t cold_host_bytes = 4ULL << 30; bool causal_scoring = false; int device = 0; ContextCacheOptions context_cache; diff --git a/src/targets/qwen3_6/impl/runtime/layouts_impl.h b/src/targets/qwen3_6/impl/runtime/layouts_impl.h index 3fb25fb154..b02ec813d9 100644 --- a/src/targets/qwen3_6/impl/runtime/layouts_impl.h +++ b/src/targets/qwen3_6/impl/runtime/layouts_impl.h @@ -137,12 +137,12 @@ PersistentLayout persistent_layout(const SequencePlanImpl& plan) { .kv_dtype = plan.kv_dtype, .kv_quant_group = plan.kv_quant_group, .enable_mtp = plan.features.mtp(), - .max_cold_pages = plan.cold_policy == ColdPolicy::Window - ? plan.cold_keep_tokens / kPagedKVPageSize + 16 - : 0, .kv_table_rows = static_cast(plan.max_concurrency), .text_physical_page_groups = physical_pages, .mtp_physical_page_groups = mtp_physical_pages, + .max_cold_pages = plan.cold_policy == ColdPolicy::Window + ? plan.cold_keep_tokens / kPagedKVPageSize + 16 + : 0, }); qwen3_6::StateImageSpec state_image_spec{ .linear = @@ -751,10 +751,10 @@ make_sequence_planner_impl(DeviceContext& device, const EngineOptions& options, .proposal_head = options.speculative.proposal_head, .features = qwen3_6::startup_features(options), .use_cuda_graph = options.use_cuda_graph, + .causal_scoring = options.purpose == EnginePurpose::CausalScoring, .cold_policy = options.cold_policy, .cold_keep_tokens = options.cold_keep_tokens, .cold_host_bytes = options.cold_host_bytes, - .causal_scoring = options.purpose == EnginePurpose::CausalScoring, .device = options.device, .context_cache = options.context_cache, }; diff --git a/src/targets/qwen3_6/impl/runtime/logical_kv_store.h b/src/targets/qwen3_6/impl/runtime/logical_kv_store.h index ba172610f0..a1b8aa8493 100644 --- a/src/targets/qwen3_6/impl/runtime/logical_kv_store.h +++ b/src/targets/qwen3_6/impl/runtime/logical_kv_store.h @@ -760,6 +760,49 @@ class LogicalKVPageStore { release_descriptor(handle, page); } + // Cold-pool transfer: detach the page's device replica (returning the + // physical page to the pool) while keeping the descriptor, so the address + // membership and its block-table slot stay stable. The page's contents + // live in a fixed cold slot; restore_from_cold brings them back on demand. + void transfer_to_cold(LogicalKVPageHandle handle, DeviceKVPageReservation& reservation) { + Page& page = require(handle); + if (page.references != 1 || page.writer_references != 1 || page.source_pins != 0 || + page.destination_pinned || page.host_replica || !page.device_replica || + page.cold_compressed) { + throw std::logic_error("logical KV page is not cold-transferable"); + } + physical_->dematerialize_one(reservation, std::move(*page.device_replica)); + page.device_replica.reset(); + page.cold_compressed = true; + } + + [[nodiscard]] bool cold_compressed(LogicalKVPageHandle handle) const noexcept { + return valid(handle) && pages_[handle.index_].cold_compressed; + } + + // Cold-pool restore: allocate a fresh physical page for a cold descriptor + // and hand it back so the caller can repopulate it from the cold slot. + // The descriptor keeps its membership position and reference counts. + [[nodiscard]] DeviceKVPageHandle restore_from_cold(LogicalKVPageHandle handle, + DeviceKVPageReservation& reservation) { + Page& page = require(handle); + if (!page.cold_compressed) { + throw std::logic_error("logical KV page is not cold"); + } + if (reservation.pages() == 0) { + if (!physical_->can_resize_reservation(reservation, 1)) { + throw std::logic_error("Paged KV pool has no capacity for a cold restore"); + } + physical_->resize_reservation(reservation, 1); + } + DeviceKVPageLease lease = physical_->materialize_one(reservation); + DeviceKVPageHandle lease_handle = lease.handle(); + page.device_replica.emplace(std::move(lease)); + page.cold_compressed = false; + page.content_epoch = next_epoch(page.content_epoch); + return lease_handle; + } + [[nodiscard]] bool release(LogicalKVPageHandle handle) noexcept { if (!valid(handle)) { return false; } Page& page = pages_[handle.index_]; @@ -778,6 +821,7 @@ class LogicalKVPageStore { std::uint8_t writer_references = 0; bool destination_pinned = false; bool occupied = false; + bool cold_compressed = false; std::optional device_replica; std::optional pending_device_replica; std::optional host_replica; @@ -1648,6 +1692,46 @@ class KVAddressSpaceStore { return pages_->physical(membership(address, logical_page)); } + // Cold-pool accessors: the window maintenance path detaches retired pages + // into fixed cold slots (transfer_to_cold + sentinel block-table entries) + // and restores them on demand (restore_from_cold + physical page publish). + [[nodiscard]] bool cold_compressed(KVAddressSpaceHandle handle, + std::uint32_t logical_page) const { + const Address& address = require(handle); + if (logical_page >= address.page_count) { + throw std::out_of_range("KV cold page is outside the address space"); + } + return pages_->cold_compressed(membership(address, logical_page)); + } + + // Whether the page may be detached into a cold slot right now: not already + // cold, exclusively referenced by its writer, and free of pins/fork ties. + [[nodiscard]] bool can_cold_transfer(KVAddressSpaceHandle handle, + std::uint32_t logical_page) const noexcept { + if (!valid(handle)) { return false; } + const Address& address = addresses_[handle.index_]; + if (logical_page >= address.page_count) { return false; } + const LogicalKVPageHandle& logical = membership(address, logical_page); + return !pages_->cold_compressed(logical) && pages_->can_dematerialize(logical); + } + + void transfer_to_cold(KVAddressSpaceHandle handle, std::uint32_t logical_page) { + Address& address = require_active(handle); + if (logical_page >= address.page_count) { + throw std::out_of_range("KV cold transfer is outside the address space"); + } + pages_->transfer_to_cold(membership(address, logical_page), address.reservation); + } + + [[nodiscard]] DeviceKVPageHandle restore_from_cold(KVAddressSpaceHandle handle, + std::uint32_t logical_page) { + Address& address = require_active(handle); + if (logical_page >= address.page_count) { + throw std::out_of_range("KV cold restore is outside the address space"); + } + return pages_->restore_from_cold(membership(address, logical_page), address.reservation); + } + [[nodiscard]] std::uint64_t content_epoch(KVAddressSpaceHandle handle, std::uint32_t logical_page) const { const Address& address = require(handle); diff --git a/src/targets/qwen3_6/impl/runtime/program.h b/src/targets/qwen3_6/impl/runtime/program.h index 4f435f77b6..f21b1c0cfc 100644 --- a/src/targets/qwen3_6/impl/runtime/program.h +++ b/src/targets/qwen3_6/impl/runtime/program.h @@ -453,6 +453,18 @@ struct SequenceState { std::vector shared_prefix_references; runtime::PrefillWork rebuild_work; std::uint32_t rebuild_tail_begin = 0; + + // Cold-pool bookkeeping: text pages currently detached into raw cold + // slots (logical page -> slot). Released with the sequence or when the + // rewrite path warms the prefix back into physical pages. + struct ColdPageEntry { + std::uint32_t page; + std::int32_t slot; + }; + std::vector cold_pages; + // First logical page not yet offered to the cold pool; compression scans + // forward from here so each round only visits the newly retired pages. + std::uint32_t cold_frontier = 0; }; struct SharedPrefixState { @@ -696,6 +708,7 @@ class ProgramImplCore { void* cold_requant_scales = nullptr; std::uint32_t cold_requant_heads = 0; void enqueue_cold_compressions(SequenceState& sequence); + void warm_cold_prefix(SequenceState& sequence, std::uint32_t end_page); std::size_t vision_handoff_peak_bytes = 0; diff --git a/src/targets/qwen3_6/impl/runtime/program_impl.h b/src/targets/qwen3_6/impl/runtime/program_impl.h index 2432be9280..2852d8e1b6 100644 --- a/src/targets/qwen3_6/impl/runtime/program_impl.h +++ b/src/targets/qwen3_6/impl/runtime/program_impl.h @@ -21,6 +21,7 @@ #include #include #include +#include #include #include #include @@ -812,6 +813,16 @@ ProgramImplCore::ProgramImplCore(const LoadedModelData& model_in, const Sequence }; decoder = std::make_unique(backing, plan.persistent.decoder); + if (cold_policy == ColdPolicy::Window) { + const std::int32_t requant_heads = decoder->text_kv.batch_layer_view(0).num_kv_heads; + if (requant_heads > 0) { + CUDA_CHECK(cudaMalloc(&cold_requant_codes, + 8192ULL * static_cast(requant_heads))); + CUDA_CHECK(cudaMalloc(&cold_requant_scales, + 1024ULL * static_cast(requant_heads))); + cold_requant_heads = static_cast(requant_heads); + } + } text_host_kv_page_stride = plan_host_kv_page_layout(decoder->text_kv.page_pool().geometry()).page_stride; text_kv_pages = std::make_unique( @@ -971,6 +982,12 @@ ProgramImplCore::ProgramImplCore(const LoadedModelData& model_in, const Sequence ProgramImplCore::~ProgramImplCore() noexcept { if (device.transfer_stream != nullptr) { (void)cudaStreamSynchronize(device.transfer_stream); } if (device.stream != nullptr) { (void)cudaStreamSynchronize(device.stream); } + if (cold_requant_codes != nullptr) { + (void)cudaFree(cold_requant_codes); + (void)cudaFree(cold_requant_scales); + cold_requant_codes = nullptr; + cold_requant_scales = nullptr; + } } std::vector ProgramImplCore::causal_score(PreparedPromptData&& prompt, @@ -9059,6 +9076,15 @@ void ProgramImplCore::start_sequence(std::uint32_t lane, SequenceState& sequence transaction.shared_source_index < shared_prefix_capacity && shared_prefix_slots[transaction.shared_source_index].role == SharedPrefixSlotRole::Catalogued; + // A retained source may carry cold-pool pages: warm them back into + // physical pages before any prefix fork touches the membership. + if (private_source_ready && cold_policy == ColdPolicy::Window) { + SequenceState& source = continuation_states[transaction.source_index]; + if (source.kv && source.kv->text.valid() && + text_kv_addresses->active(source.kv->text)) { + warm_cold_prefix(source, source.text_kv_valid); + } + } if (private_source_ready == shared_source_ready || transaction.reserved_state_count != state_slots || state_slots == 0 || !transaction.root_text_address || !transaction.text_prefix_fork || @@ -10101,6 +10127,12 @@ void ProgramImplCore::release_sequence_kv(SequenceState& sequence) noexcept { } if (text_kv_addresses) { (void)text_kv_addresses->release(sequence.kv->text); } sequence.kv.reset(); + if (decoder != nullptr) { + for (const auto& cold : sequence.cold_pages) { + decoder->text_kv.release_cold_slot(cold.slot); + } + } + sequence.cold_pages.clear(); if (host_kv_extents) { (void)host_kv_extents->release_unreferenced(); } } @@ -10147,48 +10179,55 @@ void ProgramImplCore::ordered_reset(SequenceState& sequence) { } -// Cold-pool maintenance: pack the retired tail of a sequence's text KV into +// Cold-pool maintenance: pack the retired prefix of a sequence's text KV into // raw entropy slots and detach those pages (sentinel entries in the block // table; physical pages return to the pool). Runs when the window policy is -// active and at least cold_keep_tokens lie behind the decode frontier. +// active and at least cold_keep_tokens lie behind the decode frontier. The +// decode kernels read cold pages straight from the slots, so no restore is +// needed on the steady-state path. void ProgramImplCore::enqueue_cold_compressions(SequenceState& sequence) { if (cold_policy != ColdPolicy::Window || !sequence.kv || decoder == nullptr || - !sequence.kv->text.valid()) { + !sequence.kv->text.valid() || cold_requant_codes == nullptr) { return; } KVAddressSpaceStore& store = *text_kv_addresses; KVAddressSpaceHandle text = sequence.kv->text; - const std::uint32_t total_pages = store.entitlement(text); + const std::uint32_t total_pages = store.mapped_pages(text); if (total_pages == 0) { return; } const std::uint32_t cold_pages = sequence.text_kv_valid > cold_keep_tokens ? (sequence.text_kv_valid - cold_keep_tokens) / kPagedKVPageSize : 0; - if (cold_pages == 0 || cold_pages >= total_pages) { return; } - - // Allocate cold slots for the tail (one slot per page; shared across heads - // through the flat slot addressing in the kernels). - const std::int32_t slot = decoder->text_kv.allocate_cold_slot(); - if (slot < 0) { return; } - - const std::uint32_t keep_pages = total_pages - cold_pages; - const std::int32_t kv_heads = - decoder->text_kv.batch_layer_view(0).num_kv_heads; - - // Pack every layer's cold tail pages into the raw slots. - for (std::uint32_t layer = 0; layer < decoder->text_kv.layers(); ++layer) { - const PagedKVBatchLayerView view = decoder->text_kv.batch_layer_view(layer); - const Tensor cold_slots = view.cold_slots; - if (cold_slots.data == nullptr) { continue; } - const bool int8_layer = view.dtype == DType::I8; - const auto k_mode = int8_layer ? ops::EntropyColdRequantMode::Int8G64 - : ops::EntropyColdRequantMode::Nvfp4G16; - const auto v_mode = int8_layer ? ops::EntropyColdRequantMode::Int8G64 - : ops::EntropyColdRequantMode::Iso3VG16; - for (std::uint32_t p = 0; p < cold_pages; ++p) { - const DeviceKVPageHandle ph = store.physical_page(text, keep_pages + p); - if (!ph.valid()) { continue; } + const std::uint32_t limit = cold_pages < total_pages ? cold_pages : total_pages; + if (limit == 0) { return; } + + const std::int32_t kv_heads = decoder->text_kv.batch_layer_view(0).num_kv_heads; + const std::uint32_t layers = decoder->text_kv.layers(); + // Cold slots carry requantized E2M1 planes (int8 -> E2M1 g64). The page + // stays hot unless every layer can pack, so mixed-dtype stacks skip. + for (std::uint32_t layer = 0; layer < layers; ++layer) { + if (decoder->text_kv.batch_layer_view(layer).dtype != DType::I8) { return; } + } + std::vector k_flags(static_cast(kv_heads)); + std::vector v_flags(static_cast(kv_heads)); + std::uint32_t compressed = 0; + + for (std::uint32_t page = sequence.cold_frontier; page < limit; ++page) { + if (!store.can_cold_transfer(text, page)) { continue; } + const std::int32_t slot = decoder->text_kv.allocate_cold_slot(); + if (slot < 0) { break; } // cold pool exhausted: keep the rest hot. + + bool success = true; + for (std::uint32_t layer = 0; layer < layers; ++layer) { + const PagedKVBatchLayerView view = decoder->text_kv.batch_layer_view(layer); + const Tensor cold_slots = view.cold_slots; + if (cold_slots.data == nullptr) { continue; } + if (view.dtype != DType::I8) { + success = false; // cold slots only carry int8 planes + break; + } + const DeviceKVPageHandle ph = store.physical_page(text, page); const std::int32_t physical = ph.index(); auto* k_codes = static_cast(view.k_pages.data) + physical * view.k_pages.nb[3]; @@ -10198,16 +10237,15 @@ void ProgramImplCore::enqueue_cold_compressions(SequenceState& sequence) { physical * view.k_scale_pages.nb[3]; auto* v_scales = static_cast(view.v_scale_pages.data) + physical * view.v_scale_pages.nb[3]; - const std::int64_t slot_off = - (static_cast(slot) * 2 * kv_heads) * cold_slots.nb[0]; - auto* k_slot = static_cast(cold_slots.data) + slot_off; + auto* k_slot = static_cast(cold_slots.data) + + static_cast(slot) * cold_slots.nb[3]; auto* v_slot = k_slot + cold_slots.nb[2]; auto* k_valid = static_cast(view.cold_slot_valid.data) + static_cast(slot) * view.cold_slot_valid.nb[2]; auto* v_valid = reinterpret_cast( reinterpret_cast(k_valid) + view.cold_slot_valid.nb[1]); ops::entropy_cold_requant_raw( - k_codes, k_scales, k_mode, kv_heads, 1, + k_codes, k_scales, ops::EntropyColdRequantMode::Int8G64, kv_heads, 1, static_cast(cold_requant_codes), static_cast(cold_requant_scales), device.stream); ops::cold_i8_slot_pack_raw( @@ -10215,7 +10253,7 @@ void ProgramImplCore::enqueue_cold_compressions(SequenceState& sequence) { static_cast(cold_requant_scales), kv_heads, 1, k_slot, k_valid, device.stream); ops::entropy_cold_requant_raw( - v_codes, v_scales, v_mode, kv_heads, 1, + v_codes, v_scales, ops::EntropyColdRequantMode::Int8G64, kv_heads, 1, static_cast(cold_requant_codes), static_cast(cold_requant_scales), device.stream); ops::cold_i8_slot_pack_raw( @@ -10223,20 +10261,110 @@ void ProgramImplCore::enqueue_cold_compressions(SequenceState& sequence) { static_cast(cold_requant_scales), kv_heads, 1, v_slot, v_valid, device.stream); } + if (!success) { + decoder->text_kv.release_cold_slot(slot); + continue; + } + device.synchronize(); + + // A slot only counts once every head's pack kernel committed its valid + // flag; otherwise the page would decode as garbage through the slot. + const Tensor cold_valid = decoder->text_kv.cold_slot_valid(0); + auto* k_valid = static_cast(cold_valid.data) + + static_cast(slot) * cold_valid.nb[2]; + auto* v_valid = reinterpret_cast( + reinterpret_cast(k_valid) + cold_valid.nb[1]); + CUDA_CHECK(cudaMemcpy(k_flags.data(), k_valid, + k_flags.size() * sizeof(std::int32_t), cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaMemcpy(v_flags.data(), v_valid, + v_flags.size() * sizeof(std::int32_t), cudaMemcpyDeviceToHost)); + const bool valid = + std::all_of(k_flags.begin(), k_flags.end(), + [](std::int32_t value) { return value != 0; }) && + std::all_of(v_flags.begin(), v_flags.end(), + [](std::int32_t value) { return value != 0; }); + if (!valid) { + decoder->text_kv.release_cold_slot(slot); + continue; + } + + // Publish the sentinel and return the physical page to the pool. + const std::int32_t entry = paged_kv_cold_entry(slot); + decoder->text_kv.execution_tables().publish_indices( + store.execution_row(text).handle(), page, std::span(&entry, 1), + device.stream); + store.transfer_to_cold(text, page); + sequence.cold_pages.emplace_back(page, slot); + sequence.cold_frontier = page + 1; + ++compressed; } - device.synchronize(); + if (compressed != 0) { + device.synchronize(); + std::fprintf(stderr, "[cold] compressed %u prefix pages (kept %u+)\n", compressed, + cold_keep_tokens); + } +} - // Detach the tail: shrink the address space, return the physical pages, - // and publish sentinel entries in the execution row. - const KVExecutionRowLease& row = store.execution_row(text); - std::vector sentinel(cold_pages, paged_kv_cold_entry(slot)); - decoder->text_kv.execution_tables().publish_indices( - row.handle(), keep_pages, sentinel, device.stream); - store.deactivate(text); - (void)store.activate(text, keep_pages, static_cast(row.index())); - device.synchronize(); - std::fprintf(stderr, "[cold] compressed tail %u pages -> slot %d (kept %u)\n", - cold_pages, slot, keep_pages); +// Warm-restore the cold prefix of a sequence (rewrite/resume paths only): the +// steady-state decode path reads cold pages directly from their slots, but a +// rewrite needs real physical pages so append/fork can mutate them again. +void ProgramImplCore::warm_cold_prefix(SequenceState& sequence, std::uint32_t end_page) { + if (cold_policy != ColdPolicy::Window || !sequence.kv || decoder == nullptr || + cold_requant_codes == nullptr || sequence.cold_pages.empty()) { + return; + } + KVAddressSpaceStore& store = *text_kv_addresses; + KVAddressSpaceHandle text = sequence.kv->text; + const std::uint32_t mapped = store.mapped_pages(text); + const std::uint32_t pages = std::min(end_page, mapped); + if (pages == 0) { return; } + + const int kv_heads = decoder->text_kv.batch_layer_view(0).num_kv_heads; + const std::uint32_t layers = decoder->text_kv.layers(); + std::uint32_t restored = 0; + for (std::uint32_t page = 0; page < pages; ++page) { + if (!store.cold_compressed(text, page)) { continue; } + auto entry = std::find_if(sequence.cold_pages.begin(), sequence.cold_pages.end(), + [page](const SequenceState::ColdPageEntry& e) { + return e.page == page; + }); + if (entry == sequence.cold_pages.end()) { continue; } + const std::int32_t slot = entry->slot; + const DeviceKVPageHandle physical = store.restore_from_cold(text, page); + const std::int32_t ph_index = physical.index(); + for (std::uint32_t layer = 0; layer < layers; ++layer) { + const PagedKVBatchLayerView view = decoder->text_kv.batch_layer_view(layer); + const Tensor cold_slots = view.cold_slots; + if (cold_slots.data == nullptr || view.dtype != DType::I8) { continue; } + auto* k_slot_base = static_cast(cold_slots.data); + auto* v_slot_base = k_slot_base + cold_slots.nb[2]; + auto* k_codes_i8 = static_cast(view.k_pages.data) + + static_cast(ph_index) * view.k_pages.nb[3]; + auto* v_codes_i8 = static_cast(view.v_pages.data) + + static_cast(ph_index) * view.v_pages.nb[3]; + auto* k_scales_h = static_cast( + static_cast(view.k_scale_pages.data) + + static_cast(ph_index) * view.k_scale_pages.nb[3]); + auto* v_scales_h = static_cast( + static_cast(view.v_scale_pages.data) + + static_cast(ph_index) * view.v_scale_pages.nb[3]); + ops::cold_i8_slot_restore_raw(k_slot_base + slot * cold_slots.nb[3], kv_heads, 1, + k_codes_i8, k_scales_h, device.stream); + ops::cold_i8_slot_restore_raw(v_slot_base + slot * cold_slots.nb[3], kv_heads, 1, + v_codes_i8, v_scales_h, device.stream); + } + decoder->text_kv.execution_tables().publish_indices( + store.execution_row(text).handle(), page, + std::span(&ph_index, 1), device.stream); + decoder->text_kv.release_cold_slot(slot); + sequence.cold_pages.erase(entry); + ++restored; + } + if (restored != 0) { + sequence.cold_frontier = 0; // pages are hot again; rescan from the front + device.synchronize(); + std::fprintf(stderr, "[cold] restored %u prefix pages\n", restored); + } } void ProgramImplCore::prepare_graphs() { @@ -11520,9 +11648,11 @@ ProgramImplCore::decode_raw(std::span lanes, std::span budgets, runtime::ExecutionTiming* failed_timing) { // Cold-pool maintenance at the round boundary (window policy only). - if (cold_policy == ColdPolicy::Window && lanes.size() == 1 && - sequences[lanes[0]].kv) { - enqueue_cold_compressions(sequences[lanes[0]]); + if (cold_policy == ColdPolicy::Window && lanes.size() == 1) { + SequenceState& sequence = active_sequence(lanes[0]); + if (sequence.kv) { + enqueue_cold_compressions(sequence); + } } if (speculative_backend == SpeculativeBackend::None) { return decode_ordinary_batch(lanes, budgets, failed_timing); diff --git a/src/targets/qwen3_6/impl/state/decoder_state.cpp b/src/targets/qwen3_6/impl/state/decoder_state.cpp index ddf0a03e23..7d60864357 100644 --- a/src/targets/qwen3_6/impl/state/decoder_state.cpp +++ b/src/targets/qwen3_6/impl/state/decoder_state.cpp @@ -80,12 +80,14 @@ DecoderStateLayout plan_decoder_state(LayoutBuilder& builder, const DecoderState const std::int32_t cold_slot_bytes = ops::kColdI8SlotBytes; if (spec.max_cold_pages != 0) { const std::uint32_t cold_pages = spec.max_cold_pages; + layout.text_kv.cold_slot_bytes = cold_slot_bytes; + layout.text_kv.max_cold_pages = spec.max_cold_pages; for (std::uint32_t layer = 0; layer < spec.full_attention_layers; ++layer) { - layout.cold_slots[layer] = builder.add_tensor( + layout.text_kv.cold_slots[layer] = builder.add_tensor( DType::U8, {cold_slot_bytes, static_cast(spec.kv_heads), 2, cold_pages}, 256, "cold slots L" + std::to_string(layer)); - layout.cold_slot_valid[layer] = builder.add_tensor( + layout.text_kv.cold_slot_valid[layer] = builder.add_tensor( DType::I32, {static_cast(spec.kv_heads), 2, cold_pages}, 256, "cold slot valid L" + std::to_string(layer)); } @@ -154,6 +156,9 @@ PagedKVLayerView PagedKVCache::layer_view(std::uint32_t layer, Tensor block_tabl .k_scale_pages = scaled ? pages_.plane(base + 2) : Tensor(), .v_scale_pages = scaled ? pages_.plane(base + 3) : Tensor(), .block_table = block_table, + .cold_slots = cold_slots_[layer], + .cold_slot_valid = cold_slot_valid_[layer], + .cold_slot_bytes = cold_slot_bytes_, .head_dim = head_dim_, .num_kv_heads = kv_heads_, .dtype = dtype_, @@ -174,6 +179,7 @@ PagedKVBatchLayerView PagedKVCache::batch_layer_view(std::uint32_t layer) const .block_tables = execution_tables_.matrix(), .cold_slots = cold_slots_[layer], .cold_slot_valid = cold_slot_valid_[layer], + .cold_slot_bytes = cold_slot_bytes_, .head_dim = head_dim_, .num_kv_heads = kv_heads_, .dtype = dtype_, From c0c17ff109df9c47e443282f3038f82237175b35 Mon Sep 17 00:00:00 2001 From: NInfer Agent Date: Sun, 30 Aug 2026 23:35:11 +0800 Subject: [PATCH 11/11] fix(kv): cold pool under host offload + multi-concurrency --- src/serve/generation_service.cpp | 3 + src/serve/serve_options.cpp | 17 +++ src/serve/serve_options.h | 3 + .../qwen3_6/impl/runtime/logical_kv_store.h | 28 ++++- src/targets/qwen3_6/impl/runtime/program.h | 2 + .../qwen3_6/impl/runtime/program_impl.h | 102 ++++++++++++------ .../qwen3_6/impl/runtime/request_plan_impl.h | 2 + 7 files changed, 120 insertions(+), 37 deletions(-) diff --git a/src/serve/generation_service.cpp b/src/serve/generation_service.cpp index 08d0bf7227..cba9f4575c 100644 --- a/src/serve/generation_service.cpp +++ b/src/serve/generation_service.cpp @@ -238,6 +238,9 @@ GenerationService::GenerationService(ServeOptions options, LoadProgress load_pro engine_options.use_cuda_graph = options_.use_cuda_graph; engine_options.speculative = options_.speculative; engine_options.context_cache = options_.context_cache; + engine_options.cold_policy = options_.cold_policy; + engine_options.cold_keep_tokens = options_.cold_keep_tokens; + engine_options.cold_host_bytes = options_.cold_host_bytes; engine_options.context_cost.preset_path = options_.context_cost_presets; engine_options.media_cache_bytes = options_.media_cache_bytes; engine_options.media_live_bytes = options_.media_live_bytes; diff --git a/src/serve/serve_options.cpp b/src/serve/serve_options.cpp index 66fef22197..928f11d1e1 100644 --- a/src/serve/serve_options.cpp +++ b/src/serve/serve_options.cpp @@ -77,6 +77,8 @@ std::string serve_usage_text(const char* argv0) { "[--request-log-jsonl FILE] " "[--response-store-max-records N] [--response-store-max-mib N] " "[--kv-dtype bf16|int8|fp8] [--spec mtp|dflash --draft-tokens N] " + "[--cold-policy none|window|host] [--cold-keep-tokens N] " + "[--cold-host-bytes N[g|m|k]] " "[--default-max-tokens N] [--default-thinking-budget N] " "[--vision] [--no-cuda-graph] [--no-prefix-reuse] " "[--lm-head-draft] [--no-thinking] [--preserve-thinking] [--cors] " @@ -258,6 +260,21 @@ ServeOptions parse_serve_options(int argc, char** argv) { options.device = parse_nonnegative_int(require_value("--device"), "device"); } else if (arg == "--kv-dtype") { options.kv_cache = parse_kv_dtype(require_value("--kv-dtype")); + } else if (arg == "--cold-policy") { + const std::string_view v = require_value("--cold-policy"); + if (v == "none" || v == "off") { options.cold_policy = ColdPolicy::None; } + else if (v == "window") { options.cold_policy = ColdPolicy::Window; } + else if (v == "host") { options.cold_policy = ColdPolicy::Host; } + else { throw std::invalid_argument("invalid cold-policy: " + std::string(v)); } + } else if (arg == "--cold-keep-tokens") { + options.cold_keep_tokens = + parse_u64(require_value("--cold-keep-tokens"), "cold-keep-tokens"); + if (options.cold_keep_tokens > std::numeric_limits::max()) { + throw std::invalid_argument("--cold-keep-tokens is out of range"); + } + } else if (arg == "--cold-host-bytes") { + options.cold_host_bytes = + parse_u64(require_value("--cold-host-bytes"), "cold-host-bytes"); } else if (arg == "--spec") { options.speculative.backend = product::parse_speculative_backend(require_value("--spec")); diff --git a/src/serve/serve_options.h b/src/serve/serve_options.h index c529fbc84f..8933e1b0a7 100644 --- a/src/serve/serve_options.h +++ b/src/serve/serve_options.h @@ -47,6 +47,9 @@ struct ServeOptions { bool enable_vision = false; bool use_cuda_graph = true; bool allow_prefix_reuse = true; + ColdPolicy cold_policy = ColdPolicy::None; + std::uint32_t cold_keep_tokens = 128; + std::uint64_t cold_host_bytes = 4ULL << 30; bool enable_thinking = true; // default thinking mode for the generation prompt (--no-thinking opts out) bool preserve_thinking = false; diff --git a/src/targets/qwen3_6/impl/runtime/logical_kv_store.h b/src/targets/qwen3_6/impl/runtime/logical_kv_store.h index a1b8aa8493..074e2d7dca 100644 --- a/src/targets/qwen3_6/impl/runtime/logical_kv_store.h +++ b/src/targets/qwen3_6/impl/runtime/logical_kv_store.h @@ -764,11 +764,13 @@ class LogicalKVPageStore { // physical page to the pool) while keeping the descriptor, so the address // membership and its block-table slot stay stable. The page's contents // live in a fixed cold slot; restore_from_cold brings them back on demand. + // An existing host replica is deliberately KEPT: the checkpoint restore + // path (prepare_kv_restores) needs a restorable source, and the host copy + // predates the cold requant, so it doubles as the higher-fidelity backup. void transfer_to_cold(LogicalKVPageHandle handle, DeviceKVPageReservation& reservation) { Page& page = require(handle); - if (page.references != 1 || page.writer_references != 1 || page.source_pins != 0 || - page.destination_pinned || page.host_replica || !page.device_replica || - page.cold_compressed) { + if (page.source_pins != 0 || page.destination_pinned || !page.device_replica || + page.cold_compressed || page.writer_references != 0 || page.references == 0) { throw std::logic_error("logical KV page is not cold-transferable"); } physical_->dematerialize_one(reservation, std::move(*page.device_replica)); @@ -776,13 +778,27 @@ class LogicalKVPageStore { page.cold_compressed = true; } + // Like can_dematerialize but for the cold pool: committed history pages + // (writer_references == 0) qualify, host replicas do not block (the + // transfer drops them), and catalogued/fork references are fine because + // the fork path warms cold pages before materializing them. + [[nodiscard]] bool can_cold_transfer(LogicalKVPageHandle handle) const noexcept { + if (!valid(handle)) { return false; } + const Page& page = pages_[handle.index_]; + return page.references != 0 && page.writer_references == 0 && page.source_pins == 0 && + !page.destination_pinned && page.device_replica.has_value() && + !page.cold_compressed; + } + [[nodiscard]] bool cold_compressed(LogicalKVPageHandle handle) const noexcept { return valid(handle) && pages_[handle.index_].cold_compressed; } // Cold-pool restore: allocate a fresh physical page for a cold descriptor // and hand it back so the caller can repopulate it from the cold slot. - // The descriptor keeps its membership position and reference counts. + // The descriptor keeps its membership position and reference counts. Any + // host replica is stale after the cold restore (the device copy is now + // authoritative) and is dropped; the offload path re-attaches it later. [[nodiscard]] DeviceKVPageHandle restore_from_cold(LogicalKVPageHandle handle, DeviceKVPageReservation& reservation) { Page& page = require(handle); @@ -800,6 +816,7 @@ class LogicalKVPageStore { page.device_replica.emplace(std::move(lease)); page.cold_compressed = false; page.content_epoch = next_epoch(page.content_epoch); + page.host_replica.reset(); return lease_handle; } @@ -1706,13 +1723,14 @@ class KVAddressSpaceStore { // Whether the page may be detached into a cold slot right now: not already // cold, exclusively referenced by its writer, and free of pins/fork ties. + // Host replicas do not block (transfer_to_cold drops them). [[nodiscard]] bool can_cold_transfer(KVAddressSpaceHandle handle, std::uint32_t logical_page) const noexcept { if (!valid(handle)) { return false; } const Address& address = addresses_[handle.index_]; if (logical_page >= address.page_count) { return false; } const LogicalKVPageHandle& logical = membership(address, logical_page); - return !pages_->cold_compressed(logical) && pages_->can_dematerialize(logical); + return pages_->can_cold_transfer(logical); } void transfer_to_cold(KVAddressSpaceHandle handle, std::uint32_t logical_page) { diff --git a/src/targets/qwen3_6/impl/runtime/program.h b/src/targets/qwen3_6/impl/runtime/program.h index f21b1c0cfc..3acd18ece8 100644 --- a/src/targets/qwen3_6/impl/runtime/program.h +++ b/src/targets/qwen3_6/impl/runtime/program.h @@ -709,6 +709,8 @@ class ProgramImplCore { std::uint32_t cold_requant_heads = 0; void enqueue_cold_compressions(SequenceState& sequence); void warm_cold_prefix(SequenceState& sequence, std::uint32_t end_page); + void restore_cold_page(SequenceState& sequence, std::uint32_t page, std::int32_t slot, + const DeviceKVPageHandle& physical); std::size_t vision_handoff_peak_bytes = 0; diff --git a/src/targets/qwen3_6/impl/runtime/program_impl.h b/src/targets/qwen3_6/impl/runtime/program_impl.h index 2852d8e1b6..b10322e2a3 100644 --- a/src/targets/qwen3_6/impl/runtime/program_impl.h +++ b/src/targets/qwen3_6/impl/runtime/program_impl.h @@ -1892,6 +1892,9 @@ ProgramImplCore::checkpoint_restore_requirements(const SequenceKVBundle& kv, for (std::uint32_t page = 0; page < required; ++page) { const LogicalKVPageHandle logical = addresses.logical_page(address, page); if (pages.device_resident(logical)) { continue; } + // Cold-pool pages restore in place from their raw slots; they + // need no host transfer. + if (pages.cold_compressed(logical)) { continue; } if (!pages.host_resident(logical)) { throw std::logic_error("checkpoint KV page has no restorable replica"); } @@ -4745,6 +4748,25 @@ void ProgramImplCore::prepare_materialization(MaterializationTransaction& transa for (std::uint32_t page = 0; page < mapped; ++page) { const LogicalKVPageHandle logical = addresses.logical_page(address, page); if (pages.device_resident(logical)) { continue; } + if (pages.cold_compressed(logical)) { + // Cold-pool page: its restorable source is the raw cold + // slot, not a host replica. Restore synchronously here + // (same reservation); the activation publish + // (publish_membership) republishes the physical mapping. + if (source_state == nullptr) { + throw std::logic_error("cold checkpoint page has no source bookkeeping"); + } + auto entry = std::find_if( + source_state->cold_pages.begin(), source_state->cold_pages.end(), + [page](const SequenceState::ColdPageEntry& e) { return e.page == page; }); + if (entry == source_state->cold_pages.end()) { + throw std::logic_error("cold checkpoint page has no slot record"); + } + const DeviceKVPageHandle restored = + pages.restore_from_cold(logical, reservation); + restore_cold_page(*source_state, page, entry->slot, restored); + continue; + } if (!pages.host_resident(logical) || !host_kv_extents) { throw std::logic_error("checkpoint KV page has no restorable replica"); } @@ -10305,6 +10327,44 @@ void ProgramImplCore::enqueue_cold_compressions(SequenceState& sequence) { } } +// Restore one cold page's data from its raw slot into the physical page and +// release the slot. Shared by the rewrite warm path and the checkpoint +// restore path (which must repopulate cold pages without a host replica). +void ProgramImplCore::restore_cold_page(SequenceState& sequence, std::uint32_t page, + std::int32_t slot, const DeviceKVPageHandle& physical) { + const int kv_heads = decoder->text_kv.batch_layer_view(0).num_kv_heads; + const std::uint32_t layers = decoder->text_kv.layers(); + const std::int32_t ph_index = physical.index(); + for (std::uint32_t layer = 0; layer < layers; ++layer) { + const PagedKVBatchLayerView view = decoder->text_kv.batch_layer_view(layer); + const Tensor cold_slots = view.cold_slots; + if (cold_slots.data == nullptr || view.dtype != DType::I8) { continue; } + auto* k_slot_base = static_cast(cold_slots.data); + auto* v_slot_base = k_slot_base + cold_slots.nb[2]; + auto* k_codes_i8 = static_cast(view.k_pages.data) + + static_cast(ph_index) * view.k_pages.nb[3]; + auto* v_codes_i8 = static_cast(view.v_pages.data) + + static_cast(ph_index) * view.v_pages.nb[3]; + auto* k_scales_h = static_cast( + static_cast(view.k_scale_pages.data) + + static_cast(ph_index) * view.k_scale_pages.nb[3]); + auto* v_scales_h = static_cast( + static_cast(view.v_scale_pages.data) + + static_cast(ph_index) * view.v_scale_pages.nb[3]); + ops::cold_i8_slot_restore_raw(k_slot_base + slot * cold_slots.nb[3], kv_heads, 1, + k_codes_i8, k_scales_h, device.stream); + ops::cold_i8_slot_restore_raw(v_slot_base + slot * cold_slots.nb[3], kv_heads, 1, + v_codes_i8, v_scales_h, device.stream); + } + decoder->text_kv.release_cold_slot(slot); + auto entry = std::find_if(sequence.cold_pages.begin(), sequence.cold_pages.end(), + [page](const SequenceState::ColdPageEntry& e) { + return e.page == page; + }); + if (entry != sequence.cold_pages.end()) { sequence.cold_pages.erase(entry); } + sequence.cold_frontier = 0; // pages are hot again; rescan from the front +} + // Warm-restore the cold prefix of a sequence (rewrite/resume paths only): the // steady-state decode path reads cold pages directly from their slots, but a // rewrite needs real physical pages so append/fork can mutate them again. @@ -10319,9 +10379,7 @@ void ProgramImplCore::warm_cold_prefix(SequenceState& sequence, std::uint32_t en const std::uint32_t pages = std::min(end_page, mapped); if (pages == 0) { return; } - const int kv_heads = decoder->text_kv.batch_layer_view(0).num_kv_heads; - const std::uint32_t layers = decoder->text_kv.layers(); - std::uint32_t restored = 0; + std::uint32_t restored = 0; for (std::uint32_t page = 0; page < pages; ++page) { if (!store.cold_compressed(text, page)) { continue; } auto entry = std::find_if(sequence.cold_pages.begin(), sequence.cold_pages.end(), @@ -10329,39 +10387,15 @@ void ProgramImplCore::warm_cold_prefix(SequenceState& sequence, std::uint32_t en return e.page == page; }); if (entry == sequence.cold_pages.end()) { continue; } - const std::int32_t slot = entry->slot; const DeviceKVPageHandle physical = store.restore_from_cold(text, page); const std::int32_t ph_index = physical.index(); - for (std::uint32_t layer = 0; layer < layers; ++layer) { - const PagedKVBatchLayerView view = decoder->text_kv.batch_layer_view(layer); - const Tensor cold_slots = view.cold_slots; - if (cold_slots.data == nullptr || view.dtype != DType::I8) { continue; } - auto* k_slot_base = static_cast(cold_slots.data); - auto* v_slot_base = k_slot_base + cold_slots.nb[2]; - auto* k_codes_i8 = static_cast(view.k_pages.data) + - static_cast(ph_index) * view.k_pages.nb[3]; - auto* v_codes_i8 = static_cast(view.v_pages.data) + - static_cast(ph_index) * view.v_pages.nb[3]; - auto* k_scales_h = static_cast( - static_cast(view.k_scale_pages.data) + - static_cast(ph_index) * view.k_scale_pages.nb[3]); - auto* v_scales_h = static_cast( - static_cast(view.v_scale_pages.data) + - static_cast(ph_index) * view.v_scale_pages.nb[3]); - ops::cold_i8_slot_restore_raw(k_slot_base + slot * cold_slots.nb[3], kv_heads, 1, - k_codes_i8, k_scales_h, device.stream); - ops::cold_i8_slot_restore_raw(v_slot_base + slot * cold_slots.nb[3], kv_heads, 1, - v_codes_i8, v_scales_h, device.stream); - } + restore_cold_page(sequence, page, entry->slot, physical); decoder->text_kv.execution_tables().publish_indices( store.execution_row(text).handle(), page, std::span(&ph_index, 1), device.stream); - decoder->text_kv.release_cold_slot(slot); - sequence.cold_pages.erase(entry); ++restored; } if (restored != 0) { - sequence.cold_frontier = 0; // pages are hot again; rescan from the front device.synchronize(); std::fprintf(stderr, "[cold] restored %u prefix pages\n", restored); } @@ -11648,10 +11682,14 @@ ProgramImplCore::decode_raw(std::span lanes, std::span budgets, runtime::ExecutionTiming* failed_timing) { // Cold-pool maintenance at the round boundary (window policy only). - if (cold_policy == ColdPolicy::Window && lanes.size() == 1) { - SequenceState& sequence = active_sequence(lanes[0]); - if (sequence.kv) { - enqueue_cold_compressions(sequence); + // Every active sequence maintains its own retired prefix; multi-lane + // batches compress each lane's pages independently. + if (cold_policy == ColdPolicy::Window) { + for (const std::uint32_t lane : lanes) { + SequenceState& sequence = active_sequence(lane); + if (sequence.kv) { + enqueue_cold_compressions(sequence); + } } } if (speculative_backend == SpeculativeBackend::None) { diff --git a/src/targets/qwen3_6/impl/runtime/request_plan_impl.h b/src/targets/qwen3_6/impl/runtime/request_plan_impl.h index 504725148e..1b8a98de39 100644 --- a/src/targets/qwen3_6/impl/runtime/request_plan_impl.h +++ b/src/targets/qwen3_6/impl/runtime/request_plan_impl.h @@ -802,6 +802,8 @@ std::optional ProgramImplCore::inspect_lane( for (std::uint32_t page = 0; page < required; ++page) { const LogicalKVPageHandle logical = addresses.logical_page(address, page); if (pages.device_resident(logical)) { continue; } + // Cold-pool pages restore in place from their raw slots. + if (pages.cold_compressed(logical)) { continue; } if (!pages.host_resident(logical)) { throw std::logic_error("checkpoint KV page has no restorable replica"); }