diff --git a/.agents/issues/KV-FP8/ISSUE-LOCAL-01M3QG4WWC3X0C84M9PWQZAJ10.md b/.agents/issues/KV-FP8/ISSUE-LOCAL-01M3QG4WWC3X0C84M9PWQZAJ10.md new file mode 100644 index 000000000..1f4ababbf --- /dev/null +++ b/.agents/issues/KV-FP8/ISSUE-LOCAL-01M3QG4WWC3X0C84M9PWQZAJ10.md @@ -0,0 +1,27 @@ +ID: ISSUE-LOCAL-01M3QG4WWC3X0C84M9PWQZAJ10 +Title: fp8 KV prefill runs the scalar CUDA-core flash kernel at 3.7x the bf16 cache, because FA-2 admits only bf16 q/KV/out while the fp8 store presents f32 +Row: KV-FP8 +State: OPEN +Kind: feature +GitHub: - +Mirror: PENDING +Availability: FULL +Created: 2026-09-29 +Updated: 2026-09-29 +Closed: - + +## Problem + +MEASURED on the 27B NVFP4 arm at 8k prefill: 28.0s with the fp8 KV cache against 7.8s with the same kernel on a bf16 cache, a 3.7x gap. nsys attributes it entirely to the kernel choice: 25.1s of 36s sits in `PagedFlashKernel`, the scalar CUDA-core arm, while the bf16 store runs the FA-2 split-KV kernel at 2.9 ms/layer. CAUSE: the model presents f32 q/out for a non-bf16 store (`GdnOutDType`/the KV-store route), and the vendored FA-2 admission requires bf16 q, bf16 KV and bf16 out, so every non-bf16 cache falls off the tensor-core ladder onto the per-element fp8 dequant. The W2 CUDA arm of this row (issue #1593) landed the fp8 store and the read dequant for CORRECTNESS, so this is the prefill-PERFORMANCE half the row never had. THE FIX: dequantize the fp8 cache ONCE per layer into a dense bf16 scratch and run the normal bf16 dispatch on it. The scratch is one block per request with block_size = max_seq and an identity block table, so the kernel s paged address IS the dense address and NO attention kernel changes. A shared `KvCachePresentsBf16` helper presents bf16 for an fp8 store, so FA-2 admits with zero cast kernels and any model that gains an fp8 store inherits the CUDA path unchanged. `VT_ATTN_FP8_DENSE=0` restores the per-read dequant for a same-binary A/B. TWO EARLIER SHAPES OF THE SAME LEVER WERE TRIED AND REJECTED, and the record keeps them: (a) routing fp8 through the bf16 WMMA ladder with an fp8->bf16 cast inside the K/V staging (8k 28.0 -> 9.1s) still pays a per-element dequant on every re-stream, and (b) converting the staging with `__nv_cvt_fp8_to_halfraw` made that conversion cheap (bit-identical) but did not remove it. The dense dequant removes it entirely for prefill. Numbers after: 32k prefill 34.9s / 938 tok/s against 36.4s / 901 for the bf16 store and 39.9s / 822 for llama.cpp s q8_0 KV, at 20.5 GiB against 22.5; 8k 7.6s / 1000 tok/s; decode with MTP n=3 38.6 vs 37.4. Op-level parity against the f32 reference is 1.9e-6 max abs err (the exact dequantized values), all 33 paged-attention cases pass, and benchmarks/paged_attn_prefill_ab.cpp is the isolated sweep that separated the kernel cost from the engine. + +## Resolution + +- 2026-09-30: review repairs on the W1 PR (mudler/vllm.cpp#3360). The dense bf16 + scratch and the identity table are owned by stream-aware scope guards, so the + steps that can throw between the two allocations (identity copy, dequant + launch, the bf16 dispatch, the vectors between them) release both buffers on + the same stream. A new case in `test_ops_paged_attn` injects the identity + allocation failure and the identity copy failure and checks the CUDA pool's + used bytes return to the pre-call value without the guard masking the + original exception. The benchmark header no longer claims a bf16 parity + comparison it does not perform; parity is gated by the test suite. diff --git a/.agents/specs/fp8-kv-prefill-dense-dequant.md b/.agents/specs/fp8-kv-prefill-dense-dequant.md new file mode 100644 index 000000000..17942e36f --- /dev/null +++ b/.agents/specs/fp8-kv-prefill-dense-dequant.md @@ -0,0 +1,96 @@ +# fp8 KV prefill on the bf16 dispatch: one dense dequant per layer — ISSUE-LOCAL-01M3QG4WWC3X0C84M9PWQZAJ10 + +The fp8 KV cache read sent every prefill to the scalar CUDA-core flash kernel, +3.7x slower than the same kernel on a bf16 cache, because FA-2 admits only bf16 +q/KV/out and the fp8 store presents f32. This is the prefill-performance half of +`KV-FP8`; W2 (issue [#1593](https://github.com/mudler/vllm.cpp/issues/1593)) +landed the fp8 store and read-dequant for correctness. + +Issue: [ISSUE-LOCAL-01M3QG4WWC3X0C84M9PWQZAJ10](../issues/KV-FP8/ISSUE-LOCAL-01M3QG4WWC3X0C84M9PWQZAJ10.md). +Owning row: `KV-FP8` ([engine-matrix.md](../engine-matrix.md)); the row's spec is +[fp8-kv-cache.md](fp8-kv-cache.md). + +## Premise, grounded + +| Where (line anchors at this branch's base, `b45a94273`) | What | +|---|---| +| `src/vt/cuda/cuda_paged_attn.cu:148-197` | The fp8 K/V read with the per-element dequant folded in (`Fp8E4M3ToF32Dev`), the W2 arm. | +| `src/vt/cuda/cuda_paged_attn.cu:2909`, `:2954` | The dispatch comments that name the split: the tensor-core ladder for bf16, the f32-q/out scalar arm otherwise. | +| `include/vllm/model_executor/models/kv_cache_route.h:40,73-78` | The store/read route that hands `kv_cache_dtype` to the backend. | +| `src/vllm/model_executor/models/qwen3_5.cpp` (the attention preamble) | The model-side arm this change routes. | +| `tests/vt/test_ops_paged_attn.cpp` | 33 cases on the base, including the W2 fp8 parity case. | + +Measured on the 27B NVFP4 arm at 8k prefill: 28.0 s fp8 vs 7.8 s bf16 (3.7x), +25.1 s of 36 s in `PagedFlashKernel`, against +2.9 ms/layer on the FA-2 split-KV kernel for a bf16 store. + +## Design + +Dequantize the fp8 cache ONCE per layer into a dense bf16 scratch, then run the +shipped bf16 dispatch on it: + +- The scratch is **one block per request with `block_size = max_seq` and an + identity block table**, so the kernel's paged address IS the dense address and + no attention kernel changes. +- Each scratch allocation is owned by a scope guard from the moment it + succeeds. The identity allocation and copy, the dequant launch, the bf16 + dispatch it feeds, and the vectors between them can all throw, and the guard + frees on the SAME stream, so the free is ordered behind the work that reads + the buffer. The success path still checks its own frees explicitly. +- The new shared `KvCachePresentsBf16` helper presents bf16 for an fp8 store, so + FA-2 admits with zero cast kernels. It is model-agnostic: any model that gains + an fp8 store inherits the CUDA path unchanged. +- `VT_ATTN_FP8_DENSE=0` restores the per-read dequant for a same-binary A/B. + +**Two rejected levers are recorded rather than deleted**, because both are +plausible and neither survives the numbers: + +1. Routing fp8 through the bf16 WMMA ladder with an fp8→bf16 cast inside the K/V + staging (8k 28.0 → 9.1 s, `VT_ATTN_FP8_WMMA`). It keeps the per-element + dequant on every re-stream, so the gap to bf16 (7.8 s) only narrows. +2. Converting the staging with `__nv_cvt_fp8_to_halfraw` instead of the software + decode — bit-identical and cheaper, but still per re-stream. + +The dense dequant removes the per-read dequant entirely for prefill, which is +why it is the shipped shape. The bench that separated kernel cost from engine +cost is `benchmarks/paged_attn_prefill_ab.cpp`, landed with this change. + +## Tests + +`tests/vt/test_ops_paged_attn.cpp` gains the fp8-dense parity case: the fp8 cache +against the f32 reference at **1.9e-6 max abs err** — the exact dequantized +values, tighter than the bf16-compute envelope — with the rest of the 33-case +suite unchanged. A second case injects the two failures a healthy device cannot +produce on demand (the identity-table allocation, and the identity-table copy +after both allocations succeeded) and asserts that the CUDA memory pool's used +bytes return to their pre-call value after the exception unwinds. It also +asserts the propagated message is the failing `Check`, so the guard's destructor +cannot mask the original exception. + +## What this does NOT claim + +- **No speed claim in this PR.** The numbers above are the author's measurement + on the local `sm_120a` card with host embedding on (`VT_HOST_EMBEDDING=1`), and + they are recorded as the shape's evidence, not as an operator gate. A + same-binary A/B under the GPU lease is owed by the operator, as the helper + template requires. +- The 20.5 GiB-vs-22.5 GiB peak comparison and the llama.cpp q8_0 denominator + (39.9 s / 822 tok/s) are the author's run of the same recipe; they are not + re-measured here. + +## Gates + +- `ctest --test-dir build -R test_ops_paged_attn` on a CUDA build + (`-DVLLM_CPP_CUDA=ON -DVLLM_CPP_CUTLASS_FETCH=ON` for FA-2). +- `VT_ATTN_FP8_DENSE=0` returns the base behaviour (the A/B arm). +- The full `ctest --test-dir build` on the same build. +- `scripts/agent-preflight.sh --staged`. + +## Owed + +- The operator's same-binary A/B under lease, with the 27B NVFP4 arm, both arms + in one binary, and the recipe written into `docs/BENCHMARKS.md`. +- A `sm_120a` re-measurement: the author's numbers are from the local consumer + card, and the fleet gate model runs elsewhere. +- Execution of the scratch-ownership case, which needs a CUDA device; the CUDA + lane is its gate. diff --git a/CMakeLists.txt b/CMakeLists.txt index 88716c8ad..51b4065e2 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -2905,6 +2905,14 @@ target_include_directories(vllm_music3_vocoder_conv_ab SYSTEM PRIVATE target_compile_features(vllm_music3_vocoder_conv_ab PRIVATE cxx_std_20) vllm_cpp_set_warnings(vllm_music3_vocoder_conv_ab) +# Paged-attention prefill A/B (paged vs dense KV source, bf16 vs fp8 cache). +# The isolated sweep behind the fp8 KV prefill dense dequant (KV-FP8): one +# vt::PagedAttention call per arm, so the number is the kernel cost alone. +add_executable(vllm_paged_attn_prefill_ab benchmarks/paged_attn_prefill_ab.cpp) +target_link_libraries(vllm_paged_attn_prefill_ab PRIVATE vllm) +target_compile_features(vllm_paged_attn_prefill_ab PRIVATE cxx_std_20) +vllm_cpp_set_warnings(vllm_paged_attn_prefill_ab) + # ── The `vt::Conv1d` decomposition probe (#672, #1334) ─────────────────────── # `tools/bench/conv1d_scaling_probe.cpp` is the instrument behind the scaling # curve and the residency ablation in `.agents/specs/vt-conv1d-time-block.md`. diff --git a/benchmarks/paged_attn_prefill_ab.cpp b/benchmarks/paged_attn_prefill_ab.cpp new file mode 100644 index 000000000..9734cd652 --- /dev/null +++ b/benchmarks/paged_attn_prefill_ab.cpp @@ -0,0 +1,188 @@ +// Paged-attention PREFILL A/B: paged vs dense KV source, bf16 vs fp8 cache, at +// the 27B full-attention shape (hq=32, hk=4, d=256, block_size=32). +// +// WHY AN ISOLATED SWEEP EXISTS: the engine-level fp8 prefill numbers (28s vs 8s +// at 8k) conflate the attention kernel, the per-layer dequant scratch, the D2H +// sync and the driver allocator. This runs ONE vt::PagedAttention call per arm +// on synthetic data, so the number is the kernel cost and nothing else. +// +// ARMS (one process per arm; the dispatch knobs are process-static): +// BENCH_KV=bf16 paged bf16 cache (the bf16 engine baseline) +// BENCH_KV=fp8 paged fp8 cache (dense scratch ON by default) +// BENCH_KV=fp8 VT_ATTN_FP8_DENSE=0 per-read dequant inside the kernel +// BENCH_KV=bf16 BENCH_DENSE=1 dense bf16 + identity table (layout control) +// +// ONE arm per process, so this binary cannot compare arms to each other. It +// prints a checksum of the arm's output, which makes a wildly wrong arm visible +// in the logs; it is not a parity gate. Parity is gated by +// `tests/vt/test_ops_paged_attn.cpp`, where each arm is compared against an f32 +// reference on the exact dequantized values (< 5e-2 max abs err). +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "vt/backend.h" +#include "vt/dtype.h" +#include "vt/fp8_kv.h" +#include "vt/ops.h" + +using vt::Backend; +using vt::DeviceType; +using vt::DType; +using vt::Fp8KVCacheDataType; +using vt::PagedAttentionArgs; +using vt::Queue; +using vt::Tensor; + +namespace { + +Tensor MakeT(void* data, DType dt, const std::vector& shape) { + Tensor t; + t.data = data; + t.dtype = dt; + t.device = vt::Device{DeviceType::kCUDA, 0}; + t.rank = static_cast(shape.size()); + int64_t stride = 1; + for (int i = t.rank - 1; i >= 0; --i) { + t.shape[i] = shape[static_cast(i)]; + t.stride[i] = stride; + stride *= shape[static_cast(i)]; + } + return t; +} + +std::vector RandF32(size_t n, uint32_t seed) { + std::vector v(n); + uint32_t s = seed; + for (auto& x : v) { + s = s * 1664525u + 1013904223u; + x = (static_cast(s >> 8) / static_cast(1u << 24)) * 4.0f - 2.0f; + } + return v; +} + +struct Buf { + Backend& b; + void* p = nullptr; + size_t bytes = 0; + Buf(Backend& backend, size_t n) : b(backend), bytes(n) { p = b.Alloc(n == 0 ? 1 : n); } + ~Buf() { b.Free(p); } + Buf(const Buf&) = delete; + Buf& operator=(const Buf&) = delete; +}; + +} // namespace + +int main(int argc, char** argv) { + const int64_t T = argc > 1 ? std::atoll(argv[1]) : 8192; + const int64_t Hq = argc > 2 ? std::atoll(argv[2]) : 32; + const int64_t Hk = argc > 3 ? std::atoll(argv[3]) : 4; + const int64_t D = argc > 4 ? std::atoll(argv[4]) : 256; + const int64_t BS = argc > 5 ? std::atoll(argv[5]) : 32; + const int reps = argc > 6 ? std::atoi(argv[6]) : 3; + const char* kv_env = std::getenv("BENCH_KV"); + const bool fp8 = kv_env != nullptr && std::strcmp(kv_env, "fp8") == 0; + const char* dense_env = std::getenv("BENCH_DENSE"); + const bool dense_bf16 = !fp8 && dense_env != nullptr && dense_env[0] == '1'; + const float scale = std::pow(static_cast(D), -0.5f); + const int64_t N = T / BS; // one request, identity block order + + Backend& gpu = vt::GetBackend(DeviceType::kCUDA); + Queue q = gpu.CreateQueue(); + + auto q_host = RandF32(static_cast(T * Hq * D), 2024); + auto k_host = RandF32(static_cast(N * BS * Hk * D), 137); + auto v_host = RandF32(static_cast(N * BS * Hk * D), 179); + std::vector block_table(static_cast(N)); + for (int64_t i = 0; i < N; ++i) block_table[static_cast(i)] = static_cast(i); + std::vector seq_lens = {static_cast(T)}; + std::vector qsl = {0, static_cast(T)}; + + Buf dq(gpu, static_cast(T * Hq * D) * sizeof(float)); + Buf dbt(gpu, block_table.size() * sizeof(int32_t)); + Buf dsl(gpu, seq_lens.size() * sizeof(int32_t)); + Buf dqsl(gpu, qsl.size() * sizeof(int32_t)); + Buf dout(gpu, static_cast(T * Hq * D) * sizeof(float)); + gpu.Copy(q, dq.p, q_host.data(), dq.bytes); + gpu.Copy(q, dbt.p, block_table.data(), dbt.bytes); + gpu.Copy(q, dsl.p, seq_lens.data(), dsl.bytes); + gpu.Copy(q, dqsl.p, qsl.data(), dqsl.bytes); + + // Cache: paged [N, BS, Hk, D] (bf16 or fp8), or dense [1, T, Hk, D] bf16. + const int64_t cache_elems = dense_bf16 ? T * Hk * D : N * BS * Hk * D; + Buf kcache(gpu, static_cast(cache_elems) * (fp8 ? 1 : 2)); + Buf vcache(gpu, static_cast(cache_elems) * (fp8 ? 1 : 2)); + Tensor kt, vt; + if (fp8) { + std::vector k8(k_host.size()), v8(v_host.size()); + for (size_t i = 0; i < k_host.size(); ++i) { + k8[i] = vt::StoreKvFp8E4M3(k_host[i], 1.0f); + v8[i] = vt::StoreKvFp8E4M3(v_host[i], 1.0f); + } + gpu.Copy(q, kcache.p, k8.data(), k8.size()); + gpu.Copy(q, vcache.p, v8.data(), v8.size()); + kt = MakeT(kcache.p, DType::kI8, {N, BS, Hk, D}); + vt = MakeT(vcache.p, DType::kI8, {N, BS, Hk, D}); + } else { + std::vector kb(k_host.size()), vb(v_host.size()); + for (size_t i = 0; i < k_host.size(); ++i) { + kb[i] = vt::F32ToBF16(k_host[i]); + vb[i] = vt::F32ToBF16(v_host[i]); + } + gpu.Copy(q, kcache.p, kb.data(), kb.size() * 2); + gpu.Copy(q, vcache.p, vb.data(), vb.size() * 2); + const std::vector shape = + dense_bf16 ? std::vector{1, T, Hk, D} : std::vector{N, BS, Hk, D}; + kt = MakeT(kcache.p, DType::kBF16, shape); + vt = MakeT(vcache.p, DType::kBF16, shape); + } + + // Dense arm: one block per request, block_size = T, identity table. + int32_t* bt_ptr = static_cast(dbt.p); + const int64_t bt_elems = N; + const int64_t block_size = dense_bf16 ? T : BS; + + PagedAttentionArgs args{scale, /*causal=*/true}; + if (fp8) { + args.kv_cache_dtype = Fp8KVCacheDataType::kFp8E4M3; + args.k_scale = 1.0f; + args.v_scale = 1.0f; + } + + Tensor q_t = MakeT(dq.p, DType::kF32, {T, Hq, D}); + Tensor out_t = MakeT(dout.p, DType::kF32, {T, Hq, D}); + Tensor bt_t = MakeT(bt_ptr, DType::kI32, {1, dense_bf16 ? 1 : bt_elems}); + Tensor sl_t = MakeT(dsl.p, DType::kI32, {1}); + Tensor qsl_t = MakeT(dqsl.p, DType::kI32, {2}); + + std::vector got(static_cast(T * Hq * D)); + double ms = 0.0; + for (int r = 0; r < reps + 2; ++r) { + const auto t0 = std::chrono::steady_clock::now(); + vt::PagedAttention(q, out_t, q_t, kt, vt, bt_t, sl_t, qsl_t, args); + gpu.Synchronize(q); + const double dt = + std::chrono::duration(std::chrono::steady_clock::now() - t0).count(); + if (r >= 2) ms += dt; + } + ms /= reps; + gpu.Copy(q, got.data(), dout.p, got.size() * sizeof(float)); + gpu.Synchronize(q); + + const char* arm = fp8 ? (std::getenv("VT_ATTN_FP8_DENSE") != nullptr ? "fp8/per-read" : "fp8/dense") + : (dense_bf16 ? "bf16/dense" : "bf16/paged"); std::printf("arm=%-12s T=%lld hq=%lld hk=%lld d=%lld bs=%lld %.1f ms %.0f tok/s\n", arm, + static_cast(T), static_cast(Hq), static_cast(Hk), + static_cast(D), static_cast(block_size), ms, T / (ms / 1000.0)); + // Checksum so a wrong-but-fast arm is visible. + double sum = 0.0; + for (float x : got) sum += x; + std::printf(" checksum=%.6f\n", sum); + gpu.DestroyQueue(q); + return 0; +} diff --git a/include/vllm/model_executor/models/kv_cache_route.h b/include/vllm/model_executor/models/kv_cache_route.h index 004f60d64..a5ebd0bb8 100644 --- a/include/vllm/model_executor/models/kv_cache_route.h +++ b/include/vllm/model_executor/models/kv_cache_route.h @@ -45,6 +45,15 @@ inline bool IsFp8KvCache(const PagedKvCache& kv) { return fp8_kind; } +// The dtype the ATTENTION presents for this store. An fp8 store is read through +// the bf16 dense dequant scratch (`Fp8PagedToDenseBf16Kernel`), so it presents +// bf16 exactly as a bf16 store does; a model that decides its query/out dtype +// from the store's dtype must ask THIS, or an fp8 store silently falls off the +// bf16 attention lanes (FA-2 included) onto the per-read CUDA-core path. +inline bool KvCachePresentsBf16(const PagedKvCache& kv) { + return kv.dtype == vt::DType::kBF16 || IsFp8KvCache(kv); +} + // The KV STORE. `k`/`v` are the model-dtype [T, Hkv, Dh] tensors the attention // preamble produced; `k_cache`/`v_cache` are this layer's `KvSlice` views. // diff --git a/src/vllm/model_executor/models/qwen3_5.cpp b/src/vllm/model_executor/models/qwen3_5.cpp index 0900a64a6..41646c5c0 100644 --- a/src/vllm/model_executor/models/qwen3_5.cpp +++ b/src/vllm/model_executor/models/qwen3_5.cpp @@ -6033,7 +6033,10 @@ DBuf FullAttnBlockPaged(Dev d, const FullAttnLayerWeights& w, const HfConfig& cf /*num_reqs=*/meta.num_reqs, /*uniform_spec_query_len=*/meta.uniform_spec_query_len, /*causal=*/meta.causal, - /*kv_cache_bf16=*/kv.dtype == DType::kBF16, + // An fp8 store is served through a bf16 dense scratch (the prefill + // dequant), so it presents bf16 too — that is what admits the FA-2 + // prefill lane. The decode lanes keep their own admission. + /*kv_cache_bf16=*/dense_attn::KvCachePresentsBf16(kv), /*kv_block_multiple_16=*/kv.block_size % 16 == 0, /*preamble_with_cos_sin=*/FuseAttnPreambleOn(fp4) && sdi.has_attn_cos_sin, /*fa2_platform=*/fa2_platform, @@ -6114,7 +6117,7 @@ DBuf FullAttnBlockPaged(Dev d, const FullAttnLayerWeights& w, const HfConfig& cf Tensor vw = v3; DBuf kbf(d, DType::kBF16, {T, Hkv, Dh}); DBuf vbf(d, DType::kBF16, {T, Hkv, Dh}); - if (kv.dtype == DType::kBF16 || dense_attn::IsFp8KvCache(kv)) { + if (dense_attn::KvCachePresentsBf16(kv)) { // K may already be bf16 (an FA2 preamble emits bf16 k directly — // the RN round of the same f32 value this CastBf16 would produce); only // down-cast when the preamble/fallback produced f32 K. diff --git a/src/vt/cuda/cuda_paged_attn.cu b/src/vt/cuda/cuda_paged_attn.cu index 752a54d33..5e712de31 100644 --- a/src/vt/cuda/cuda_paged_attn.cu +++ b/src/vt/cuda/cuda_paged_attn.cu @@ -26,10 +26,13 @@ // Prefill is never CUDA-graph-captured (see cuda_matmul_nvfp4.cu), so the // launcher may read query_start_loc D2H to build per-request query tiles. #include +#include #include #include #include +#include +#include #include #include #include @@ -106,6 +109,54 @@ void Check(cudaError_t err, const char* what) { cudaStream_t AsStream(const Queue& q) { return static_cast(q.handle); } +// ─── Test-only fault injection for the fp8 dense-prefill scratch ──────────── +// The OWNERSHIP of that scratch is what the release-on-unwind cases prove, and +// the failures they need to inject are CUDA failures a healthy device never +// produces on demand. Two independent bits, armed only through +// `vt::cuda::testing` (defined at the bottom of this file); production code +// never sets them. +constexpr int kFp8DenseFailIdentityAlloc = 1 << 0; +constexpr int kFp8DenseFailIdentityCopy = 1 << 1; + +std::atomic& Fp8DenseFailpoints() { + static std::atomic failpoints{0}; + return failpoints; +} + +bool Fp8DenseFailpointArmed(int failpoint) { + return (Fp8DenseFailpoints().load(std::memory_order_relaxed) & failpoint) != 0; +} + +// OWNER for a stream-ordered scratch allocation. The destructor frees on the +// SAME stream the consumers were enqueued on, so the free is ordered behind +// them, and it runs while an exception unwinds — the case the explicit frees at +// the end of a success path cannot reach. +// +// The destructor must not throw: `Check` here would replace the exception +// already in flight (or terminate, during unwinding), which is the failure this +// guard exists to prevent. A failed free is therefore dropped; the success path +// still checks its OWN frees explicitly through `release()`. +class StreamScratch { + public: + StreamScratch(cudaStream_t stream, void* ptr) : stream_(stream), ptr_(ptr) {} + ~StreamScratch() { + if (ptr_ != nullptr) (void)cudaFreeAsync(ptr_, stream_); + } + StreamScratch(const StreamScratch&) = delete; + StreamScratch& operator=(const StreamScratch&) = delete; + + // Hands ownership back so the caller can check the free explicitly. + void* release() { + void* p = ptr_; + ptr_ = nullptr; + return p; + } + + private: + cudaStream_t stream_; + void* ptr_; +}; + // Opt a kernel into more dynamic shared memory than the 48 KiB every CUDA // architecture guarantees without an opt-in. // @@ -2578,6 +2629,16 @@ bool PrefillFlash2VecEnabled() { return enabled; } +// fp8 KV prefill via the one-time dense dequant scratch. Default ON; +// VT_ATTN_FP8_DENSE=0 restores the per-read dequant for a same-binary A/B. +bool PrefillFp8DenseEnabled() { + static const bool enabled = [] { + const char* e = std::getenv("VT_ATTN_FP8_DENSE"); + return !(e != nullptr && e[0] == '0'); + }(); + return enabled; +} + // Default query-tile size when VT_ATTN_PREFILL_BM is unset. Microbench (GB10, // 8x1024..2x4096 prefill) picked BM=16 with the *restructured* BM kernel: it is // result-identical to the old Vec kernel but fills more of the 8 warps in QKᵀ at @@ -3093,18 +3154,63 @@ void LaunchPaged(cudaStream_t s, Tensor& out, const Tensor& query, const Tensor& // read is dequantized as Dequant(fp8) * k_scale|v_scale inside LoadKv, mirroring // upstream's `scaled_vec_conversion` (quant_utils.cuh:419-429). // -// SCOPE, argued rather than assumed. Only the two CORRECTNESS-GRADE kernels are -// reachable from here — the tiled flash prefill and the block decode — and that -// is not a shortcut, it is what the ladder above already implies. Every faster -// arm is bf16-NATIVE by construction: the WMMA prefill ladder stages -// `__nv_bfloat16` fragments, the vendored FA-2 launchers take bf16 pointers, and -// the vectorized decode-opt/GQA kernels read the cache through LoadRowN/LoadRow8, -// which are 128-bit `uint4` loads specialized for bf16 and f32 only. Upstream -// draws the same line from the other side: FlashAttention only serves a -// quantized KV cache when `flash_attn_supports_kv_cache_dtype` says so -// (flash_attn.py:181-187,796-805) and otherwise the backend refuses. A tensor- -// core fp8 read is a PERFORMANCE brick, not this one; W2's gate is parity with -// the W1 CPU reference, and W4 owns the memory/throughput measurement. +// SCOPE. PREFILL takes the dense dequant scratch below when the geometry fits, +// so its attention is served by the same lane a bf16 cache gets (FA-2 in the +// engine); the per-read kernels remain for decode and for shapes outside that +// gate, and decode is bandwidth-bound, so its in-register dequant is already the +// right shape. +// fp8 paged cache -> dense bf16 (per request), ONCE per layer: the attention +// kernel re-reads K/V once per query tile, so dequantizing in its staging loop +// pays the conversion thousands of times per element. One block per (request, +// key position), `blockIdx.y` selects the KV head, 16-byte vector reads. +__global__ void Fp8PagedToDenseBf16Kernel(const uint8_t* k_cache, const uint8_t* v_cache, + __nv_bfloat16* k_dense, __nv_bfloat16* v_dense, + const int32_t* block_table, const int32_t* seq_lens, + int max_seq, int hk, int d, int block_size, + int64_t bt_row, int64_t bt_col, int64_t kc_blk, + int64_t kc_pg, int64_t kc_hd, int64_t vc_blk, + int64_t vc_pg, int64_t vc_hd, float k_scale, + float v_scale) { + const int64_t t = blockIdx.x; + const int r = static_cast(t / max_seq); + const int j = static_cast(t % max_seq); + const int g = blockIdx.y; + if (j >= seq_lens[r]) return; + const int blk = + block_table[static_cast(r) * bt_row + (j / block_size) * bt_col]; + const int off = j % block_size; + const uint8_t* ksrc = k_cache + static_cast(blk) * kc_blk + + static_cast(off) * kc_pg + static_cast(g) * kc_hd; + const uint8_t* vsrc = v_cache + static_cast(blk) * vc_blk + + static_cast(off) * vc_pg + static_cast(g) * vc_hd; + __nv_bfloat16* kdst = + k_dense + (static_cast(r) * max_seq + j) * hk * d + static_cast(g) * d; + __nv_bfloat16* vdst = + v_dense + (static_cast(r) * max_seq + j) * hk * d + static_cast(g) * d; + const int n4 = d >> 4; + for (int i = threadIdx.x; i < n4; i += blockDim.x) { + const uint4 raw_k = reinterpret_cast(ksrc)[i]; + const uint4 raw_v = reinterpret_cast(vsrc)[i]; + const uint8_t* kb = reinterpret_cast(&raw_k); + const uint8_t* vb = reinterpret_cast(&raw_v); + alignas(16) __nv_bfloat16 tk[16], tv[16]; +#pragma unroll + for (int e = 0; e < 16; ++e) { + // e4m3 is exactly representable in bf16, so this is the whole decode. + tk[e] = __float2bfloat16(__half2float(__nv_cvt_fp8_to_halfraw(kb[e], __NV_E4M3)) * k_scale); + tv[e] = __float2bfloat16(__half2float(__nv_cvt_fp8_to_halfraw(vb[e], __NV_E4M3)) * v_scale); + } + const uint4* tk4 = reinterpret_cast(tk); + const uint4* tv4 = reinterpret_cast(tv); + uint4* kd4 = reinterpret_cast(kdst + (i << 4)); + uint4* vd4 = reinterpret_cast(vdst + (i << 4)); + kd4[0] = tk4[0]; + kd4[1] = tk4[1]; + vd4[0] = tv4[0]; + vd4[1] = tv4[1]; + } +} + template void LaunchPagedFp8Out(cudaStream_t s, Tensor& out, const Tensor& query, const Tensor& k_cache, const Tensor& v_cache, const Tensor& block_table, const Tensor& seq_lens, @@ -3116,6 +3222,102 @@ void LaunchPagedFp8Out(cudaStream_t s, Tensor& out, const Tensor& query, const T // Same predicate LaunchPaged uses to pick the tiled prefill kernel. const bool is_prefill = num_tokens > num_reqs; if (is_prefill && d <= kMaxEpl * 32 && PrefillFlashEnabled()) { + // Dequantize the cache ONCE per layer into a dense bf16 scratch and run the + // normal bf16 dispatch on it, so the attention is served by the same lane a + // bf16 cache gets. The scratch is one block per request with + // block_size = max_seq and an identity block table, so the kernel's paged + // address IS the dense address and no attention kernel changes. The dequant + // is exact (bf16 holds every e4m3 value), so tokens match the per-read path. + // VT_ATTN_FP8_DENSE=0 restores that path for a same-binary A/B. + if (PrefillFp8DenseEnabled() && d % 16 == 0) { + // `max_seq_len` is the runner's host-known max over seq_lens; only callers + // that leave it 0 (op-level tests) pay the readback. + int64_t max_seq = args.max_seq_len; + std::vector sl; + if (max_seq <= 0) { + sl.resize(static_cast(num_reqs)); + Check(cudaMemcpyAsync(sl.data(), seq_lens.Ptr(), + static_cast(num_reqs) * sizeof(int32_t), + cudaMemcpyDeviceToHost, s), + "fp8 prefill seq_lens D2H"); + Check(cudaStreamSynchronize(s), "fp8 prefill seq_lens sync"); + max_seq = *std::max_element(sl.begin(), sl.end()); + } + if (max_seq > 0) { + const size_t dense_bytes = static_cast(num_reqs) * + static_cast(max_seq) * + static_cast(num_kv_heads) * + static_cast(d) * sizeof(__nv_bfloat16); + uint8_t* dense = nullptr; + Check(cudaMallocAsync(&dense, dense_bytes * 2, s), "fp8 prefill dense scratch"); + // Every later step can throw — the identity allocation and copy, the + // dequant launch, the bf16 dispatch it feeds, and the vectors between + // them — so each allocation is owned by a scope guard the moment it + // exists. On the success path the frees are still explicit and checked. + StreamScratch dense_scratch(s, dense); + int32_t* d_ids = nullptr; + if (Fp8DenseFailpointArmed(kFp8DenseFailIdentityAlloc)) { + Check(cudaErrorMemoryAllocation, "fp8 prefill identity table"); + } else { + Check(cudaMallocAsync(&d_ids, static_cast(num_reqs) * sizeof(int32_t), s), + "fp8 prefill identity table"); + } + StreamScratch ids_scratch(s, d_ids); + std::vector ids(static_cast(num_reqs)); + for (int64_t r = 0; r < num_reqs; ++r) + ids[static_cast(r)] = static_cast(r); + if (Fp8DenseFailpointArmed(kFp8DenseFailIdentityCopy)) { + Check(cudaErrorMemoryAllocation, "fp8 prefill identity table H2D"); + } else { + Check(cudaMemcpyAsync(d_ids, ids.data(), static_cast(num_reqs) * sizeof(int32_t), + cudaMemcpyHostToDevice, s), + "fp8 prefill identity table H2D"); + } + const dim3 dense_grid(static_cast(num_reqs * max_seq), + static_cast(num_kv_heads)); + Fp8PagedToDenseBf16Kernel<<>>( + k_cache.Ptr(), v_cache.Ptr(), + reinterpret_cast<__nv_bfloat16*>(dense), + reinterpret_cast<__nv_bfloat16*>(dense + dense_bytes), + block_table.Ptr(), seq_lens.Ptr(), static_cast(max_seq), + static_cast(num_kv_heads), static_cast(d), static_cast(block_size), + block_table.stride[0], block_table.stride[1], k_cache.stride[0], k_cache.stride[1], + k_cache.stride[2], v_cache.stride[0], v_cache.stride[1], v_cache.stride[2], + args.k_scale, args.v_scale); + Check(cudaGetLastError(), "fp8 prefill dense dequant launch"); + Tensor kd{}; + kd.data = dense; + kd.dtype = DType::kBF16; + kd.device = k_cache.device; + kd.rank = 4; + kd.shape[0] = num_reqs; + kd.shape[1] = max_seq; + kd.shape[2] = num_kv_heads; + kd.shape[3] = d; + kd.stride[0] = max_seq * num_kv_heads * d; + kd.stride[1] = num_kv_heads * d; + kd.stride[2] = d; + kd.stride[3] = 1; + Tensor vd = kd; + vd.data = dense + dense_bytes; + Tensor idt{}; + idt.data = d_ids; + idt.dtype = DType::kI32; + idt.device = k_cache.device; + idt.rank = 2; + idt.shape[0] = num_reqs; + idt.shape[1] = 1; + idt.stride[0] = 1; + idt.stride[1] = 1; + // The standard bf16 dispatch, so the bf16 ladder (FA-2 in the engine) + // serves the fp8 cache through the scratch. + LaunchPaged(s, out, query, kd, vd, idt, seq_lens, + query_start_loc, args); + Check(cudaFreeAsync(ids_scratch.release(), s), "fp8 prefill identity table free"); + Check(cudaFreeAsync(dense_scratch.release(), s), "fp8 prefill dense scratch free"); + return; + } + } LaunchPrefillFlash(s, out, query, k_cache, v_cache, block_table, seq_lens, query_start_loc, args, hq, d, num_reqs, num_kv_heads, block_size); @@ -3198,4 +3400,44 @@ struct Registrar { } registrar; } // namespace + +// ─── Test seams for the scratch-ownership cases ───────────────────────────── +// Declared by `tests/vt/test_ops_paged_attn.cpp` (the same arrangement as the +// FA2 decode seams). A CPU-only build never compiles this file, and the test +// supplies inline no-op stubs so its cases still compile there. +namespace testing { + +void SetFp8DenseIdentityAllocFailureForTesting(bool fail) { + std::atomic& failpoints = Fp8DenseFailpoints(); + if (fail) failpoints.fetch_or(kFp8DenseFailIdentityAlloc, std::memory_order_relaxed); + else failpoints.fetch_and(~kFp8DenseFailIdentityAlloc, std::memory_order_relaxed); +} + +void SetFp8DenseIdentityCopyFailureForTesting(bool fail) { + std::atomic& failpoints = Fp8DenseFailpoints(); + if (fail) failpoints.fetch_or(kFp8DenseFailIdentityCopy, std::memory_order_relaxed); + else failpoints.fetch_and(~kFp8DenseFailIdentityCopy, std::memory_order_relaxed); +} + +void ClearFp8DenseFailureForTesting() { + Fp8DenseFailpoints().store(0, std::memory_order_relaxed); +} + +// Bytes currently in use in the CUDA memory pool the scratch is allocated from. +// The release cases bracket one call with this: a buffer the guard did not free +// keeps its bytes counted after the stream drains, so the observed value stays +// above the pre-call one. A failed query throws rather than reporting 0, which +// would make the comparison pass vacuously. +size_t Fp8DensePoolUsedBytesForTesting() { + int device = 0; + Check(cudaGetDevice(&device), "pool used: cudaGetDevice"); + cudaMemPool_t pool = nullptr; + Check(cudaDeviceGetMemPool(&pool, device), "pool used: cudaDeviceGetMemPool"); + size_t used = 0; + Check(cudaMemPoolGetAttribute(pool, cudaMemPoolAttrUsedMemCurrent, &used), + "pool used: cudaMemPoolGetAttribute"); + return used; +} + +} // namespace testing } // namespace vt::cuda diff --git a/tests/vt/test_ops_paged_attn.cpp b/tests/vt/test_ops_paged_attn.cpp index 46473e94a..93fa5dab2 100644 --- a/tests/vt/test_ops_paged_attn.cpp +++ b/tests/vt/test_ops_paged_attn.cpp @@ -34,6 +34,7 @@ #include "vt/backend.h" #include "vt/dtype.h" +#include "vt/fp8_kv.h" #include "vt/ops.h" #ifdef VLLM_CPP_FLASH_ATTN @@ -52,6 +53,25 @@ size_t Fa2DecodeScratchShapeCountForTesting(int device, void* stream); } // namespace vt::cuda::testing #endif +// The scratch-ownership seams behind the fp8 dense-prefill block of +// `LaunchPagedFp8Out` (cuda_paged_attn.cu). The CUDA build defines them; a +// CPU-only build cannot link them, so the inline no-op stubs keep the +// release-on-unwind case COMPILED there (it skips at runtime through HasCuda), +// the same reason the FA2 seams above are declared in this TU. +namespace vt::cuda::testing { +#ifdef VLLM_CPP_CUDA +void SetFp8DenseIdentityAllocFailureForTesting(bool fail); +void SetFp8DenseIdentityCopyFailureForTesting(bool fail); +void ClearFp8DenseFailureForTesting(); +size_t Fp8DensePoolUsedBytesForTesting(); +#else +inline void SetFp8DenseIdentityAllocFailureForTesting(bool /*fail*/) {} +inline void SetFp8DenseIdentityCopyFailureForTesting(bool /*fail*/) {} +inline void ClearFp8DenseFailureForTesting() {} +inline size_t Fp8DensePoolUsedBytesForTesting() { return 0; } +#endif +} // namespace vt::cuda::testing + using vt::AttentionArgs; using vt::AttentionWindow; using vt::Backend; @@ -1066,6 +1086,184 @@ TEST_CASE("paged_attention CUDA WMMA (bf16 cache) matches f32 ref at head_dim 25 CHECK(local_max_abs < 5e-2); } +// =========================================================================== +// fp8-KV PREFILL through the one-time dense dequant: the op dequantizes the +// 1-byte cache into a dense bf16 scratch and runs the standard bf16 dispatch on +// it. This call presents an f32 query, so that dispatch takes its CUDA-core +// flash arm (the engine presents bf16, which admits FA-2). The reference runs on +// the dequantized values, so the error is the flash kernel's own rounding. +// VT_ATTN_FP8_DENSE=0 restores the per-read path. +// =========================================================================== +TEST_CASE("paged_attention CUDA fp8-KV prefill (one-time dense dequant) matches f32 ref at head_dim 256") { + if (!HasCuda()) { + MESSAGE("no CUDA backend; skipping paged_attention fp8-KV WMMA parity (dgx-pending)"); + return; + } + const int64_t Hq = 16, Hk = 2, D = 256, block_size = 16; + const float scale = std::pow(static_cast(D), -0.5f); + const float k_scale = 1.0f, v_scale = 1.0f; + std::vector qsl = {0, 100, 101, 104}; + std::vector seq_lens = {100, 133, 140}; + const int64_t num_tokens = 104; + const int64_t num_reqs = 3; + const int64_t num_blocks = 64, page = Hk * D, max_blocks = 9; + auto q = RandF32(static_cast(num_tokens * Hq * D), 2024); + auto kc = RandF32(static_cast(num_blocks * block_size * page), 137); + auto vc = RandF32(static_cast(num_blocks * block_size * page), 179); + std::vector block_table = {5, 0, 11, 3, 0, 0, 0, 0, 0, + 2, 17, 9, 20, 1, 8, 0, 0, 0, + 30, 4, 22, 15, 6, 19, 7, 12, 0}; + + // fp8-e4m3 cache (per-tensor scale 1.0, the uncalibrated default) and the f32 + // reference on the exact dequantized values. + std::vector kc_f(kc.size()), vc_f(vc.size()); + std::vector kc_r(kc.size()), vc_r(vc.size()); + for (size_t i = 0; i < kc.size(); ++i) { + kc_f[i] = vt::StoreKvFp8E4M3(kc[i], k_scale); + kc_r[i] = vt::LoadKvFp8E4M3(kc_f[i], k_scale); + vc_f[i] = vt::StoreKvFp8E4M3(vc[i], v_scale); + vc_r[i] = vt::LoadKvFp8E4M3(vc_f[i], v_scale); + } + std::vector ref = ComposedPagedRef(q, kc_r, vc_r, block_table, max_blocks, seq_lens, qsl, + Hq, Hk, D, block_size, scale, true); + + const int64_t within = block_size * Hk * D; + std::vector combined(static_cast(num_blocks * 2 * within), 0); + for (int64_t b = 0; b < num_blocks; ++b) + for (int64_t e = 0; e < within; ++e) { + combined[static_cast((b * 2 + 0) * within + e)] = + kc_f[static_cast(b * within + e)]; + combined[static_cast((b * 2 + 1) * within + e)] = + vc_f[static_cast(b * within + e)]; + } + + Backend& gpu = vt::GetBackend(DeviceType::kCUDA); + QueueGuard g(gpu); + DeviceTensor dq(gpu, g.q, DType::kF32, {num_tokens, Hq, D}, q.data()); + DeviceTensor dcache(gpu, g.q, DType::kI8, {num_blocks * 2 * within}, combined.data()); + auto SliceView = [&](int which) { + Tensor t = dcache.tensor(); + t.data = static_cast(t.data) + + static_cast(which) * static_cast(within) * vt::SizeOf(DType::kI8); + t.rank = 4; + t.shape[0] = num_blocks; + t.shape[1] = block_size; + t.shape[2] = Hk; + t.shape[3] = D; + t.stride[0] = 2 * within; + t.stride[1] = Hk * D; + t.stride[2] = D; + t.stride[3] = 1; + return t; + }; + Tensor kview = SliceView(0); + Tensor vview = SliceView(1); + DeviceTensor dbt(gpu, g.q, DType::kI32, {num_reqs, max_blocks}, block_table.data()); + DeviceTensor dsl(gpu, g.q, DType::kI32, {num_reqs}, seq_lens.data()); + DeviceTensor dqsl(gpu, g.q, DType::kI32, {num_reqs + 1}, qsl.data()); + DeviceTensor dout(gpu, g.q, DType::kF32, {num_tokens, Hq, D}); + PagedAttentionArgs args{scale, true}; + args.kv_cache_dtype = vt::Fp8KVCacheDataType::kFp8E4M3; + args.k_scale = k_scale; + args.v_scale = v_scale; + vt::PagedAttention(g.q, dout.tensor(), dq.tensor(), kview, vview, dbt.tensor(), dsl.tensor(), + dqsl.tensor(), args); + std::vector got(static_cast(num_tokens * Hq * D), 0.0f); + dout.Download(g.q, got.data()); + + double max_abs = 0.0; + for (size_t i = 0; i < ref.size(); ++i) + max_abs = std::max(max_abs, std::abs(static_cast(got[i]) - ref[i])); + MESSAGE("fp8-KV dense prefill max_abs_err vs f32 ref = " << max_abs); + CHECK(max_abs < 5e-2); +} + +// =========================================================================== +// fp8-KV PREFILL scratch ownership: the dense scratch and the identity table +// live only for the duration of one dispatch, and the steps that can throw +// between the two cudaMallocAsync calls (the identity copy, the dequant launch, +// the bf16 dispatch it feeds, the vectors between them) must not strand either +// buffer. The two injections below are failures a healthy device never produces +// on demand; each pass proves the memory pool gives the bytes back after the +// exception unwinds, and that the propagated exception is the Check that +// failed rather than something a destructor raised in its place. +// =========================================================================== +TEST_CASE("paged_attention CUDA fp8-KV prefill returns both scratch buffers to the pool when a later step throws") { + if (!HasCuda()) { + MESSAGE("no CUDA backend; skipping fp8-KV scratch release (dgx-pending)"); + return; + } + const int64_t Hq = 2, Hk = 1, D = 256, block_size = 4; + const float scale = std::pow(static_cast(D), -0.5f); + const std::vector qsl = {0, 4, 10}; + const std::vector seq_lens = {4, 6}; + const int64_t num_tokens = 10, num_reqs = 2, num_blocks = 4, max_blocks = 2; + auto q = RandF32(static_cast(num_tokens * Hq * D), 7); + auto kv = RandF32(static_cast(num_blocks * block_size * Hk * D), 11); + std::vector cache(kv.size()); + for (size_t i = 0; i < kv.size(); ++i) cache[i] = vt::StoreKvFp8E4M3(kv[i], 1.0f); + std::vector block_table = {0, 1, 2, 3}; + + const int64_t within = block_size * Hk * D; + std::vector combined(static_cast(num_blocks * 2 * within), 0); + for (int64_t b = 0; b < num_blocks; ++b) + for (int64_t e = 0; e < within; ++e) { + combined[static_cast((b * 2 + 0) * within + e)] = + cache[static_cast(b * within + e)]; + combined[static_cast((b * 2 + 1) * within + e)] = + cache[static_cast(b * within + e)]; + } + + Backend& gpu = vt::GetBackend(DeviceType::kCUDA); + QueueGuard g(gpu); + DeviceTensor dq(gpu, g.q, DType::kF32, {num_tokens, Hq, D}, q.data()); + DeviceTensor dcache(gpu, g.q, DType::kI8, {num_blocks * 2 * within}, combined.data()); + auto SliceView = [&](int which) { + Tensor t = dcache.tensor(); + t.data = static_cast(t.data) + + static_cast(which) * static_cast(within) * vt::SizeOf(DType::kI8); + t.rank = 4; + t.shape[0] = num_blocks; + t.shape[1] = block_size; + t.shape[2] = Hk; + t.shape[3] = D; + t.stride[0] = 2 * within; + t.stride[1] = Hk * D; + t.stride[2] = D; + t.stride[3] = 1; + return t; + }; + Tensor kview = SliceView(0); + Tensor vview = SliceView(1); + DeviceTensor dbt(gpu, g.q, DType::kI32, {num_reqs, max_blocks}, block_table.data()); + DeviceTensor dsl(gpu, g.q, DType::kI32, {num_reqs}, seq_lens.data()); + DeviceTensor dqsl(gpu, g.q, DType::kI32, {num_reqs + 1}, qsl.data()); + DeviceTensor dout(gpu, g.q, DType::kF32, {num_tokens, Hq, D}); + PagedAttentionArgs args{scale, true}; + args.kv_cache_dtype = vt::Fp8KVCacheDataType::kFp8E4M3; + args.k_scale = 1.0f; + args.v_scale = 1.0f; + + auto pass = [&](const char* what, const char* expected) { + CAPTURE(what); + const size_t used_before = vt::cuda::testing::Fp8DensePoolUsedBytesForTesting(); + CHECK_THROWS_WITH_AS( + vt::PagedAttention(g.q, dout.tensor(), dq.tensor(), kview, vview, dbt.tensor(), + dsl.tensor(), dqsl.tensor(), args), + doctest::Contains(expected), std::runtime_error); + vt::cuda::testing::ClearFp8DenseFailureForTesting(); + gpu.Synchronize(g.q); + CHECK(vt::cuda::testing::Fp8DensePoolUsedBytesForTesting() == used_before); + }; + + // (1) The identity-table allocation throws after the dense scratch exists. + vt::cuda::testing::SetFp8DenseIdentityAllocFailureForTesting(true); + pass("identity allocation", "fp8 prefill identity table:"); + // (2) The identity-table copy throws after BOTH allocations succeeded. + vt::cuda::testing::SetFp8DenseIdentityCopyFailureForTesting(true); + pass("identity copy", "fp8 prefill identity table H2D:"); +} + // =========================================================================== // CUDA parity at the GATE-MODEL config: head_dim 256, GQA 16q/2kv. This is the // M2.4 flash prefill path's real shape — head_dim 256 means the warp-per-row