diff --git a/Cargo.lock b/Cargo.lock index ec7492f70..6555182ce 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3773,6 +3773,8 @@ dependencies = [ "log", "pegainfer-build", "serde", + "serde_json", + "sha2 0.11.0", "tvm-ffi", ] diff --git a/README.md b/README.md index f30b2b665..cf3be38de 100644 --- a/README.md +++ b/README.md @@ -71,7 +71,7 @@ pegainfer --model-path models/Qwen3-4B - Rust (2024 edition), CUDA Toolkit (nvcc, cuBLAS), CUDA-capable GPU - NVIDIA driver R545 (CUDA 12.3) or newer; `cuFuncGetName` sets this floor, while per-symbol lazy loading keeps the `cuda-12090` cudarc binding from requiring a CUDA 12.9 driver - The default build (Qwen3-4B / 8B) is pure Rust + CUDA — no Python at all -- Python 3 + Triton for `qwen35` feature builds (build-time only — no Python at runtime) +- Python 3 + Triton for `qwen35` feature builds (build-time only — no Python at runtime). The supported Qwen3.5-4B SM120/Hv32 path can additionally link a pre-generated FlashInfer/CuTe GDN AOT candidate; generation is separate from serving - The `kimi-k2` EP path additionally needs NCCL ≥ 2.27 at runtime (`ncclAlltoAll`) ### Build & Run @@ -105,7 +105,8 @@ curl -N http://localhost:8000/v1/completions \ More options ```bash -# Qwen3.5 requires the feature-gated Triton AOT kernels (Python + Triton at build time) +# Qwen3.5 uses build-time Triton AOT; a validated SM120/Hv32 FlashInfer GDN +# candidate can additionally be selected through PEGAINFER_QWEN35_GDN_AOT_BUNDLE uv venv && uv pip install triton export PEGAINFER_TRITON_PYTHON=.venv/bin/python cargo run --release --features qwen35 -- --model-path models/Qwen3.5-4B @@ -120,6 +121,7 @@ cargo run --release -- --cuda-graph=false |----------|-------------| | `CUDA_HOME` | CUDA Toolkit path (default: `/usr/local/cuda`) | | `PEGAINFER_TRITON_PYTHON` | Python with Triton for `qwen35` build-time AOT compilation | +| `PEGAINFER_QWEN35_GDN_AOT_BUNDLE` | Validated Qwen3.5-4B SM120/Hv32 FlashInfer GDN candidate directory linked by `pegainfer-kernels/build.rs` | | `PEGAINFER_TILELANG_PYTHON` | Python with TileLang for `k3` build-time kernel generation | | `PEGAINFER_CUDA_SM` | GPU SM target override when `nvidia-smi` unavailable (e.g. `120`) | @@ -303,14 +305,14 @@ flowchart TB **Key design decisions:** - **GPU-first runtime** — model execution stays in native Rust/CUDA paths -- **Custom GPU kernels** — CUDA for decode-critical paths, Triton AOT for Qwen3.5 compatibility kernels, FlashInfer for paged attention/sampling, NCCL for multi-GPU reductions, and cuBLAS for matrix multiplication +- **Custom GPU kernels** — CUDA for decode-critical paths; Triton AOT for the general Qwen3.5 GDN path; a statically linked FlashInfer/CuTe AOT specialization for Qwen3.5-4B SM120/Hv32 when a validated candidate is supplied; FlashInfer for paged attention/sampling; NCCL for multi-GPU reductions; and cuBLAS for matrix multiplication - **CUDA Graph** on Qwen decode paths — eliminates kernel launch overhead where enabled - **Per-model crate boundary** — Qwen3-4B owns its config, weights, scheduler/executor, tests, benches, and kernel plan in `pegainfer-qwen3` **Model details:** - **Qwen3**: 32 Q heads, 8 KV heads (GQA 4:1), head_dim=128 -- **Qwen3.5**: hybrid — 24 linear attention layers (Gated Delta Rule) + 8 full attention layers, head_dim=256 +- **Qwen3.5**: hybrid — 24 linear attention layers (Gated Delta Rule) + 8 full attention layers, head_dim=256. Qwen3.5-4B on single-GPU SM120 uses the validated FlashInfer GDN AOT specialization when linked; unsupported geometry/SM/TP configurations retain the explicit Triton path - **DeepSeek V2-Lite**: feature-gated 2-GPU EP2 correctness/attribution path for the HF/host-staged/NCCL narrow greedy gate ### What's not (yet) implemented @@ -350,6 +352,8 @@ PEGAINFER_TEST_MODEL_PATH=models/Qwen3.5-4B cargo test --release -p pegainfer-qw PEGAINFER_TEST_MODEL_PATH=models/DeepSeek-V2-Lite cargo test --release -p pegainfer-deepseek-v2-lite --features deepseek-v2-lite --test e2e_ep2 -- --nocapture ``` +The SM120 FlashInfer GDN production boundary has a separate fail-closed five-gate runner because it requires a real generated candidate, the pinned Qwen3.5-4B snapshot, and an SM120 GPU. See [`pegainfer-kernels/tools/flashinfer_gdn/README.md`](pegainfer-kernels/tools/flashinfer_gdn/README.md); the non-default `gdn-validation` feature exists only for that runner and adds no counters or validation API to a default serving build. + The DeepSeek-V2-Lite E2E is a correctness/integration gate. Direct diagnostics and HTTP SLO report commands live in [`benchmarking.md`](docs/models/deepseek-v2-lite/benchmarking.md). ## License diff --git a/pegainfer-kernels/Cargo.toml b/pegainfer-kernels/Cargo.toml index 4849b68ae..bbf3ab350 100644 --- a/pegainfer-kernels/Cargo.toml +++ b/pegainfer-kernels/Cargo.toml @@ -15,6 +15,8 @@ tvm-ffi = { version = "0.1.0-alpha.0", optional = true } [build-dependencies] cc = { workspace = true } pegainfer-build = { workspace = true } +serde_json = { workspace = true, optional = true } +sha2 = { workspace = true, optional = true } [features] default = [] @@ -24,7 +26,7 @@ tvm-ffi-triton-cubin = ["dep:tvm-ffi", "qwen35"] deepseek-v2-lite = [] # Gemma 4: NVFP4 dequantization, whose conversion intrinsics need CUDA >= 12.8. gemma4 = [] -qwen35 = [] +qwen35 = ["dep:serde_json", "dep:sha2"] # Shared MoE/MLA third-party substrate: DeepEP, DeepGEMM, and FlashMLA. glm52 = ["moe"] # Kimi K3: TileLang-generated kernels, FP8xFP4 masked grouped GEMM, and the diff --git a/pegainfer-kernels/KERNELS.md b/pegainfer-kernels/KERNELS.md index 57b207b2d..00fc00020 100644 --- a/pegainfer-kernels/KERNELS.md +++ b/pegainfer-kernels/KERNELS.md @@ -231,7 +231,7 @@ The crate still builds CUDA/Triton symbols needed by the current root binary: - Qwen3.5 HD256 full-attention kernels: `csrc/qwen35/prefill_attention_hd256.cu`, `csrc/shared/paged_attention.cu`. - Qwen3.5 linear-attention decode kernels: `csrc/qwen35/conv1d.cu`, `csrc/qwen35/gated_delta_rule.cu`. -- Qwen3.5 chunk-wise GDR prefill Triton AOT kernels: `tools/triton/gated_delta_rule_chunkwise_kernels.py`. +- Qwen3.5 chunk-wise GDR prefill uses the Triton AOT kernels in `tools/triton/gated_delta_rule_chunkwise_kernels.py` generally. A validated build-linked candidate selects the FlashInfer/CuTe AOT specialization in `csrc/qwen35/flashinfer_gdn_aot.c` for single-GPU SM120 Qwen3.5-4B (`Hq/Hk/Hv/D=16/16/32/128`); unsupported SM, geometry, and TP configurations retain Triton. These are preserved for build compatibility. They are not part of the Qwen3-4B Phase 1 API surface. diff --git a/pegainfer-kernels/build.rs b/pegainfer-kernels/build.rs index 949be5c06..d83035e69 100644 --- a/pegainfer-kernels/build.rs +++ b/pegainfer-kernels/build.rs @@ -9,6 +9,9 @@ use std::sync::Mutex; use std::thread; use std::time::Instant; +#[cfg(feature = "qwen35")] +use sha2::Digest as _; + struct TritonKernelSpec { artifact_dir: &'static str, kernel_path: &'static str, @@ -39,6 +42,177 @@ struct FlashInferIncludes { cccl: Vec, } +#[cfg(feature = "qwen35")] +const QWEN35_GDN_AOT_ABI_VERSION: u64 = 1; +#[cfg(feature = "qwen35")] +const QWEN35_GDN_AOT_ENV: &str = "PEGAINFER_QWEN35_GDN_AOT_BUNDLE"; + +#[cfg(feature = "qwen35")] +fn sha256_file(path: &Path) -> String { + let bytes = + fs::read(path).unwrap_or_else(|error| panic!("failed to read {}: {error}", path.display())); + let digest = sha2::Sha256::digest(bytes); + let mut hex = String::with_capacity(digest.len() * 2); + for byte in digest { + write!(&mut hex, "{byte:02x}").expect("write SHA-256 hex to String"); + } + hex +} + +#[cfg(feature = "qwen35")] +fn json_u64(value: &serde_json::Value, path: &[&str]) -> u64 { + let mut cursor = value; + for key in path { + cursor = &cursor[*key]; + } + cursor.as_u64().unwrap_or_else(|| { + panic!( + "GDN AOT manifest field {} must be an unsigned integer", + path.join(".") + ) + }) +} + +#[cfg(feature = "qwen35")] +fn json_str<'a>(value: &'a serde_json::Value, path: &[&str]) -> &'a str { + let mut cursor = value; + for key in path { + cursor = &cursor[*key]; + } + cursor + .as_str() + .unwrap_or_else(|| panic!("GDN AOT manifest field {} must be a string", path.join("."))) +} + +/// Validate and attach the release-provided Qwen3.5 GDN object. The generated +/// object and its native CuTe runtime archive are linked statically; serving +/// never reads a manifest, loads PTX, or discovers a Python wheel. +#[cfg(feature = "qwen35")] +fn build_qwen35_flashinfer_gdn_aot( + root: &Path, + out_dir: &Path, + cuda_include: &Path, +) -> (Vec, Option) { + println!("cargo:rerun-if-env-changed={QWEN35_GDN_AOT_ENV}"); + let shim = root.join("csrc/qwen35/flashinfer_gdn_aot.c"); + let shim_header = root.join("csrc/qwen35/flashinfer_gdn_aot.h"); + println!("cargo:rerun-if-changed={}", shim.display()); + println!("cargo:rerun-if-changed={}", shim_header.display()); + + let config_header = out_dir.join("flashinfer_gdn_build_config.h"); + let mut includes = vec![root.join("csrc/qwen35"), out_dir.to_path_buf()]; + let mut linked_objects = Vec::new(); + let mut runtime_dir = None; + let mut config = String::from( + "#pragma once\n#define PEGAINFER_QWEN35_GDN_ARTIFACT_SHA256 \"unavailable\"\n#define PEGAINFER_QWEN35_GDN_WORKSPACE_BYTES_PER_SM 128u\n", + ); + + if let Some(bundle) = std::env::var_os(QWEN35_GDN_AOT_ENV) { + let bundle = PathBuf::from(bundle); + let manifest_path = bundle.join("manifest.json"); + let manifest_bytes = fs::read(&manifest_path).unwrap_or_else(|error| { + panic!("read GDN AOT manifest {}: {error}", manifest_path.display()) + }); + let manifest: serde_json::Value = serde_json::from_slice(&manifest_bytes) + .unwrap_or_else(|error| panic!("parse GDN AOT manifest: {error}")); + assert_eq!(json_u64(&manifest, &["schema_version"]), 3); + assert_eq!(json_str(&manifest, &["variant"]), "qwen35_4b_candidate"); + assert_eq!(json_str(&manifest, &["target", "arch"]), "sm_120a"); + assert_eq!( + json_str(&manifest, &["target", "code_object"]), + "embedded_cubin" + ); + assert_eq!( + json_u64(&manifest, &["abi", "version"]), + QWEN35_GDN_AOT_ABI_VERSION + ); + assert_eq!(json_u64(&manifest, &["geometry", "h_q"]), 16); + assert_eq!(json_u64(&manifest, &["geometry", "h_k"]), 16); + assert_eq!(json_u64(&manifest, &["geometry", "h_v"]), 32); + assert_eq!(json_u64(&manifest, &["geometry", "head_dim"]), 128); + assert_eq!(json_str(&manifest, &["tokens", "extent"]), "dynamic"); + assert_eq!(json_u64(&manifest, &["tokens", "minimum"]), 1); + assert_eq!( + json_str(&manifest, &["abi", "state_layout"]), + "openinfer_hkv_v_contiguous" + ); + assert_eq!(json_str(&manifest, &["workspace", "kind"]), "per_sm"); + let workspace_bytes_per_sm = json_u64(&manifest, &["workspace", "bytes_per_sm"]); + assert_eq!(workspace_bytes_per_sm, 128); + assert_eq!(json_u64(&manifest, &["workspace", "alignment_bytes"]), 128); + + let header = bundle.join("kernel.h"); + let object = bundle.join("kernel.o"); + let runtime = bundle.join("libcuda_dialect_runtime_static.a"); + for (label, path, hash_path, size_path) in [ + ( + "header", + &header, + ["artifact", "header", "sha256"], + ["artifact", "header", "size_bytes"], + ), + ( + "object", + &object, + ["artifact", "object", "sha256"], + ["artifact", "object", "size_bytes"], + ), + ( + "native runtime", + &runtime, + ["artifact", "native_runtime", "sha256"], + ["artifact", "native_runtime", "size_bytes"], + ), + ] { + assert!( + path.is_file(), + "GDN AOT {label} is missing: {}", + path.display() + ); + assert_eq!(sha256_file(path), json_str(&manifest, &hash_path)); + assert_eq!( + fs::metadata(path).expect("read GDN AOT metadata").len(), + json_u64(&manifest, &size_path) + ); + println!("cargo:rerun-if-changed={}", path.display()); + } + println!("cargo:rerun-if-changed={}", manifest_path.display()); + + let object_hash = json_str(&manifest, &["artifact", "object", "sha256"]); + config = format!( + "#pragma once\n#define PEGAINFER_QWEN35_GDN_ARTIFACT_SHA256 \"{object_hash}\"\n#define PEGAINFER_QWEN35_GDN_WORKSPACE_BYTES_PER_SM {workspace_bytes_per_sm}u\n" + ); + includes.push(bundle); + linked_objects.push(object); + runtime_dir = runtime.parent().map(Path::to_path_buf); + } + fs::write(&config_header, config).expect("write GDN AOT build config"); + + let shim_obj = out_dir.join("qwen35_flashinfer_gdn_aot.o"); + let compiler = cc::Build::new().get_compiler(); + let mut command = compiler.to_command(); + command + .arg("-c") + .arg(&shim) + .arg("-o") + .arg(&shim_obj) + .arg("-O3") + .arg("-std=c11") + .arg("-fPIC") + .arg("-isystem") + .arg(cuda_include); + for include in includes { + command.arg("-I").arg(include); + } + if runtime_dir.is_some() { + command.arg("-DPEGAINFER_QWEN35_GDN_AOT"); + } + let status = command.status().expect("compile Qwen3.5 GDN AOT shim"); + assert!(status.success(), "Qwen3.5 GDN AOT shim compilation failed"); + linked_objects.push(shim_obj); + (linked_objects, runtime_dir) +} + const GLM52_TRTLLM_FMHA_CUBINS: &[(&str, &str)] = &[ ( "kGlm52FmhaSparseSeedQ8", @@ -1904,6 +2078,14 @@ fn main() { // --- k3: DeepGEMM-only, no DeepEP/NCCL dependency --- let k3_enabled = cfg!(feature = "k3"); let qwen35_enabled = cfg!(feature = "qwen35"); + #[cfg(feature = "qwen35")] + let (qwen35_gdn_objects, qwen35_gdn_runtime_dir) = if qwen35_enabled { + build_qwen35_flashinfer_gdn_aot(&crate_root(), &out_dir, &cuda_include) + } else { + (Vec::new(), None) + }; + #[cfg(not(feature = "qwen35"))] + let (qwen35_gdn_objects, qwen35_gdn_runtime_dir) = (Vec::new(), None::); if glm52_enabled { generate_glm52_trtllm_fmha_cubins(&crate_root(), &out_dir); build_glm52_cutedsl_fp8_dsl(&crate_root(), &out_dir, &cuda_include); @@ -1970,6 +2152,9 @@ fn main() { if !kimi_k2_enabled && is_kimi_k2_source(&csrc_dir, path) { return None; } + if !qwen35_enabled && file_name == "gdn_prepare.cu" { + return None; + } // --- k3 --- if !k3_enabled && is_k3_source(&csrc_dir, path) { return None; @@ -2486,6 +2671,7 @@ fn main() { ar_args.extend( obj_files .into_iter() + .chain(qwen35_gdn_objects) .map(|path| path.to_string_lossy().to_string()), ); @@ -2525,6 +2711,10 @@ fn main() { toolkit.link_search(); } println!("cargo:rustc-link-lib=static=kernels_cuda"); + if let Some(runtime_dir) = qwen35_gdn_runtime_dir { + println!("cargo:rustc-link-search=native={}", runtime_dir.display()); + println!("cargo:rustc-link-lib=static=cuda_dialect_runtime_static"); + } println!("cargo:rustc-link-lib=cudart"); println!("cargo:rustc-link-lib=cublas"); println!("cargo:rustc-link-lib=cublasLt"); diff --git a/pegainfer-kernels/csrc/qwen35/flashinfer_gdn_aot.c b/pegainfer-kernels/csrc/qwen35/flashinfer_gdn_aot.c new file mode 100644 index 000000000..9018f9b55 --- /dev/null +++ b/pegainfer-kernels/csrc/qwen35/flashinfer_gdn_aot.c @@ -0,0 +1,184 @@ +#include "flashinfer_gdn_aot.h" + +#include +#include +#include + +#include "flashinfer_gdn_build_config.h" + +#ifdef PEGAINFER_QWEN35_GDN_AOT +static int32_t status_from_cuda(cudaError_t error) { + if (error == cudaSuccess) return PEGAINFER_QWEN35_GDN_OK; + if (error == cudaErrorNotSupported) + return PEGAINFER_QWEN35_GDN_NOT_SUPPORTED; + if (error == cudaErrorInvalidValue) + return PEGAINFER_QWEN35_GDN_INVALID_ARGUMENT; + return PEGAINFER_QWEN35_GDN_CUDA_ERROR; +} +#include "kernel.h" + +typedef struct { + pegainfer_qwen35_gdn_qwen35_4b_candidate_Kernel_Module_t module; + int32_t device; + size_t workspace_bytes; +} gdn_handle_t; + +static int32_t load_current_device( + pegainfer_qwen35_gdn_qwen35_4b_candidate_Kernel_Module_t *module, + int32_t device) { + cudaError_t ret = cudaSuccess; + cudaLibrary_t *library = &module->module; + struct { + cudaLibrary_t **library; + cudaError_t *ret; + } init_args = {&library, &ret}; + _mlir_pegainfer_qwen35_gdn_qwen35_4b_candidate_cuda_init( + (void **)&init_args); + if (ret != cudaSuccess) return (int32_t)ret; + struct { + cudaLibrary_t **library; + int32_t *device; + cudaError_t *ret; + } load_args = {&library, &device, &ret}; + _mlir_pegainfer_qwen35_gdn_qwen35_4b_candidate_cuda_load_to_device( + (void **)&load_args); + return (int32_t)ret; +} +#endif + +uint32_t pegainfer_qwen35_gdn_abi_version(void) { + return PEGAINFER_QWEN35_GDN_ABI_VERSION; +} + +const char *pegainfer_qwen35_gdn_artifact_sha256(void) { + return PEGAINFER_QWEN35_GDN_ARTIFACT_SHA256; +} + +int32_t pegainfer_qwen35_gdn_aot_available(void) { +#ifdef PEGAINFER_QWEN35_GDN_AOT + return 1; +#else + return 0; +#endif +} + +int32_t pegainfer_qwen35_gdn_create(void **handle, int32_t device) { + if (handle == NULL) return PEGAINFER_QWEN35_GDN_INVALID_ARGUMENT; + *handle = NULL; +#ifdef PEGAINFER_QWEN35_GDN_AOT + cudaError_t ret = cudaSetDevice(device); + if (ret != cudaSuccess) return status_from_cuda(ret); + int32_t major = 0, minor = 0; + ret = cudaDeviceGetAttribute(&major, cudaDevAttrComputeCapabilityMajor, + device); + if (ret != cudaSuccess) return status_from_cuda(ret); + ret = cudaDeviceGetAttribute(&minor, cudaDevAttrComputeCapabilityMinor, + device); + if (ret != cudaSuccess) return status_from_cuda(ret); + if (major != 12 || minor != 0) return PEGAINFER_QWEN35_GDN_NOT_SUPPORTED; + int32_t sm_count = 0; + ret = cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, + device); + if (ret != cudaSuccess) return status_from_cuda(ret); + gdn_handle_t *owner = (gdn_handle_t *)calloc(1, sizeof(*owner)); + if (owner == NULL) return PEGAINFER_QWEN35_GDN_CUDA_ERROR; + owner->device = device; + owner->workspace_bytes = + (size_t)sm_count * PEGAINFER_QWEN35_GDN_WORKSPACE_BYTES_PER_SM; + int32_t rc = load_current_device(&owner->module, device); + if (rc != (int32_t)cudaSuccess) { + free(owner); + return PEGAINFER_QWEN35_GDN_CUDA_ERROR; + } + *handle = owner; + return PEGAINFER_QWEN35_GDN_OK; +#else + (void)device; + return PEGAINFER_QWEN35_GDN_NOT_SUPPORTED; +#endif +} + +int32_t pegainfer_qwen35_gdn_workspace_bytes(void *handle, + size_t *workspace_bytes) { + if (handle == NULL || workspace_bytes == NULL) + return PEGAINFER_QWEN35_GDN_INVALID_ARGUMENT; +#ifdef PEGAINFER_QWEN35_GDN_AOT + gdn_handle_t *owner = (gdn_handle_t *)handle; + *workspace_bytes = owner->workspace_bytes; + return PEGAINFER_QWEN35_GDN_OK; +#else + return PEGAINFER_QWEN35_GDN_NOT_SUPPORTED; +#endif +} + +int32_t pegainfer_qwen35_gdn_launch(void *handle, + const pegainfer_qwen35_gdn_args_t *args) { +#ifdef PEGAINFER_QWEN35_GDN_AOT + if (handle == NULL || args == NULL || + args->struct_size != sizeof(*args) || args->tokens == 0 || + args->tokens > (uint32_t)(INT32_MAX / 32) || + args->workspace_bytes > (size_t)INT32_MAX || + args->q == NULL || args->k == NULL || args->v == NULL || + args->output == NULL || args->alpha == NULL || args->beta == NULL || + args->state == NULL || args->initial_state == NULL || + args->workspace == NULL || args->cu_seqlens == NULL || + args->stream == NULL) { + return PEGAINFER_QWEN35_GDN_INVALID_ARGUMENT; + } + if (args->abi_version != PEGAINFER_QWEN35_GDN_ABI_VERSION) + return PEGAINFER_QWEN35_GDN_ABI_MISMATCH; + gdn_handle_t *owner = (gdn_handle_t *)handle; + cudaError_t cuda_rc = cudaSetDevice(owner->device); + if (cuda_rc != cudaSuccess) return status_from_cuda(cuda_rc); + if (args->workspace_bytes < owner->workspace_bytes) + return PEGAINFER_QWEN35_GDN_INVALID_ARGUMENT; + + int32_t tokens = (int32_t)args->tokens; + int32_t gates = tokens * 32; + int32_t workspace_bytes = (int32_t)args->workspace_bytes; + int32_t cu_count = 2; + pegainfer_qwen35_gdn_qwen35_4b_candidate_Tensor_g_q_t q = { + (void *)args->q, {tokens}}; + pegainfer_qwen35_gdn_qwen35_4b_candidate_Tensor_g_k_t k = { + (void *)args->k, {tokens}}; + pegainfer_qwen35_gdn_qwen35_4b_candidate_Tensor_g_v_t v = { + (void *)args->v, {tokens}}; + pegainfer_qwen35_gdn_qwen35_4b_candidate_Tensor_g_o_t output = { + args->output, {tokens}}; + pegainfer_qwen35_gdn_qwen35_4b_candidate_Tensor_g_alpha_t alpha = { + (void *)args->alpha, {gates}}; + pegainfer_qwen35_gdn_qwen35_4b_candidate_Tensor_g_beta_t beta = { + (void *)args->beta, {gates}}; + pegainfer_qwen35_gdn_qwen35_4b_candidate_Tensor_g_state_t state = { + args->state}; + pegainfer_qwen35_gdn_qwen35_4b_candidate_Tensor_g_init_state_t initial = { + (void *)args->initial_state}; + pegainfer_qwen35_gdn_qwen35_4b_candidate_Tensor_g_tensormaps_t workspace = { + args->workspace, {workspace_bytes}}; + pegainfer_qwen35_gdn_qwen35_4b_candidate_Tensor_cu_seqlens_t cu_seqlens = { + (void *)args->cu_seqlens, {cu_count}}; + int32_t rc = cute_dsl_pegainfer_qwen35_gdn_qwen35_4b_candidate_wrapper( + &owner->module, &q, &k, &v, &output, &alpha, &beta, &state, + &initial, &workspace, &cu_seqlens, 0.08838834764831845f, + 16, 16, 32, 32, 1, 1, 0, 32, (cudaStream_t)args->stream); + return rc == 0 ? PEGAINFER_QWEN35_GDN_OK + : PEGAINFER_QWEN35_GDN_CUDA_ERROR; +#else + (void)handle; + (void)args; + return PEGAINFER_QWEN35_GDN_NOT_SUPPORTED; +#endif +} + +void pegainfer_qwen35_gdn_destroy(void *handle) { +#ifdef PEGAINFER_QWEN35_GDN_AOT + if (handle != NULL) { + gdn_handle_t *owner = (gdn_handle_t *)handle; + cudaSetDevice(owner->device); + cudaLibraryUnload(owner->module.module); + free(owner); + } +#else + (void)handle; +#endif +} diff --git a/pegainfer-kernels/csrc/qwen35/flashinfer_gdn_aot.h b/pegainfer-kernels/csrc/qwen35/flashinfer_gdn_aot.h new file mode 100644 index 000000000..f7d952c7b --- /dev/null +++ b/pegainfer-kernels/csrc/qwen35/flashinfer_gdn_aot.h @@ -0,0 +1,86 @@ +#pragma once + +#include +#include + +#ifdef __cplusplus +extern "C" { +#endif + +#define PEGAINFER_QWEN35_GDN_ABI_VERSION 1u + +typedef enum { + PEGAINFER_QWEN35_GDN_OK = 0, + PEGAINFER_QWEN35_GDN_NOT_SUPPORTED = 1, + PEGAINFER_QWEN35_GDN_INVALID_ARGUMENT = 2, + PEGAINFER_QWEN35_GDN_ABI_MISMATCH = 3, + PEGAINFER_QWEN35_GDN_CUDA_ERROR = 4, +} pegainfer_qwen35_gdn_status_t; + +typedef struct { + uint32_t abi_version; + uint32_t struct_size; + const void *q; + const void *k; + const void *v; + void *output; + const void *alpha; + const void *beta; + void *state; + const void *initial_state; + void *workspace; + size_t workspace_bytes; + const int64_t *cu_seqlens; + uint32_t tokens; + void *stream; +} pegainfer_qwen35_gdn_args_t; + +#if defined(__cplusplus) +#define PEGAINFER_GDN_STATIC_ASSERT static_assert +#define PEGAINFER_GDN_ALIGNOF alignof +#else +#define PEGAINFER_GDN_STATIC_ASSERT _Static_assert +#define PEGAINFER_GDN_ALIGNOF _Alignof +#endif + +#define PEGAINFER_GDN_ASSERT_OFFSET(type, field, expected) \ + PEGAINFER_GDN_STATIC_ASSERT(offsetof(type, field) == (expected), \ + #type "." #field " ABI offset changed") + +PEGAINFER_GDN_STATIC_ASSERT(sizeof(pegainfer_qwen35_gdn_args_t) == 112, + "GDN args ABI size changed"); +PEGAINFER_GDN_STATIC_ASSERT(PEGAINFER_GDN_ALIGNOF(pegainfer_qwen35_gdn_args_t) == 8, + "GDN args ABI alignment changed"); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, abi_version, 0); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, struct_size, 4); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, q, 8); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, k, 16); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, v, 24); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, output, 32); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, alpha, 40); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, beta, 48); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, state, 56); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, initial_state, 64); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, workspace, 72); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, workspace_bytes, 80); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, cu_seqlens, 88); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, tokens, 96); +PEGAINFER_GDN_ASSERT_OFFSET(pegainfer_qwen35_gdn_args_t, stream, 104); + +#undef PEGAINFER_GDN_ASSERT_OFFSET +#undef PEGAINFER_GDN_STATIC_ASSERT +#undef PEGAINFER_GDN_ALIGNOF + +uint32_t pegainfer_qwen35_gdn_abi_version(void); +const char *pegainfer_qwen35_gdn_artifact_sha256(void); +int32_t pegainfer_qwen35_gdn_aot_available(void); +int32_t pegainfer_qwen35_gdn_workspace_bytes(void *handle, + size_t *workspace_bytes); +int32_t pegainfer_qwen35_gdn_create(void **handle, int32_t device); +int32_t pegainfer_qwen35_gdn_launch(void *handle, + const pegainfer_qwen35_gdn_args_t *args); +void pegainfer_qwen35_gdn_destroy(void *handle); + +#ifdef __cplusplus +} +#endif diff --git a/pegainfer-kernels/csrc/qwen35/gdn_prepare.cu b/pegainfer-kernels/csrc/qwen35/gdn_prepare.cu new file mode 100644 index 000000000..323a9704f --- /dev/null +++ b/pegainfer-kernels/csrc/qwen35/gdn_prepare.cu @@ -0,0 +1,155 @@ +#include "common.cuh" + +#include +#include + +namespace { + +constexpr int kHeadDim = 128; +constexpr int kThreads = 128; + +__device__ __forceinline__ float block_sum_128(float value) { + __shared__ float warp_sums[4]; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + value = warp_reduce_sum(value); + if (lane == 0) { + warp_sums[warp] = value; + } + __syncthreads(); + return warp_sums[0] + warp_sums[1] + warp_sums[2] + warp_sums[3]; +} + +__device__ __forceinline__ void record_non_finite(float value, uint32_t* status) { + if (!isfinite(value)) { + atomicExch(status, 1u); + } +} + +// Production Hv32 specialization. One block owns one native Q or K head and +// the corresponding V head: +// +// item [0,16) -> Q[item] + V[item] +// item [16,32) -> K[item-16] + V[item] +// +// Q and K retain independent reductions and output layouts. Pairing each with +// one V head removes the separate 32 V CTAs without expanding Q/K. A block +// reports all non-finite Q/K/V/gate inputs with at most one atomic update. +__global__ void gdn_prefill_native_prepare_hv32_kernel( + const __nv_bfloat16* __restrict__ qkv, // [T, 64*D] + const __nv_bfloat16* __restrict__ b_proj, // [T, 32] + const __nv_bfloat16* __restrict__ a_proj, // [T, 32] + const __nv_bfloat16* __restrict__ dt_bias, // [32] + const float* __restrict__ a_log, // [32] + __nv_bfloat16* __restrict__ q_out, // [T, 16, D] + __nv_bfloat16* __restrict__ k_out, // [T, 16, D] + __nv_bfloat16* __restrict__ v_out, // [T, 32, D] + float* __restrict__ alpha_out, // [T, 32] + float* __restrict__ beta_out, // [T, 32] + uint32_t* __restrict__ non_finite_status, + int qkv_dim, + int tokens) { + const int token = blockIdx.x; + const int item = blockIdx.y; + const int d = threadIdx.x; + if (token >= tokens) { + return; + } + + constexpr int kHq = 16; + constexpr int kHk = 16; + constexpr int kHv = 32; + const bool is_q = item < kHq; + const int qk_head = is_q ? item : item - kHq; + const int v_head = item; + const size_t token_base = static_cast(token) * qkv_dim; + const size_t qk_base = is_q ? 0 : static_cast(kHq) * kHeadDim; + const size_t v_base = static_cast(kHq + kHk) * kHeadDim; + + const float qk_value = + __bfloat162float(qkv[token_base + qk_base + qk_head * kHeadDim + d]); + const __nv_bfloat16 v = qkv[token_base + v_base + v_head * kHeadDim + d]; + const float v_value = __bfloat162float(v); + bool non_finite = !isfinite(qk_value) || !isfinite(v_value); + + const float inv_norm = rsqrtf(block_sum_128(qk_value * qk_value) + 1.0e-12f); + const __nv_bfloat16 normalized = __float2bfloat16(qk_value * inv_norm); + if (is_q) { + q_out[(static_cast(token) * kHq + qk_head) * kHeadDim + d] = + normalized; + } else { + k_out[(static_cast(token) * kHk + qk_head) * kHeadDim + d] = + normalized; + } + v_out[(static_cast(token) * kHv + v_head) * kHeadDim + d] = v; + + if (d == 0) { + const size_t gate_offset = static_cast(token) * kHv + v_head; + const float a = __bfloat162float(a_proj[gate_offset]); + const float b = __bfloat162float(b_proj[gate_offset]); + const float bias = __bfloat162float(dt_bias[v_head]); + const float log_a = a_log[v_head]; + non_finite |= + !isfinite(a) || !isfinite(b) || !isfinite(bias) || !isfinite(log_a); + + const float x = a + bias; + const float softplus = + x > 20.0f ? x : (x < -20.0f ? expf(x) : log1pf(expf(x))); + const float log_alpha = -expf(log_a) * softplus; + alpha_out[gate_offset] = expf(log_alpha); + const float exp_b = expf(b < 0.0f ? b : -b); + beta_out[gate_offset] = + b >= 0.0f ? 1.0f / (1.0f + exp_b) : exp_b / (1.0f + exp_b); + } + + if (__syncthreads_or(non_finite) && d == 0) { + atomicExch(non_finite_status, 1u); + } +} + +CUresult map_cuda_error(cudaError_t error) { + if (error == cudaSuccess) { + return CUDA_SUCCESS; + } + if (error == cudaErrorInvalidValue || error == cudaErrorInvalidDevicePointer) { + return CUDA_ERROR_INVALID_VALUE; + } + return CUDA_ERROR_UNKNOWN; +} + +} // namespace + +extern "C" CUresult gated_delta_rule_prefill_native_prepare_cuda( + const __nv_bfloat16* qkv, + const __nv_bfloat16* b_proj, + const __nv_bfloat16* a_proj, + const __nv_bfloat16* dt_bias, + const float* a_log, + __nv_bfloat16* q_out, + __nv_bfloat16* k_out, + __nv_bfloat16* v_out, + float* alpha_out, + float* beta_out, + uint32_t* non_finite_status, + int tokens, + cudaStream_t stream) { + constexpr int kHq = 16; + constexpr int kHk = 16; + constexpr int kHv = 32; + constexpr int kQkvDim = (kHq + kHk + kHv) * kHeadDim; + if (qkv == nullptr || b_proj == nullptr || a_proj == nullptr || dt_bias == nullptr || + a_log == nullptr || q_out == nullptr || k_out == nullptr || v_out == nullptr || + alpha_out == nullptr || beta_out == nullptr || non_finite_status == nullptr || + tokens <= 0) { + return CUDA_ERROR_INVALID_VALUE; + } + + // The chunk owner allocates this status word zeroed once. Every layer ORs + // into the same sticky status so the host can validate once at the chunk + // boundary instead of introducing one D2H synchronization per layer. + const dim3 grid(tokens, kHv); + gdn_prefill_native_prepare_hv32_kernel<<>>( + qkv, b_proj, a_proj, dt_bias, a_log, q_out, k_out, v_out, alpha_out, + beta_out, non_finite_status, kQkvDim, tokens); + return map_cuda_error(cudaGetLastError()); +} diff --git a/pegainfer-kernels/src/ffi/qwen35.rs b/pegainfer-kernels/src/ffi/qwen35.rs index c949dc186..669ea4da2 100644 --- a/pegainfer-kernels/src/ffi/qwen35.rs +++ b/pegainfer-kernels/src/ffi/qwen35.rs @@ -1,13 +1,85 @@ +#[cfg(feature = "qwen35")] +use std::ffi::c_char; +#[cfg(feature = "qwen35")] +use std::ffi::c_void; + #[cfg(feature = "qwen35")] use cudarc::driver::sys::CUresult; use cudarc::driver::sys::CUstream; use super::Half; +/// Kernels-private Rust mirror of the stable C ABI. Model crates never import +/// this struct: the safe `ops::Qwen35GdnAot` wrapper owns validation, workspace, +/// handle lifetime, and conversion from semantic tensors to device addresses. +#[cfg(feature = "qwen35")] +#[repr(C)] +#[derive(Clone, Copy, Debug)] +pub struct FlashInferGdnPrefillArgs { + pub abi_version: u32, + pub struct_size: u32, + pub q: u64, + pub k: u64, + pub v: u64, + pub output: u64, + pub alpha: u64, + pub beta: u64, + pub state: u64, + pub initial_state: u64, + pub workspace: u64, + pub workspace_bytes: u64, + pub cu_seqlens: u64, + pub tokens: u32, + pub stream: CUstream, +} + // Qwen3.5-4B private kernels (hybrid linear + HD256 full attention). // Sources: csrc/qwen35/*.cu. The paged HD256 attention entry points are shared // with Gemma 4 and are declared in `shared.rs`. unsafe extern "C" { + #[cfg(feature = "qwen35")] + pub fn pegainfer_qwen35_gdn_abi_version() -> u32; + #[cfg(feature = "qwen35")] + pub fn pegainfer_qwen35_gdn_artifact_sha256() -> *const c_char; + #[cfg(feature = "qwen35")] + pub fn pegainfer_qwen35_gdn_aot_available() -> i32; + #[cfg(feature = "qwen35")] + pub fn pegainfer_qwen35_gdn_create(handle: *mut *mut c_void, device: i32) -> i32; + #[cfg(feature = "qwen35")] + pub fn pegainfer_qwen35_gdn_workspace_bytes( + handle: *mut c_void, + workspace_bytes: *mut usize, + ) -> i32; + #[cfg(feature = "qwen35")] + pub fn pegainfer_qwen35_gdn_launch( + handle: *mut c_void, + args: *const FlashInferGdnPrefillArgs, + ) -> i32; + #[cfg(feature = "qwen35")] + pub fn pegainfer_qwen35_gdn_destroy(handle: *mut c_void); + + /// Native, non-expanded FlashInfer-GDN input preparation. + /// + /// `q_out`, `k_out`, and `v_out` are token-major `[T,H,D]`; alpha/beta are + /// FP32 `[T,Hv]`. `non_finite_status` is zeroed asynchronously and set to + /// one by the kernel if any consumed input is non-finite. + #[cfg(feature = "qwen35")] + pub fn gated_delta_rule_prefill_native_prepare_cuda( + qkv: *const Half, + b_proj: *const Half, + a_proj: *const Half, + dt_bias: *const Half, + a_log: *const f32, + q_out: *mut Half, + k_out: *mut Half, + v_out: *mut Half, + alpha_out: *mut f32, + beta_out: *mut f32, + non_finite_status: *mut u32, + tokens: i32, + stream: CUstream, + ) -> CUresult; + // Qwen3.5 full-attention prefill prep that writes K/V directly into paged KV. pub fn prefill_attention_hd256_prep_paged_cuda( q_full_batch: *const Half, diff --git a/pegainfer-kernels/src/ops.rs b/pegainfer-kernels/src/ops.rs index 085e9e53f..9820009c1 100644 --- a/pegainfer-kernels/src/ops.rs +++ b/pegainfer-kernels/src/ops.rs @@ -20,6 +20,8 @@ mod kimi_k2; mod linear; mod lora; mod norm; +#[cfg(feature = "qwen35")] +mod qwen35; mod sampling; pub use attention::Hd512DecodeMetadata; @@ -178,6 +180,8 @@ pub use norm::rms_norm_gated_batch_into; pub use norm::rms_norm_into; pub use norm::rms_norm_offset_into; pub use norm::rms_norm_rows_into; +#[cfg(feature = "qwen35")] +pub use qwen35::*; pub use sampling::BatchSamplingRow; pub use sampling::BatchSamplingScratch; pub use sampling::argmax; diff --git a/pegainfer-kernels/src/ops/qwen35.rs b/pegainfer-kernels/src/ops/qwen35.rs new file mode 100644 index 000000000..5d3e086b8 --- /dev/null +++ b/pegainfer-kernels/src/ops/qwen35.rs @@ -0,0 +1,267 @@ +//! Stable Qwen3.5 GDN prefill boundary. +//! +//! Generated CuTe symbols, tensor wrappers, TMA descriptors, module lifetime, +//! and the low-level launch ABI stop below this module. Model crates see only +//! the semantic geometry and device buffers used by Gated DeltaNet prefill. + +use std::ffi::CStr; +use std::ffi::c_void; +use std::ptr::NonNull; + +use anyhow::Context; +use anyhow::Result; +use anyhow::ensure; +use cudarc::driver::CudaSlice; +use cudarc::driver::DevicePtr; +use cudarc::driver::DevicePtrMut; + +use crate::ffi; +use crate::tensor::DeviceContext; +use crate::tensor::HiddenStates; + +const QWEN35_GDN_ABI_VERSION: u32 = 1; +const STATUS_OK: i32 = 0; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct Qwen35GdnGeometry { + pub h_q: usize, + pub h_k: usize, + pub h_v: usize, + pub head_dim: usize, +} + +impl Qwen35GdnGeometry { + pub const PRODUCTION: Self = Self { + h_q: 16, + h_k: 16, + h_v: 32, + head_dim: 128, + }; +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum Qwen35GdnSupport { + Supported, + UnsupportedSm, + UnsupportedGeometry, +} + +fn qwen35_gdn_capability(sm: i32, geometry: Qwen35GdnGeometry) -> Qwen35GdnSupport { + if sm != 120 { + Qwen35GdnSupport::UnsupportedSm + } else if geometry != Qwen35GdnGeometry::PRODUCTION { + Qwen35GdnSupport::UnsupportedGeometry + } else { + Qwen35GdnSupport::Supported + } +} + +#[derive(Debug)] +pub struct Qwen35GdnAot { + handle: NonNull, + device_ordinal: usize, + geometry: Qwen35GdnGeometry, + workspace_bytes: usize, +} + +pub struct Qwen35GdnWorkspace { + workspace: CudaSlice, + cu_seqlens: CudaSlice, + tokens: usize, +} + +// The handle is bound to one CUDA device and all launches are issued by the +// owning model thread on its DeviceContext stream. +unsafe impl Send for Qwen35GdnAot {} + +impl Qwen35GdnAot { + pub fn load_for_production( + ctx: &DeviceContext, + geometry: Qwen35GdnGeometry, + ) -> Result> { + let (major, minor) = ctx.ctx.compute_capability()?; + let sm = major * 10 + minor; + if qwen35_gdn_capability(sm, geometry) != Qwen35GdnSupport::Supported { + return Ok(None); + } + ensure!( + unsafe { ffi::pegainfer_qwen35_gdn_abi_version() } == QWEN35_GDN_ABI_VERSION, + "Qwen3.5 GDN stable C ABI version mismatch" + ); + ensure!( + unsafe { ffi::pegainfer_qwen35_gdn_aot_available() } == 1, + "SM120/Hv32 selects FlashInfer GDN, but the validated prebuilt AOT artifact was not linked; set PEGAINFER_QWEN35_GDN_AOT_BUNDLE at build time" + ); + let mut raw = std::ptr::null_mut(); + let status = + unsafe { ffi::pegainfer_qwen35_gdn_create(&raw mut raw, ctx.device_ordinal as i32) }; + ensure!( + status == STATUS_OK, + "Qwen3.5 GDN preload failed with stable ABI status {status}" + ); + let handle = NonNull::new(raw).context("Qwen3.5 GDN preload returned a null handle")?; + let mut workspace_bytes = 0; + let status = unsafe { + ffi::pegainfer_qwen35_gdn_workspace_bytes(handle.as_ptr(), &raw mut workspace_bytes) + }; + if status != STATUS_OK { + unsafe { ffi::pegainfer_qwen35_gdn_destroy(handle.as_ptr()) }; + anyhow::bail!("Qwen3.5 GDN workspace query failed with stable ABI status {status}"); + } + Ok(Some(Self { + handle, + device_ordinal: ctx.device_ordinal, + geometry, + workspace_bytes, + })) + } + + pub fn artifact_sha256(&self) -> &'static str { + let pointer = unsafe { ffi::pegainfer_qwen35_gdn_artifact_sha256() }; + if pointer.is_null() { + return "unavailable"; + } + unsafe { CStr::from_ptr(pointer) } + .to_str() + .unwrap_or("invalid-utf8") + } + + pub fn allocate_workspace( + &self, + ctx: &DeviceContext, + tokens: usize, + ) -> Result { + ensure!(tokens > 0, "Qwen3.5 GDN workspace requires T>=1"); + let workspace = ctx + .stream + .alloc_zeros(self.workspace_bytes) + .map_err(|error| anyhow::anyhow!("allocate Qwen3.5 GDN workspace: {error}"))?; + let end = i64::try_from(tokens).context("Qwen3.5 GDN T exceeds i64")?; + let cu_seqlens = ctx + .stream + .clone_htod(&[0_i64, end]) + .map_err(|error| anyhow::anyhow!("upload Qwen3.5 GDN sequence metadata: {error}"))?; + Ok(Qwen35GdnWorkspace { + workspace, + cu_seqlens, + tokens, + }) + } + + #[allow(clippy::too_many_arguments)] + pub fn launch_in_place( + &self, + ctx: &DeviceContext, + q: &HiddenStates, + k: &HiddenStates, + v: &HiddenStates, + alpha: &CudaSlice, + beta: &CudaSlice, + state: &mut CudaSlice, + output: &mut HiddenStates, + launch_workspace: &mut Qwen35GdnWorkspace, + ) -> Result<()> { + let state_elements = self.geometry.h_v * self.geometry.head_dim * self.geometry.head_dim; + ensure!( + state.len() == state_elements, + "Qwen3.5 GDN state length mismatch" + ); + let (state_ptr, _state) = state.device_ptr_mut(&ctx.stream); + self.launch_with_state_pointers( + ctx, + q, + k, + v, + alpha, + beta, + state_ptr, + state_ptr, + output, + launch_workspace, + ) + } + + #[allow(clippy::too_many_arguments)] + fn launch_with_state_pointers( + &self, + ctx: &DeviceContext, + q: &HiddenStates, + k: &HiddenStates, + v: &HiddenStates, + alpha: &CudaSlice, + beta: &CudaSlice, + state_ptr: u64, + initial_state_ptr: u64, + output: &mut HiddenStates, + launch_workspace: &mut Qwen35GdnWorkspace, + ) -> Result<()> { + let t = q.seq_len; + let g = self.geometry; + ensure!( + ctx.device_ordinal == self.device_ordinal, + "Qwen3.5 GDN device mismatch" + ); + ensure!( + t > 0 && k.seq_len == t && v.seq_len == t && output.seq_len == t, + "Qwen3.5 GDN token extents do not match" + ); + ensure!( + q.hidden_dim == g.h_q * g.head_dim + && k.hidden_dim == g.h_k * g.head_dim + && v.hidden_dim == g.h_v * g.head_dim + && output.hidden_dim == g.h_v * g.head_dim, + "Qwen3.5 GDN tensor geometry mismatch" + ); + ensure!( + alpha.len() == t * g.h_v + && beta.len() == t * g.h_v + && launch_workspace.workspace.len() >= self.workspace_bytes + && launch_workspace.cu_seqlens.len() == 2 + && launch_workspace.tokens == t, + "Qwen3.5 GDN buffer contract mismatch" + ); + + let (q_ptr, _q) = q.data.device_ptr(&ctx.stream); + let (k_ptr, _k) = k.data.device_ptr(&ctx.stream); + let (v_ptr, _v) = v.data.device_ptr(&ctx.stream); + let (alpha_ptr, _alpha) = alpha.device_ptr(&ctx.stream); + let (beta_ptr, _beta) = beta.device_ptr(&ctx.stream); + let (output_ptr, _output) = output.data.device_ptr_mut(&ctx.stream); + let workspace_bytes = launch_workspace.workspace.len() as u64; + let (workspace_ptr, _workspace) = launch_workspace.workspace.device_ptr_mut(&ctx.stream); + let (cu_ptr, _cu) = launch_workspace.cu_seqlens.device_ptr(&ctx.stream); + let args = ffi::FlashInferGdnPrefillArgs { + abi_version: QWEN35_GDN_ABI_VERSION, + struct_size: size_of::() as u32, + q: q_ptr, + k: k_ptr, + v: v_ptr, + output: output_ptr, + alpha: alpha_ptr, + beta: beta_ptr, + state: state_ptr, + initial_state: initial_state_ptr, + workspace: workspace_ptr, + workspace_bytes, + cu_seqlens: cu_ptr, + tokens: t.try_into().context("Qwen3.5 GDN T exceeds u32")?, + stream: ctx.stream.cu_stream(), + }; + let status = + unsafe { ffi::pegainfer_qwen35_gdn_launch(self.handle.as_ptr(), &raw const args) }; + ensure!( + status == STATUS_OK, + "Qwen3.5 GDN launch failed with stable ABI status {status}" + ); + Ok(()) + } +} + +impl Drop for Qwen35GdnAot { + fn drop(&mut self) { + unsafe { ffi::pegainfer_qwen35_gdn_destroy(self.handle.as_ptr()) }; + } +} + +#[cfg(test)] +mod tests; diff --git a/pegainfer-kernels/src/ops/qwen35/tests.rs b/pegainfer-kernels/src/ops/qwen35/tests.rs new file mode 100644 index 000000000..b85bec869 --- /dev/null +++ b/pegainfer-kernels/src/ops/qwen35/tests.rs @@ -0,0 +1,228 @@ +use half::bf16; + +use super::*; + +impl Qwen35GdnAot { + #[allow(clippy::too_many_arguments)] + fn launch_separate_for_test( + &self, + ctx: &DeviceContext, + q: &HiddenStates, + k: &HiddenStates, + v: &HiddenStates, + alpha: &CudaSlice, + beta: &CudaSlice, + initial_state: &CudaSlice, + state: &mut CudaSlice, + output: &mut HiddenStates, + launch_workspace: &mut Qwen35GdnWorkspace, + ) -> Result<()> { + let state_elements = self.geometry.h_v * self.geometry.head_dim * self.geometry.head_dim; + ensure!( + initial_state.len() == state_elements && state.len() == state_elements, + "Qwen3.5 GDN separate-state length mismatch" + ); + let (initial_state_ptr, _initial_state) = initial_state.device_ptr(&ctx.stream); + let (state_ptr, _state) = state.device_ptr_mut(&ctx.stream); + self.launch_with_state_pointers( + ctx, + q, + k, + v, + alpha, + beta, + state_ptr, + initial_state_ptr, + output, + launch_workspace, + ) + } +} + +fn ensure_bitwise_f32(label: &str, expected: &[f32], actual: &[f32]) -> Result<()> { + ensure!( + expected.len() == actual.len(), + "{label} length mismatch: expected {}, actual {}", + expected.len(), + actual.len() + ); + if let Some(index) = expected + .iter() + .zip(actual) + .position(|(expected, actual)| expected.to_bits() != actual.to_bits()) + { + anyhow::bail!( + "{label} first bitwise mismatch at {index}: expected={} actual={}", + expected[index], + actual[index] + ); + } + eprintln!("{label}: elements={} bitwise_mismatches=0", expected.len()); + Ok(()) +} + +fn assert_stable_c_struct_layout() { + macro_rules! assert_offsets { + ($ty:ty, {$($field:ident: $offset:expr),+ $(,)?}) => { + $(assert_eq!(std::mem::offset_of!($ty, $field), $offset);)+ + }; + } + + assert_eq!(size_of::(), 112); + assert_eq!(align_of::(), 8); + assert_offsets!(ffi::FlashInferGdnPrefillArgs, { + abi_version: 0, struct_size: 4, q: 8, k: 16, v: 24, output: 32, + alpha: 40, beta: 48, state: 56, initial_state: 64, workspace: 72, + workspace_bytes: 80, cu_seqlens: 88, tokens: 96, stream: 104, + }); +} + +#[test] +fn stable_c_struct_layout_is_frozen() { + assert_stable_c_struct_layout(); +} + +#[test] +#[ignore = "requires an SM120 GPU and a build-linked validated FlashInfer GDN AOT bundle"] +fn sm120_stable_abi_alias_and_separate_state_are_bitwise_identical() -> Result<()> { + assert_stable_c_struct_layout(); + let ctx = DeviceContext::new()?; + let unsupported = Qwen35GdnGeometry { + h_v: 48, + ..Qwen35GdnGeometry::PRODUCTION + }; + ensure!( + Qwen35GdnAot::load_for_production(&ctx, unsupported)?.is_none(), + "production load boundary accepted unsupported Hv48 geometry on SM120" + ); + + let geometry = Qwen35GdnGeometry::PRODUCTION; + let backend = Qwen35GdnAot::load_for_production(&ctx, geometry)? + .context("validated FlashInfer GDN AOT bundle is not available on SM120")?; + ensure!( + backend.artifact_sha256() != "unavailable" + && backend.artifact_sha256() != "invalid-utf8" + && backend.artifact_sha256().len() == 64, + "production boundary did not expose a linked object SHA-256" + ); + let bf16_values = |elements: usize, modulus: usize, scale: f32| { + (0..elements) + .map(|index| { + let signed = (index % modulus) as i32 - (modulus / 2) as i32; + bf16::from_f32(signed as f32 * scale) + }) + .collect::>() + }; + let state_elements = geometry.h_v * geometry.head_dim * geometry.head_dim; + let initial_host = (0..geometry.h_v) + .flat_map(|head| { + (0..geometry.head_dim).flat_map(move |key| { + (0..geometry.head_dim) + .map(move |value| (head * 100_000 + key * 100 + value) as f32 * 1.0e-6) + }) + }) + .collect::>(); + ensure!( + initial_host.len() == state_elements, + "HKV fixture size mismatch" + ); + + for tokens in [1_usize, 63, 64, 65, 128] { + let q = HiddenStates::from_host( + &ctx, + &bf16_values(tokens * geometry.h_q * geometry.head_dim, 127, 1.0 / 1024.0), + geometry.h_q * geometry.head_dim, + tokens, + )?; + let k = HiddenStates::from_host( + &ctx, + &bf16_values(tokens * geometry.h_k * geometry.head_dim, 113, 1.0 / 1024.0), + geometry.h_k * geometry.head_dim, + tokens, + )?; + let v = HiddenStates::from_host( + &ctx, + &bf16_values(tokens * geometry.h_v * geometry.head_dim, 97, 1.0 / 128.0), + geometry.h_v * geometry.head_dim, + tokens, + )?; + let alpha = ctx + .stream + .clone_htod(&vec![0.9921875_f32; tokens * geometry.h_v])?; + let beta = ctx + .stream + .clone_htod(&vec![0.5_f32; tokens * geometry.h_v])?; + + let initial_state = ctx.stream.clone_htod(&initial_host)?; + let mut separate_state: CudaSlice = ctx.stream.alloc_zeros(state_elements)?; + let mut separate_output = + HiddenStates::zeros(&ctx, geometry.h_v * geometry.head_dim, tokens)?; + let mut separate_workspace = backend.allocate_workspace(&ctx, tokens)?; + backend.launch_separate_for_test( + &ctx, + &q, + &k, + &v, + &alpha, + &beta, + &initial_state, + &mut separate_state, + &mut separate_output, + &mut separate_workspace, + )?; + + let separate_output = separate_output.to_host(&ctx)?; + let separate_state = ctx.stream.clone_dtoh(&separate_state)?; + ctx.sync()?; + + ensure!( + separate_output.iter().all(|value| value.is_finite()), + "stable C ABI output contains a non-finite value at T={tokens}" + ); + ensure!( + separate_state.iter().all(|value| value.is_finite()), + "stable C ABI final state contains a non-finite value at T={tokens}" + ); + ensure!( + separate_output.iter().any(|&value| value != 0.0), + "stable C ABI output remained zero at T={tokens}" + ); + ensure!( + separate_state != initial_host, + "stable C ABI recurrent state did not update at T={tokens}" + ); + + if tokens == 65 { + let mut alias_state = ctx.stream.clone_htod(&initial_host)?; + let mut alias_output = + HiddenStates::zeros(&ctx, geometry.h_v * geometry.head_dim, tokens)?; + let mut alias_workspace = backend.allocate_workspace(&ctx, tokens)?; + backend.launch_in_place( + &ctx, + &q, + &k, + &v, + &alpha, + &beta, + &mut alias_state, + &mut alias_output, + &mut alias_workspace, + )?; + let alias_output = alias_output.to_host(&ctx)?; + let alias_state = ctx.stream.clone_dtoh(&alias_state)?; + ctx.sync()?; + ensure_bitwise_f32( + "stable C ABI alias/separate output [T=65,Hv=32,D=128,bf16]", + &separate_output, + &alias_output, + )?; + ensure_bitwise_f32( + "stable C ABI alias/separate final state [T=65,Hv=32,D=128,f32,HKV]", + &separate_state, + &alias_state, + )?; + } + } + + Ok(()) +} diff --git a/pegainfer-kernels/third_party/flashinfer b/pegainfer-kernels/third_party/flashinfer index 19f1a41e6..a0efa0adf 160000 --- a/pegainfer-kernels/third_party/flashinfer +++ b/pegainfer-kernels/third_party/flashinfer @@ -1 +1 @@ -Subproject commit 19f1a41e6b21f0c422d775e377b6fdf9a1fc9d23 +Subproject commit a0efa0adfe49bb836ab1a147d6572980b870f3d4 diff --git a/pegainfer-kernels/tools/flashinfer_gdn/README.md b/pegainfer-kernels/tools/flashinfer_gdn/README.md new file mode 100644 index 000000000..cee5c7eb1 --- /dev/null +++ b/pegainfer-kernels/tools/flashinfer_gdn/README.md @@ -0,0 +1,130 @@ +# FlashInfer GDN SM120 AOT candidate + +This directory owns the generation-only FlashInfer/CuTe environment for the +Qwen3.5 GDN prefill specialization. Serving does not import Python, CuTe, +FlashInfer, or Triton for this kernel and does not load PTX. The release build +validates the manifest, then statically links the exported native object and +`libcuda_dialect_runtime_static.a` behind the stable PegaInfer C ABI. + +The source lock pins FlashInfer and the small HKV state-layout specialization. +The generator emits only the production Hv32 candidate. SM120 + +Hq/Hk/Hv/D=`16/16/32/128`, BF16 inputs, FP32 HKV state, single GPU is eligible +for production selection. Other capabilities retain the Triton path. An +eligible configuration requires the validated AOT object and does not silently +fall back when that object was not linked. + +## Contract-validated local generation + +Run these commands from the repository root. The artifact contract currently +requires Python 3.12.3 exactly; use an interpreter with that version rather +than the serving or Triton environment. + +Initialize the pinned FlashInfer source and create an isolated generation +environment: + +```bash +git submodule update --init \ + pegainfer-kernels/third_party/flashinfer + +python3 -c \ + 'import sys; assert sys.version.split()[0] == "3.12.3", sys.version' +python3 -m venv target/flashinfer-gdn-cu13-venv + +export PEGAINFER_GDN_AOT_PYTHON="$PWD/target/flashinfer-gdn-cu13-venv/bin/python" + +"$PEGAINFER_GDN_AOT_PYTHON" -m pip install \ + -r pegainfer-kernels/tools/flashinfer_gdn/requirements-cu13.lock +``` + +The version assertion stops immediately unless `python3` is exactly Python +3.12.3. The generation-only CUDA 13 packages are pinned in +`requirements-cu13.lock`; they are not serving dependencies. Retired CUDA +12.8/PTX generation workflows are not supported. + +Generate a fresh production-only candidate. The generator refuses to overwrite +an existing output directory, so remove or rename an old local output before +reusing the same path. + +```bash +"$PEGAINFER_GDN_AOT_PYTHON" \ + pegainfer-kernels/tools/flashinfer_gdn/generate.py \ + --python "$PEGAINFER_GDN_AOT_PYTHON" \ + --flashinfer-dir pegainfer-kernels/third_party/flashinfer \ + --output target/flashinfer-gdn-sm120 +``` + +The output directory itself is the only generated candidate and directly +contains `manifest.json`, `kernel.h`, `kernel.o`, and the static runtime archive. + +The contract pins and validates the source, patch, generator, package versions, +compiler metadata, ABI, geometry, and hashes recorded by the candidate. Repeated +generation has been byte-identical on the same host, but cross-host object +identity is not currently guaranteed. Release distribution must therefore +preserve and validate the complete candidate and its manifest rather than assume a +globally fixed object hash. + +Validate a generated or downloaded candidate against its pinned source: + +```bash +python3 pegainfer-kernels/tools/flashinfer_gdn/artifact_contract.py \ + validate-candidate target/flashinfer-gdn-sm120 \ + --flashinfer-dir pegainfer-kernels/third_party/flashinfer +``` + +For a Qwen3.5 release build, point the kernel build at the validated variant. +Qwen3.5 still needs its normal build-time Triton AOT environment; see +[`../triton/README.md`](../triton/README.md). + +```bash +export PEGAINFER_QWEN35_GDN_AOT_BUNDLE="$PWD/target/flashinfer-gdn-sm120" +export PEGAINFER_CUDA_SM=120 +export PEGAINFER_TRITON_PYTHON="$PWD/.venv/bin/python" + +cargo build --release \ + -p pegainfer-server \ + --no-default-features \ + --features qwen35 \ + --bin pegainfer +``` + +The production confidence gate is not the Python packager validating itself. +On an SM120 runner with the pinned model snapshot, invoke the canonical runner; +it validates the real candidate, builds through production `build.rs`, and runs +the five exact GPU gates with fail-on-skip/test-count checks. Use a separate +Python 3.12 environment with Triton 3.7.1 for the production Qwen3.5 build; +do not reuse the Torch 2.7.1/CuTe generation environment for Triton AOT: + +```bash +env \ + PEGAINFER_CUDA_SM=120 \ + PEGAINFER_TRITON_PYTHON="$PWD/target/flashinfer-gdn-triton-venv/bin/python" \ + PEGAINFER_QWEN35_GDN_AOT_BUNDLE="$PWD/target/flashinfer-gdn-sm120" \ + PEGAINFER_TEST_MODEL_PATH="$PWD/models/Qwen3.5-4B" \ + PEGAINFER_TEST_MODEL_REVISION=851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a \ + CARGO_TARGET_DIR="$PWD/target/gdn-production-gates" \ + pegainfer-qwen35/tools/run_gdn_production_gates.sh +``` + +Set the optional `PEGAINFER_GDN_EXPECT_BRANCH` when the run must be pinned to a +specific local review branch. The runner rejects a mismatch rather than +silently validating another checkout. + +When `PEGAINFER_QWEN35_GDN_AOT_BUNDLE` is set, `pegainfer-kernels/build.rs` +rechecks schema, SM, geometry, ABI, object/header/runtime hashes and sizes. A +missing, incomplete, or incompatible selected path fails the build instead of +silently linking a different kernel. + +When the variable is not set, the build contains no FlashInfer GDN object. +Unsupported SM, geometry, or tensor-parallel configurations still use the +explicit Triton capability fallback. The supported SM120/Hv32/single-GPU +configuration instead fails model startup with a missing-AOT error; it does not +silently change backend. The model crate never receives the candidate path and +sees only a semantic GDN operation. + +CUDA Graph and successful-GDN-launch evidence used by the five-gate runner is +compiled only with the non-default `pegainfer-qwen35/gdn-validation` feature. +Default serving objects contain neither those counters nor their public +validation API. + +Generated headers, objects, static archives, candidates, model weights, `target/`, +logs, and benchmark JSON are release/build artifacts and must not be committed. diff --git a/pegainfer-kernels/tools/flashinfer_gdn/__init__.py b/pegainfer-kernels/tools/flashinfer_gdn/__init__.py new file mode 100644 index 000000000..f9e40e2db --- /dev/null +++ b/pegainfer-kernels/tools/flashinfer_gdn/__init__.py @@ -0,0 +1 @@ +"""Reproducible FlashInfer GDN SM120 artifact tooling.""" diff --git a/pegainfer-kernels/tools/flashinfer_gdn/artifact_contract.py b/pegainfer-kernels/tools/flashinfer_gdn/artifact_contract.py new file mode 100644 index 000000000..1fda4c8ef --- /dev/null +++ b/pegainfer-kernels/tools/flashinfer_gdn/artifact_contract.py @@ -0,0 +1,502 @@ +#!/usr/bin/env python3 +"""Prepare, package, and validate the single production GDN AOT candidate.""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import re +import shutil +import subprocess +import sys +from pathlib import Path +from typing import Any + + +SCHEMA_VERSION = 3 +VARIANT = "qwen35_4b_candidate" +TARGET_ARCH = "sm_120a" +FROZEN_FLASHINFER_COMMIT = "a0efa0adfe49bb836ab1a147d6572980b870f3d4" +GEOMETRY = {"h_q": 16, "h_k": 16, "h_v": 32, "head_dim": 128} +TOKENS = {"extent": "dynamic", "minimum": 1} +WORKSPACE = {"kind": "per_sm", "bytes_per_sm": 128, "alignment_bytes": 128} +DTYPES = { + "q": "bfloat16", + "k": "bfloat16", + "v": "bfloat16", + "o": "bfloat16", + "alpha": "float32", + "beta": "float32", + "state": "float32", + "cu_seqlens": "int64", + "workspace": "uint8", +} +PINNED_TOOLCHAIN = { + "python": "3.12.3", + "ptx_compiler_release": "13.1", + "ptx_compiler_version": "13.1.66", + "ptx_isa": "9.1", + "cutlass_dsl": "4.5.0", + "cutlass_dsl_libs_base": "4.5.0", + "torch": "2.7.1", + "cuda_python": "13.0.1", + "cuda_bindings": "13.0.3", + "cuda_pathfinder": "1.6.0", +} +KERNEL_SOURCE = "flashinfer/gdn_kernels/delta_rule_dsl/delta_rule_sm120.py" +ARTIFACT_FILES = { + "header": "kernel.h", + "object": "kernel.o", + "native_runtime": "libcuda_dialect_runtime_static.a", +} +FORBIDDEN_TMA_CLUSTER_LOAD = ( + "cp.async.bulk.tensor.3d.shared::cluster.global.tile." + "mbarrier::complete_tx::bytes.L2::cache_hint" +) +class ContractError(RuntimeError): + """An artifact or source contract is invalid.""" + + +def sha256_bytes(data: bytes) -> str: + return hashlib.sha256(data).hexdigest() + + +def sha256_file(path: Path) -> str: + return sha256_bytes(path.read_bytes()) + + +def read_json(path: Path) -> dict[str, Any]: + try: + value = json.loads(path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise ContractError(f"cannot read JSON {path}: {exc}") from exc + if not isinstance(value, dict): + raise ContractError(f"expected a JSON object in {path}") + return value + + +def write_json(path: Path, value: dict[str, Any]) -> None: + path.write_text( + json.dumps(value, indent=2, sort_keys=True, ensure_ascii=True) + "\n", + encoding="utf-8", + ) + + +def source_lock_path() -> Path: + return Path(__file__).with_name("source-lock.json") + + +def requirements_lock_path() -> Path: + return Path(__file__).with_name("requirements-cu13.lock") + + +def compiler_path() -> Path: + return Path(__file__).with_name("compile_sm120.py") + + +def load_source_lock(path: Path | None = None) -> tuple[dict[str, Any], str]: + path = path or source_lock_path() + lock = read_json(path) + if lock.get("schema_version") != SCHEMA_VERSION: + raise ContractError("source lock schema_version mismatch") + if lock.get("flashinfer_commit") != FROZEN_FLASHINFER_COMMIT: + raise ContractError("source lock FlashInfer commit mismatch") + patches = lock.get("patches") + if not isinstance(patches, list) or len(patches) != 1: + raise ContractError("source lock must contain exactly one HKV patch") + patch = patches[0] + if not isinstance(patch, dict): + raise ContractError("source lock patch entry must be an object") + patch_relative = patch.get("path") + if patch_relative != "patches/0001-openinfer-hkv-state-layout.patch": + raise ContractError("source lock HKV patch path mismatch") + patch_path = path.parent / patch_relative + if not patch_path.is_file(): + raise ContractError(f"source lock patch is missing: {patch_path}") + if patch.get("sha256") != sha256_file(patch_path): + raise ContractError("source lock HKV patch hash mismatch") + patched_kernel_sha256 = lock.get("patched_kernel_sha256") + if not isinstance(patched_kernel_sha256, str) or len(patched_kernel_sha256) != 64: + raise ContractError("source lock patched kernel hash is missing") + return lock, sha256_file(path) + + +def run_git(flashinfer_dir: Path, *args: str) -> str: + result = subprocess.run( + ["git", "-C", str(flashinfer_dir), *args], + check=False, + capture_output=True, + text=True, + ) + if result.returncode != 0: + detail = result.stderr.strip() or result.stdout.strip() + raise ContractError(f"git {' '.join(args)} failed for {flashinfer_dir}: {detail}") + return result.stdout.strip() + + +def verify_flashinfer_base(flashinfer_dir: Path) -> str: + flashinfer_dir = flashinfer_dir.resolve() + commit = run_git(flashinfer_dir, "rev-parse", "HEAD") + if commit != FROZEN_FLASHINFER_COMMIT: + raise ContractError( + f"FlashInfer SHA mismatch: expected {FROZEN_FLASHINFER_COMMIT}, got {commit}" + ) + dirty = run_git(flashinfer_dir, "status", "--porcelain", "--untracked-files=no") + if dirty: + raise ContractError("FlashInfer tracked source is dirty") + return commit + + +def inspect_kernel_source(source_dir: Path, commit: str) -> dict[str, Any]: + kernel_path = source_dir / KERNEL_SOURCE + if not kernel_path.is_file(): + raise ContractError(f"patched GDN kernel is missing: {kernel_path}") + lock, source_lock_sha256 = load_source_lock() + kernel_sha256 = sha256_file(kernel_path) + _require_equal( + kernel_sha256, lock["patched_kernel_sha256"], "patched GDN kernel hash" + ) + return { + "flashinfer_commit": commit, + "kernel_source_sha256": kernel_sha256, + "source_lock_sha256": source_lock_sha256, + } + + +def prepare_flashinfer_source(flashinfer_dir: Path, destination: Path) -> dict[str, Any]: + commit = verify_flashinfer_base(flashinfer_dir) + lock, _ = load_source_lock() + if destination.exists(): + raise ContractError(f"refusing to overwrite prepared source: {destination}") + shutil.copytree(flashinfer_dir / "flashinfer", destination / "flashinfer") + for patch in lock["patches"]: + patch_path = source_lock_path().parent / patch["path"] + result = subprocess.run( + ["git", "apply", "--unsafe-paths", str(patch_path)], + cwd=destination, + check=False, + capture_output=True, + text=True, + ) + if result.returncode != 0: + detail = result.stderr.strip() or result.stdout.strip() + raise ContractError(f"failed to apply HKV patch: {detail}") + return inspect_kernel_source(destination, commit) + + +def verify_prepared_flashinfer_source( + source_dir: Path, flashinfer_dir: Path +) -> dict[str, Any]: + commit = verify_flashinfer_base(flashinfer_dir) + lock, _ = load_source_lock() + source = inspect_kernel_source(source_dir, commit) + _require_equal( + source["kernel_source_sha256"], + lock["patched_kernel_sha256"], + "prepared HKV kernel hash", + ) + return source + + +def normalize_ptx(ptx: str) -> str: + """Normalize harmless path/debug text without changing PTX instructions.""" + normalized_lines: list[str] = [] + file_directive = re.compile(r'^(\s*\.file\s+\d+\s+")([^"]+)(".*)$') + for raw_line in ptx.replace("\r\n", "\n").replace("\r", "\n").splitlines(): + line = raw_line.rstrip() + match = file_directive.match(line) + if match: + name = Path(match.group(2).replace("\\", "/")).name + line = f"{match.group(1)}{name}{match.group(3)}" + normalized_lines.append(line) + return "\n".join(normalized_lines) + "\n" + + +def expected_spec(variant: str) -> dict[str, Any]: + _require_equal(variant, VARIANT, "artifact variant") + return { + "variant": VARIANT, + "target_arch": TARGET_ARCH, + "geometry": dict(GEOMETRY), + "dtypes": dict(DTYPES), + "tokens": dict(TOKENS), + } + + +def _require_equal(actual: Any, expected: Any, label: str) -> None: + if actual != expected: + raise ContractError(f"{label} mismatch: expected {expected!r}, got {actual!r}") + + +def validate_compile_metadata( + metadata: dict[str, Any], source: dict[str, Any] +) -> None: + _require_equal( + set(metadata), + { + "flashinfer_commit", + "kernel_source_sha256", + "source_lock_sha256", + "generator_sha256", + "requirements_lock_sha256", + "toolchain", + "aot", + }, + "compile metadata keys", + ) + for key in ("flashinfer_commit", "kernel_source_sha256", "source_lock_sha256"): + _require_equal(metadata.get(key), source[key], f"compile metadata {key}") + _require_equal( + metadata.get("generator_sha256"), + sha256_file(compiler_path()), + "compile metadata generator hash", + ) + _require_equal( + metadata.get("requirements_lock_sha256"), + sha256_file(requirements_lock_path()), + "compile metadata requirements lock hash", + ) + aot = metadata.get("aot") + if not isinstance(aot, dict): + raise ContractError("compile metadata is missing AOT export metadata") + toolchain = metadata.get("toolchain") + if not isinstance(toolchain, dict): + raise ContractError("compile metadata is missing toolchain") + _require_equal(toolchain, PINNED_TOOLCHAIN, "compile metadata toolchain") + + +def build_manifest( + *, + variant: str, + header_bytes: bytes, + object_bytes: bytes, + runtime_bytes: bytes, + compile_metadata: dict[str, Any], + source: dict[str, Any], +) -> dict[str, Any]: + spec = expected_spec(variant) + return { + "schema_version": SCHEMA_VERSION, + "artifact_kind": "flashinfer_cute_gdn_prefill_aot_object", + "variant": variant, + "target": {"arch": TARGET_ARCH, "code_object": "embedded_cubin"}, + "geometry": spec["geometry"], + "dtypes": spec["dtypes"], + "tokens": spec["tokens"], + "abi": { + "version": 1, + "function_prefix": compile_metadata["aot"]["function_prefix"], + "geometry_binding": "stable_project_c_wrapper", + "q_view": {"shape": ["T", 128, spec["geometry"]["h_q"]], "stride": [spec["geometry"]["h_q"] * 128, 1, 128]}, + "k_view": {"shape": [128, "T", spec["geometry"]["h_k"]], "stride": [1, spec["geometry"]["h_k"] * 128, 128]}, + "v_view": {"shape": [128, "T", spec["geometry"]["h_v"]], "stride": [1, spec["geometry"]["h_v"] * 128, 128]}, + "o_view": {"shape": [128, "T", spec["geometry"]["h_v"]], "stride": [1, spec["geometry"]["h_v"] * 128, 128]}, + "state_layout": "openinfer_hkv_v_contiguous", + }, + "workspace": dict(WORKSPACE), + "source": { + **source, + "generator_sha256": compile_metadata["generator_sha256"], + "requirements_lock_sha256": compile_metadata["requirements_lock_sha256"], + }, + "toolchain": compile_metadata["toolchain"], + "artifact": { + "format": "elf_relocatable_with_embedded_cubin", + "header": { + "sha256": sha256_bytes(header_bytes), + "size_bytes": len(header_bytes), + }, + "object": { + "sha256": sha256_bytes(object_bytes), + "size_bytes": len(object_bytes), + }, + "native_runtime": { + "sha256": sha256_bytes(runtime_bytes), + "size_bytes": len(runtime_bytes), + }, + }, + "distribution": { + "cute_runtime_linkage": "static", + }, + } + + +def package_candidate( + *, + raw_aot_dir: Path, + compile_metadata_path: Path, + output_dir: Path, + source: dict[str, Any], +) -> Path: + if output_dir.exists(): + raise ContractError(f"refusing to overwrite existing output directory: {output_dir}") + metadata = read_json(compile_metadata_path) + validate_compile_metadata(metadata, source) + + aot = metadata["aot"] + header_path = raw_aot_dir / aot["header"] + object_path = raw_aot_dir / aot["object"] + if not header_path.is_file() or not object_path.is_file(): + raise ContractError("AOT export header/object is missing") + header_bytes = header_path.read_bytes() + object_bytes = object_path.read_bytes() + runtime_path = Path(aot["native_runtime"]) + if not runtime_path.is_file(): + raise ContractError("CuTe static runtime archive is missing") + runtime_bytes = runtime_path.read_bytes() + _require_equal(aot["header_sha256"], sha256_bytes(header_bytes), "AOT header hash") + _require_equal(aot["header_size_bytes"], len(header_bytes), "AOT header size") + _require_equal(aot["object_sha256"], sha256_bytes(object_bytes), "AOT object hash") + _require_equal(aot["object_size_bytes"], len(object_bytes), "AOT object size") + _require_equal( + aot["native_runtime_sha256"], + sha256_bytes(runtime_bytes), + "CuTe static runtime hash", + ) + _require_equal( + aot["native_runtime_size_bytes"], + len(runtime_bytes), + "CuTe static runtime size", + ) + + output_dir.mkdir(parents=True) + header_name = ARTIFACT_FILES["header"] + object_name = ARTIFACT_FILES["object"] + runtime_name = ARTIFACT_FILES["native_runtime"] + (output_dir / header_name).write_bytes(header_bytes) + (output_dir / object_name).write_bytes(object_bytes) + (output_dir / runtime_name).write_bytes(runtime_bytes) + manifest = build_manifest( + variant=VARIANT, + header_bytes=header_bytes, + object_bytes=object_bytes, + runtime_bytes=runtime_bytes, + compile_metadata=metadata, + source=source, + ) + manifest_path = output_dir / "manifest.json" + write_json(manifest_path, manifest) + return manifest_path + + +def validate_manifest( + manifest_path: Path, + *, + flashinfer_dir: Path | None = None, + expected_variant: str | None = None, +) -> dict[str, Any]: + manifest = read_json(manifest_path) + _require_equal(manifest.get("schema_version"), SCHEMA_VERSION, "schema_version") + variant = expected_variant or manifest.get("variant") + if not isinstance(variant, str): + raise ContractError("manifest variant is missing") + spec = expected_spec(variant) + _require_equal(manifest.get("variant"), variant, "variant") + _require_equal(manifest.get("target"), {"arch": TARGET_ARCH, "code_object": "embedded_cubin"}, "target") + _require_equal(manifest.get("geometry"), spec["geometry"], "geometry") + _require_equal(manifest.get("dtypes"), spec["dtypes"], "dtypes") + _require_equal(manifest.get("tokens"), spec["tokens"], "dynamic token contract") + + source_manifest = manifest.get("source") + if not isinstance(source_manifest, dict): + raise ContractError("manifest source is missing") + _require_equal(source_manifest.get("flashinfer_commit"), FROZEN_FLASHINFER_COMMIT, "FlashInfer SHA") + lock, source_lock_sha256 = load_source_lock() + _require_equal( + source_manifest.get("source_lock_sha256"), + source_lock_sha256, + "source lock hash", + ) + _require_equal( + source_manifest.get("kernel_source_sha256"), + lock["patched_kernel_sha256"], + "patched kernel hash", + ) + _require_equal(source_manifest.get("generator_sha256"), sha256_file(compiler_path()), "generator hash") + _require_equal( + source_manifest.get("requirements_lock_sha256"), + sha256_file(requirements_lock_path()), + "requirements lock hash", + ) + + if flashinfer_dir is not None: + verify_flashinfer_base(flashinfer_dir) + workspace = manifest.get("workspace") + if not isinstance(workspace, dict): + raise ContractError("workspace is missing") + _require_equal(workspace, WORKSPACE, "workspace") + + artifact = manifest.get("artifact") + if not isinstance(artifact, dict): + raise ContractError("artifact metadata is missing") + _require_equal(artifact.get("format"), "elf_relocatable_with_embedded_cubin", "artifact format") + for component in ("header", "object", "native_runtime"): + entry = artifact.get(component) + if not isinstance(entry, dict): + raise ContractError(f"artifact {component} metadata is missing") + _require_equal( + set(entry), + {"sha256", "size_bytes"}, + f"artifact {component} metadata keys", + ) + name = ARTIFACT_FILES[component] + path = manifest_path.parent / name + if not path.is_file(): + raise ContractError(f"artifact {component} file is missing: {path}") + data = path.read_bytes() + _require_equal(entry.get("size_bytes"), len(data), f"artifact {component} size") + _require_equal(entry.get("sha256"), sha256_bytes(data), f"artifact {component} hash") + manifest_toolchain = manifest.get("toolchain") + if not isinstance(manifest_toolchain, dict): + raise ContractError("manifest toolchain is missing") + _require_equal(manifest_toolchain, PINNED_TOOLCHAIN, "manifest toolchain") + abi = manifest.get("abi") + if not isinstance(abi, dict): + raise ContractError("ABI metadata is missing") + _require_equal(abi.get("version"), 1, "stable C ABI version") + _require_equal(abi.get("function_prefix"), f"pegainfer_qwen35_gdn_{variant}", "AOT function prefix") + _require_equal( + abi.get("geometry_binding"), + "stable_project_c_wrapper", + "geometry binding", + ) + _require_equal( + abi.get("state_layout"), + "openinfer_hkv_v_contiguous", + "state layout", + ) + + distribution = manifest.get("distribution") + _require_equal( + distribution, + {"cute_runtime_linkage": "static"}, + "distribution metadata", + ) + return manifest + + +def default_flashinfer_dir() -> Path: + return Path(__file__).resolve().parents[2] / "third_party" / "flashinfer" + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + subparsers = parser.add_subparsers(dest="command", required=True) + candidate_parser = subparsers.add_parser("validate-candidate") + candidate_parser.add_argument("candidate", type=Path) + candidate_parser.add_argument("--flashinfer-dir", type=Path) + args = parser.parse_args() + + try: + manifest = args.candidate / "manifest.json" + validate_manifest(manifest, flashinfer_dir=args.flashinfer_dir) + print(f"validated {args.candidate}") + except ContractError as exc: + print(f"error: {exc}", file=sys.stderr) + return 2 + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/pegainfer-kernels/tools/flashinfer_gdn/compile_sm120.py b/pegainfer-kernels/tools/flashinfer_gdn/compile_sm120.py new file mode 100644 index 000000000..e6f1a4058 --- /dev/null +++ b/pegainfer-kernels/tools/flashinfer_gdn/compile_sm120.py @@ -0,0 +1,253 @@ +#!/usr/bin/env python3 +"""AOT-export the production FlashInfer GDN specialization to a C header/object.""" + +from __future__ import annotations + +import argparse +import importlib +import importlib.metadata +import json +import re +import sys +import types +from pathlib import Path + +from artifact_contract import ( + FORBIDDEN_TMA_CLUSTER_LOAD, + TARGET_ARCH, + VARIANT, + expected_spec, + compiler_path, + normalize_ptx, + requirements_lock_path, + sha256_file, + verify_prepared_flashinfer_source, + write_json, +) + + +def package_version(distribution: str) -> str: + try: + return importlib.metadata.version(distribution) + except importlib.metadata.PackageNotFoundError as exc: + raise RuntimeError(f"required generation package is missing: {distribution}") from exc + + +def ptx_metadata(ptx: str) -> dict[str, str]: + compiler_match = re.search( + r"Cuda compilation tools, release\s+([0-9.]+),\s+V([0-9.]+)", ptx + ) + isa_match = re.search(r"^\.version\s+([0-9.]+)$", ptx, re.MULTILINE) + if not compiler_match or not isa_match: + raise RuntimeError("cannot derive CUDA compiler/PTX ISA from generated PTX") + return { + "ptx_compiler_release": compiler_match.group(1), + "ptx_compiler_version": compiler_match.group(2), + "ptx_isa": isa_match.group(1), + } + + +def read_compiled_ptx(compiled: object) -> str: + artifact = getattr(compiled, "__ptx__", None) + if isinstance(artifact, str) and ".version" in artifact: + return artifact + if isinstance(artifact, str) and Path(artifact).is_file(): + return Path(artifact).read_text(encoding="utf-8") + raise RuntimeError("CuTe compile did not expose a readable PTX artifact") + + +def find_static_cuda_dialect_runtime() -> Path: + """Locate the runtime archive shipped by the pinned CuTe DSL wheel.""" + import cutlass + + cutlass_file = Path(cutlass.__file__).resolve() + roots = { + Path(entry).resolve() + for entry in sys.path + if entry and ("site-packages" in entry or "dist-packages" in entry) + } + # Stay inside the installed wheel/package tree. `Path.parents` eventually + # reaches `/`; recursively globbing that root made generation appear hung. + roots.update((cutlass_file.parent, cutlass_file.parent.parent)) + matches: list[Path] = [] + for root in roots: + if not root.is_dir(): + continue + matches.extend(root.glob("**/libcuda_dialect_runtime_static.a")) + unique = sorted({path.resolve() for path in matches if path.is_file()}) + if len(unique) != 1: + raise RuntimeError( + "expected exactly one libcuda_dialect_runtime_static.a in the pinned " + f"generation environment, found {[str(path) for path in unique]}" + ) + return unique[0] + + +def import_frozen_kernel(flashinfer_dir: Path): + """Import only the frozen kernel package, without FlashInfer's top-level API.""" + package_paths = { + "flashinfer": flashinfer_dir / "flashinfer", + "flashinfer.gdn_kernels": flashinfer_dir / "flashinfer" / "gdn_kernels", + "flashinfer.gdn_kernels.delta_rule_dsl": ( + flashinfer_dir / "flashinfer" / "gdn_kernels" / "delta_rule_dsl" + ), + } + for name, path in package_paths.items(): + package = types.ModuleType(name) + package.__path__ = [str(path)] + package.__package__ = name + sys.modules[name] = package + + # delta_rule_sm120 imports these helpers for its public Torch wrapper. The + # offline fake-tensor compiler never calls them, so avoid importing the + # rest of FlashInfer and its unrelated pynvml/JIT dependencies. + utils = types.ModuleType("flashinfer.utils") + + def generation_only_stub(*_args, **_kwargs): + raise RuntimeError("runtime-only FlashInfer helper called by offline compiler") + + utils.get_device_sm_count = generation_only_stub + utils._get_cache_buf = generation_only_stub + sys.modules["flashinfer.utils"] = utils + + cache_module = importlib.import_module( + "flashinfer.gdn_kernels.delta_rule_dsl.custom_compile_cache" + ) + kernel_module = importlib.import_module( + "flashinfer.gdn_kernels.delta_rule_dsl.delta_rule_sm120" + ) + return cache_module.cached_compile, kernel_module._FullyFusedDeltaRuleSm120 + + +def compile_variant(variant: str, flashinfer_dir: Path) -> tuple[object, str]: + spec = expected_spec(variant) + geometry = spec["geometry"] + import cutlass + import cutlass.cute as cute + + cached_compile, kernel_type = import_frozen_kernel(flashinfer_dir) + + h_q = geometry["h_q"] + h_k = geometry["h_k"] + h_v = geometry["h_v"] + d = geometry["head_dim"] + t = cute.sym_int() + flat_tokens = cute.sym_int() + workspace_bytes = cute.sym_int() + cu_count = cute.sym_int() + + q = cute.runtime.make_fake_tensor( + cutlass.BFloat16, (t, d, h_q), stride=(h_q * d, 1, d), assumed_align=16 + ) + k = cute.runtime.make_fake_tensor( + cutlass.BFloat16, (d, t, h_k), stride=(1, h_k * d, d), assumed_align=16 + ) + v = cute.runtime.make_fake_tensor( + cutlass.BFloat16, (d, t, h_v), stride=(1, h_v * d, d), assumed_align=16 + ) + o = cute.runtime.make_fake_tensor( + cutlass.BFloat16, (d, t, h_v), stride=(1, h_v * d, d), assumed_align=16 + ) + alpha = cute.runtime.make_fake_compact_tensor(cutlass.Float32, (flat_tokens,), assumed_align=16) + beta = cute.runtime.make_fake_compact_tensor(cutlass.Float32, (flat_tokens,), assumed_align=16) + state = cute.runtime.make_fake_compact_tensor(cutlass.Float32, (h_v * d * d,), assumed_align=16) + init_state = cute.runtime.make_fake_compact_tensor(cutlass.Float32, (h_v * d * d,), assumed_align=16) + workspace = cute.runtime.make_fake_compact_tensor(cutlass.Uint8, (workspace_bytes,), assumed_align=128) + cu_seqlens = cute.runtime.make_fake_compact_tensor(cutlass.Int64, (cu_count,), assumed_align=8) + stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True) + + kernel = kernel_type( + needs_alpha=True, + needs_beta=True, + needs_init_state=True, + needs_checkpointing=False, + dtype=cutlass.BFloat16, + ) + args = ( + q, + k, + v, + o, + alpha, + beta, + state, + init_state, + None, + None, + workspace, + cu_seqlens, + cutlass.Float32(1.0 / (d**0.5)), + cutlass.Int32(h_q), + cutlass.Int32(h_k), + cutlass.Int32(h_v), + cutlass.Int32(max(h_q, h_v)), + cutlass.Int32(1), + cutlass.Int32(1), + cutlass.Int32(0), + cutlass.Int32(max(h_q, h_v)), + stream, + ) + compiled = cached_compile(kernel, *args, compile_options=(cute.GPUArch(TARGET_ARCH),)) + ptx = normalize_ptx(read_compiled_ptx(compiled)) + if FORBIDDEN_TMA_CLUSTER_LOAD in ptx: + raise RuntimeError("upstream SM120 TMA workaround was not applied") + return compiled, ptx + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--variant", required=True, choices=(VARIANT,)) + parser.add_argument("--flashinfer-dir", required=True, type=Path) + parser.add_argument("--base-flashinfer-dir", required=True, type=Path) + parser.add_argument("--aot-out", required=True, type=Path) + parser.add_argument("--metadata-out", required=True, type=Path) + args = parser.parse_args() + + source = verify_prepared_flashinfer_source( + args.flashinfer_dir, args.base_flashinfer_dir + ) + compiled, ptx = compile_variant(args.variant, args.flashinfer_dir.resolve()) + prefix = f"pegainfer_qwen35_gdn_{args.variant}" + args.aot_out.mkdir(parents=True, exist_ok=True) + compiled.export_to_c(str(args.aot_out), prefix, prefix) + header = args.aot_out / f"{prefix}.h" + object_file = args.aot_out / f"{prefix}.o" + if not header.is_file() or not object_file.is_file(): + raise RuntimeError("CuTe export_to_c did not produce the expected .h/.o pair") + runtime_archive = find_static_cuda_dialect_runtime() + metadata = { + "flashinfer_commit": source["flashinfer_commit"], + "kernel_source_sha256": source["kernel_source_sha256"], + "source_lock_sha256": source["source_lock_sha256"], + "generator_sha256": sha256_file(compiler_path()), + "requirements_lock_sha256": sha256_file(requirements_lock_path()), + "toolchain": { + "python": sys.version.split()[0], + **ptx_metadata(ptx), + "cutlass_dsl": package_version("nvidia-cutlass-dsl"), + "cutlass_dsl_libs_base": package_version("nvidia-cutlass-dsl-libs-base"), + "torch": package_version("torch"), + "cuda_python": package_version("cuda-python"), + "cuda_bindings": package_version("cuda-bindings"), + "cuda_pathfinder": package_version("cuda-pathfinder"), + }, + "aot": { + "function_prefix": prefix, + "header": header.name, + "header_sha256": sha256_file(header), + "header_size_bytes": header.stat().st_size, + "object": object_file.name, + "object_sha256": sha256_file(object_file), + "object_size_bytes": object_file.stat().st_size, + "native_runtime": str(runtime_archive), + "native_runtime_sha256": sha256_file(runtime_archive), + "native_runtime_size_bytes": runtime_archive.stat().st_size, + }, + } + write_json(args.metadata_out, metadata) + print(json.dumps({"variant": args.variant, "aot": str(args.aot_out), "metadata": str(args.metadata_out)}, sort_keys=True)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/pegainfer-kernels/tools/flashinfer_gdn/generate.py b/pegainfer-kernels/tools/flashinfer_gdn/generate.py new file mode 100644 index 000000000..a709da722 --- /dev/null +++ b/pegainfer-kernels/tools/flashinfer_gdn/generate.py @@ -0,0 +1,81 @@ +#!/usr/bin/env python3 +"""Generate and package the production FlashInfer GDN SM120 AOT object.""" + +from __future__ import annotations + +import argparse +import json +import shutil +import subprocess +import sys +import tempfile +from pathlib import Path + +from artifact_contract import ( + VARIANT, + ContractError, + default_flashinfer_dir, + package_candidate, + prepare_flashinfer_source, + validate_manifest, +) + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--python", type=Path, default=Path(sys.executable)) + parser.add_argument("--flashinfer-dir", type=Path, default=default_flashinfer_dir()) + parser.add_argument("--output", type=Path, default=Path("target/flashinfer-gdn-sm120")) + args = parser.parse_args() + + output = args.output.resolve() + if output.exists(): + print(f"error: refusing to overwrite existing output directory: {output}", file=sys.stderr) + return 2 + try: + with tempfile.TemporaryDirectory(prefix="openinfer-gdn-sm120-") as temp_name: + temp = Path(temp_name) + prepared = temp / "patched-flashinfer" + source = prepare_flashinfer_source(args.flashinfer_dir, prepared) + staged = temp / "candidate" + compiler = Path(__file__).with_name("compile_sm120.py") + raw_dir = temp / "raw" + metadata_path = raw_dir / "compile-metadata.json" + subprocess.run( + [ + str(args.python), + str(compiler), + "--variant", + VARIANT, + "--flashinfer-dir", + str(prepared), + "--base-flashinfer-dir", + str(args.flashinfer_dir), + "--aot-out", + str(raw_dir), + "--metadata-out", + str(metadata_path), + ], + check=True, + ) + package_candidate( + raw_aot_dir=raw_dir, + compile_metadata_path=metadata_path, + output_dir=staged, + source=source, + ) + validate_manifest( + staged / "manifest.json", flashinfer_dir=args.flashinfer_dir + ) + output.parent.mkdir(parents=True, exist_ok=True) + shutil.move(str(staged), output) + except (ContractError, OSError, subprocess.CalledProcessError) as exc: + print(f"error: generation failed: {exc}", file=sys.stderr) + return 2 + + print(json.dumps({"candidate": str(output), "variant": VARIANT}, sort_keys=True)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/pegainfer-kernels/tools/flashinfer_gdn/patches/0001-openinfer-hkv-state-layout.patch b/pegainfer-kernels/tools/flashinfer_gdn/patches/0001-openinfer-hkv-state-layout.patch new file mode 100644 index 000000000..84e2f7a51 --- /dev/null +++ b/pegainfer-kernels/tools/flashinfer_gdn/patches/0001-openinfer-hkv-state-layout.patch @@ -0,0 +1,25 @@ +diff --git a/flashinfer/gdn_kernels/delta_rule_dsl/delta_rule_sm120.py b/flashinfer/gdn_kernels/delta_rule_dsl/delta_rule_sm120.py +--- a/flashinfer/gdn_kernels/delta_rule_dsl/delta_rule_sm120.py ++++ b/flashinfer/gdn_kernels/delta_rule_dsl/delta_rule_sm120.py +@@ -5,2 +5,3 @@ + import cutlass.cute as cute ++import cuda.bindings.driver as cuda + import cutlass.pipeline as pipeline +@@ -673,3 +674,3 @@ + (self.D, self.D, num_sab_heads, total_checkpoints), +- order=(0, 1, 2, 3), ++ order=(1, 0, 2, 3), # OpenInfer [H,K,V]: V is contiguous. + ) +@@ -1232,3 +1233,4 @@ + state_layout = cute.make_ordered_layout( +- (self.D, self.D, num_sab_heads, num_seqs), order=(0, 1, 2, 3) ++ (self.D, self.D, num_sab_heads, num_seqs), ++ order=(1, 0, 2, 3), # OpenInfer [H,K,V]: V is contiguous. + ) +@@ -1473,4 +1475,4 @@ + checkpoint_every_n_tokens: cutlass.Int32, +- grid_x: int, +- stream, ++ grid_x: cutlass.Int32, ++ stream: cuda.CUstream, + ): diff --git a/pegainfer-kernels/tools/flashinfer_gdn/requirements-cu13.lock b/pegainfer-kernels/tools/flashinfer_gdn/requirements-cu13.lock new file mode 100644 index 000000000..9f65a000d --- /dev/null +++ b/pegainfer-kernels/tools/flashinfer_gdn/requirements-cu13.lock @@ -0,0 +1,8 @@ +# Generation-only CUDA 13 environment. Serving statically links the exported +# object and libcuda_dialect_runtime_static.a; it never imports these packages. +nvidia-cutlass-dsl[cu13]==4.5.0 +nvidia-cutlass-dsl-libs-base==4.5.0 +cuda-python==13.0.1 +cuda-bindings==13.0.3 +cuda-pathfinder==1.6.0 +torch==2.7.1 diff --git a/pegainfer-kernels/tools/flashinfer_gdn/source-lock.json b/pegainfer-kernels/tools/flashinfer_gdn/source-lock.json new file mode 100644 index 000000000..6442cbab0 --- /dev/null +++ b/pegainfer-kernels/tools/flashinfer_gdn/source-lock.json @@ -0,0 +1,11 @@ +{ + "schema_version": 3, + "flashinfer_commit": "a0efa0adfe49bb836ab1a147d6572980b870f3d4", + "patches": [ + { + "path": "patches/0001-openinfer-hkv-state-layout.patch", + "sha256": "76aff3ef6d5fc1ecb9895640ee88d29d63ebbeecdabd622264e638335d3c6f22" + } + ], + "patched_kernel_sha256": "4e3c6f81edf39b5444f20353b1307c8028b2496d702b7bfb9ebfbcebf4f7b35b" +} diff --git a/pegainfer-qwen35/Cargo.toml b/pegainfer-qwen35/Cargo.toml index f308934bf..25726284c 100644 --- a/pegainfer-qwen35/Cargo.toml +++ b/pegainfer-qwen35/Cargo.toml @@ -33,6 +33,7 @@ vllm-text = { workspace = true } [features] default = [] +gdn-validation = ["qwen35"] qwen35 = ["pegainfer-kernels/qwen35"] [lints] @@ -50,10 +51,6 @@ required-features = ["qwen35"] name = "sampling_behavior" required-features = ["qwen35"] -[[test]] -name = "chunked_prefill" -required-features = ["qwen35"] - [[test]] name = "serving_tp2" required-features = ["qwen35"] diff --git a/pegainfer-qwen35/src/batch_decode.rs b/pegainfer-qwen35/src/batch_decode.rs index 8d1c373c6..23ab24f3b 100644 --- a/pegainfer-qwen35/src/batch_decode.rs +++ b/pegainfer-qwen35/src/batch_decode.rs @@ -329,6 +329,8 @@ impl Qwen35Model { ); if !self.config.decode_group_is_compiled() { + #[cfg(feature = "gdn-validation")] + graph_state.evidence.record_graph_eager_fallback(); LOG_UNCOMPILED_DECODE_ROUTE.call_once(|| { let group = self.config.num_attention_heads / self.config.num_key_value_heads; log::info!( @@ -389,6 +391,8 @@ impl Qwen35Model { let mut graphs = std::mem::take(&mut graph_state.graphs); let linear_state_ptrs = &graph_state.linear_pointer_tables.state_ptrs; let linear_conv_state_ptrs = &graph_state.linear_pointer_tables.conv_state_ptrs; + #[cfg(feature = "gdn-validation")] + let was_captured = graphs[bucket_idx].is_captured(); let result = graphs[bucket_idx].run_or_capture(&self.ctx, || { self.batch_decode_kernels_graph( kv_buffer, @@ -399,6 +403,14 @@ impl Qwen35Model { &mut graph_state.buffers, ) }); + #[cfg(feature = "gdn-validation")] + if result.is_ok() { + if was_captured { + graph_state.evidence.record_graph_replay(); + } else { + graph_state.evidence.record_graph_capture(); + } + } graph_state.graphs = graphs; result } diff --git a/pegainfer-qwen35/src/batch_decode_graph.rs b/pegainfer-qwen35/src/batch_decode_graph.rs index eb88ba691..0c7be347c 100644 --- a/pegainfer-qwen35/src/batch_decode_graph.rs +++ b/pegainfer-qwen35/src/batch_decode_graph.rs @@ -8,6 +8,8 @@ use pegainfer_core::tensor::DeviceContext; use super::config::Config35; use super::config::LocalGeometry; use super::decode_buffers::BatchDecodeBuffers35; +#[cfg(feature = "gdn-validation")] +use super::gdn_validation::GdnValidationEvidenceHandle; use super::recurrent_state::LinearStatePointerTables; use super::recurrent_state::RecurrentState; @@ -54,6 +56,8 @@ pub(crate) struct BatchDecodeGraphState { pub(crate) linear_pointer_tables: LinearStatePointerTables, /// One `CudaGraphState` per BATCH_BUCKETS entry (indexed by position). pub(crate) graphs: Vec, + #[cfg(feature = "gdn-validation")] + pub(crate) evidence: GdnValidationEvidenceHandle, } impl BatchDecodeGraphState { @@ -102,9 +106,20 @@ impl BatchDecodeGraphState { slot_states, linear_pointer_tables, graphs, + #[cfg(feature = "gdn-validation")] + evidence: Default::default(), }) } + #[cfg(feature = "gdn-validation")] + pub(crate) fn with_validation_evidence( + mut self, + evidence: GdnValidationEvidenceHandle, + ) -> Self { + self.evidence = evidence; + self + } + /// D2D copy `src` recurrent state into slot `slot_idx`. /// /// Call once when a request joins the batch (after prefill finishes). @@ -117,6 +132,8 @@ impl BatchDecodeGraphState { slot_idx: usize, ) -> Result<()> { let dst = &mut self.slot_states[slot_idx]; + #[cfg(feature = "gdn-validation")] + let reused = dst.seq_len != 0; for (dst_layer, src_layer) in dst.layers.iter_mut().zip(src.layers.iter()) { ctx.stream .memcpy_dtod(&src_layer.state, &mut dst_layer.state) @@ -126,6 +143,13 @@ impl BatchDecodeGraphState { .map_err(|e| anyhow::anyhow!("copy conv state to slot {slot_idx}: {e}"))?; } dst.seq_len = src.seq_len; + #[cfg(feature = "gdn-validation")] + self.evidence.record_state_slot_copy(reused); Ok(()) } + + #[cfg(feature = "gdn-validation")] + pub(crate) fn record_slot_compaction(&self) { + self.evidence.record_slot_compaction(); + } } diff --git a/pegainfer-qwen35/src/executor.rs b/pegainfer-qwen35/src/executor.rs index 60e021330..a084db619 100644 --- a/pegainfer-qwen35/src/executor.rs +++ b/pegainfer-qwen35/src/executor.rs @@ -14,6 +14,8 @@ use pegainfer_frontend::sampler::SamplingParams; use crate::batch_decode_graph::BatchDecodeGraphState; use crate::decode_buffers::BatchDecodeBuffers35; +#[cfg(feature = "gdn-validation")] +use crate::gdn_validation::GdnPrefillRuntimeEvidence; use crate::logprobs::snapshot_requested_logprobs; use crate::recurrent_state::RecurrentState; use crate::weights::Qwen35Model; @@ -123,6 +125,13 @@ impl Qwen35Executor { }) } + /// Return production backend identity and launch proof when Auto selected + /// the build-linked FlashInfer specialization. + #[cfg(feature = "gdn-validation")] + pub fn flashinfer_gdn_runtime_evidence(&self) -> Result { + self.model.flashinfer_gdn_runtime_evidence() + } + pub fn execute_prefill(&mut self, plan: PrefillPlan<'_>) -> Result { anyhow::ensure!( !plan.requests.is_empty(), @@ -309,6 +318,8 @@ impl Qwen35Executor { })?; } self.graph_state.slot_states[idx].seq_len = self.graph_state.slot_states[last].seq_len; + #[cfg(feature = "gdn-validation")] + self.graph_state.record_slot_compaction(); self.active[idx].graph_slot_idx = idx; } Ok(()) diff --git a/pegainfer-qwen35/src/flashinfer_gdn.rs b/pegainfer-qwen35/src/flashinfer_gdn.rs new file mode 100644 index 000000000..4d2e422eb --- /dev/null +++ b/pegainfer-qwen35/src/flashinfer_gdn.rs @@ -0,0 +1,109 @@ +//! Model-side semantic boundary for Qwen3.5 GDN prefill. +//! +//! CuTe/generated-symbol/TMA/module/workspace details belong exclusively to +//! `pegainfer-kernels`. This module owns only model policy, prepared tensors, +//! recurrent state, and observable backend evidence. + +use anyhow::Context; +use anyhow::Result; +use anyhow::ensure; +use cudarc::driver::CudaSlice; +use pegainfer_core::tensor::DeviceContext; +use pegainfer_core::tensor::HiddenStates; +use pegainfer_kernels::ops::Qwen35GdnAot; +use pegainfer_kernels::ops::Qwen35GdnGeometry; +use pegainfer_kernels::ops::Qwen35GdnWorkspace; + +use crate::config::Config35; +use crate::prefill_buffers::GdnPrepareScratch35; +use crate::weights::Qwen35Model; + +/// Backend selected once at the production prefill boundary. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) enum GdnPrefillBackend { + Triton, + FlashInfer, +} + +pub(crate) struct FlashInferGdnChunkResources { + pub(crate) prepare: GdnPrepareScratch35, + pub(crate) output: HiddenStates, + launch: Qwen35GdnWorkspace, +} + +impl FlashInferGdnChunkResources { + pub(crate) fn new( + ctx: &DeviceContext, + config: &Config35, + backend: &Qwen35GdnAot, + tokens: usize, + ) -> Result { + let geometry = model_geometry(config); + ensure!( + geometry == Qwen35GdnGeometry::PRODUCTION, + "FlashInfer GDN is not supported for model geometry {geometry:?}" + ); + Ok(Self { + prepare: GdnPrepareScratch35::new(ctx, config, tokens)?, + output: HiddenStates::zeros(ctx, geometry.h_v * geometry.head_dim, tokens)?, + launch: backend.allocate_workspace(ctx, tokens)?, + }) + } + + pub(crate) fn ensure_prepare_inputs_finite(&self, ctx: &DeviceContext) -> Result<()> { + let status = ctx + .stream + .clone_dtoh(&self.prepare.non_finite_status) + .map_err(|error| anyhow::anyhow!("read native GDN finite-status failed: {error}"))?; + ctx.sync()?; + ensure!( + status == [0], + "native GDN prepare rejected non-finite qkv/gate input" + ); + Ok(()) + } + + pub(crate) fn launch_in_place( + &mut self, + ctx: &DeviceContext, + backend: &Qwen35GdnAot, + state: &mut CudaSlice, + ) -> Result<()> { + backend.launch_in_place( + ctx, + &self.prepare.q, + &self.prepare.k, + &self.prepare.v, + &self.prepare.alpha, + &self.prepare.beta, + state, + &mut self.output, + &mut self.launch, + ) + } +} + +pub(crate) fn model_geometry(config: &Config35) -> Qwen35GdnGeometry { + Qwen35GdnGeometry { + h_q: config.linear_num_key_heads, + h_k: config.linear_num_key_heads, + h_v: config.linear_num_value_heads, + head_dim: config.linear_key_head_dim, + } +} + +impl Qwen35Model { + pub(crate) fn resolved_gdn_backend(&self) -> GdnPrefillBackend { + if self.flashinfer_gdn.is_some() { + GdnPrefillBackend::FlashInfer + } else { + GdnPrefillBackend::Triton + } + } + + pub(super) fn flashinfer_gdn(&self) -> Result<&Qwen35GdnAot> { + self.flashinfer_gdn + .as_ref() + .context("FlashInfer GDN is not selected for this model capability") + } +} diff --git a/pegainfer-qwen35/src/gdn_validation.rs b/pegainfer-qwen35/src/gdn_validation.rs new file mode 100644 index 000000000..01028f7c8 --- /dev/null +++ b/pegainfer-qwen35/src/gdn_validation.rs @@ -0,0 +1,121 @@ +#![cfg(feature = "gdn-validation")] + +//! Non-default runtime evidence for the required SM120 GDN validation gates. + +use std::sync::Arc; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; + +use anyhow::Result; + +use crate::weights::Qwen35Model; + +#[derive(Debug, Default)] +struct GdnValidationEvidenceCounters { + successful_launches: AtomicU64, + graph_captures: AtomicU64, + graph_replays: AtomicU64, + graph_eager_fallbacks: AtomicU64, + state_slot_copies: AtomicU64, + state_slot_reuses: AtomicU64, + slot_compactions: AtomicU64, +} + +#[derive(Clone, Debug, Default)] +pub(crate) struct GdnValidationEvidenceHandle { + counters: Arc, +} + +impl GdnValidationEvidenceHandle { + pub(crate) fn record_successful_launch(&self) { + self.counters + .successful_launches + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn record_graph_capture(&self) { + self.counters.graph_captures.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn record_graph_replay(&self) { + self.counters.graph_replays.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn record_graph_eager_fallback(&self) { + self.counters + .graph_eager_fallbacks + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn record_state_slot_copy(&self, reused: bool) { + self.counters + .state_slot_copies + .fetch_add(1, Ordering::Relaxed); + if reused { + self.counters + .state_slot_reuses + .fetch_add(1, Ordering::Relaxed); + } + } + + pub(crate) fn record_slot_compaction(&self) { + self.counters + .slot_compactions + .fetch_add(1, Ordering::Relaxed); + } +} + +/// Runtime proof for production dispatch and same-path validation gates. +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct GdnPrefillRuntimeEvidence { + pub selected_backend: String, + pub artifact_sha256: String, + pub successful_launches: u64, + pub graph_captures: u64, + pub graph_replays: u64, + pub graph_eager_fallbacks: u64, + pub state_slot_copies: u64, + pub state_slot_reuses: u64, + pub slot_compactions: u64, +} + +#[derive(Clone, Debug)] +pub struct GdnPrefillRuntimeEvidenceHandle { + selected_backend: &'static str, + artifact_sha256: String, + validation: GdnValidationEvidenceHandle, +} + +impl GdnPrefillRuntimeEvidenceHandle { + pub fn snapshot(&self) -> GdnPrefillRuntimeEvidence { + let counters = &self.validation.counters; + GdnPrefillRuntimeEvidence { + selected_backend: self.selected_backend.to_owned(), + artifact_sha256: self.artifact_sha256.clone(), + successful_launches: counters.successful_launches.load(Ordering::Relaxed), + graph_captures: counters.graph_captures.load(Ordering::Relaxed), + graph_replays: counters.graph_replays.load(Ordering::Relaxed), + graph_eager_fallbacks: counters.graph_eager_fallbacks.load(Ordering::Relaxed), + state_slot_copies: counters.state_slot_copies.load(Ordering::Relaxed), + state_slot_reuses: counters.state_slot_reuses.load(Ordering::Relaxed), + slot_compactions: counters.slot_compactions.load(Ordering::Relaxed), + } + } +} + +impl Qwen35Model { + pub fn flashinfer_gdn_runtime_evidence(&self) -> Result { + Ok(self.flashinfer_gdn_runtime_evidence_handle()?.snapshot()) + } + + pub fn flashinfer_gdn_runtime_evidence_handle( + &self, + ) -> Result { + let backend = self.flashinfer_gdn()?; + Ok(GdnPrefillRuntimeEvidenceHandle { + selected_backend: "flashinfer", + artifact_sha256: backend.artifact_sha256().to_owned(), + validation: self.gdn_validation_evidence.clone(), + }) + } +} diff --git a/pegainfer-qwen35/src/lib.rs b/pegainfer-qwen35/src/lib.rs index 0638d222c..308113fa7 100644 --- a/pegainfer-qwen35/src/lib.rs +++ b/pegainfer-qwen35/src/lib.rs @@ -10,6 +10,9 @@ pub(crate) mod config; mod decode_buffers; mod executor; mod ffi; +mod flashinfer_gdn; +#[cfg(feature = "gdn-validation")] +mod gdn_validation; mod logprobs; pub mod model_line; mod ops; @@ -63,7 +66,13 @@ pub mod runtime { pub use crate::executor::PrefillStepItem; pub use crate::executor::Qwen35Executor; pub use crate::executor::RequestId; + #[cfg(feature = "gdn-validation")] + pub use crate::gdn_validation::GdnPrefillRuntimeEvidence; + #[cfg(feature = "gdn-validation")] + pub use crate::gdn_validation::GdnPrefillRuntimeEvidenceHandle; pub use crate::scheduler::start_with_capacity; + #[cfg(feature = "gdn-validation")] + pub use crate::start_engine_with_flashinfer_gdn_for_accuracy; pub use crate::tp_executor::DropExpectation; pub use crate::tp_executor::Qwen35TpExecutor; pub use crate::weights::Qwen35Model; @@ -111,6 +120,31 @@ pub fn start_engine( ) } +/// Start the normal single-GPU production scheduler and expose build-linked +/// FlashInfer launch evidence to end-to-end accuracy tests. +#[cfg(feature = "gdn-validation")] +pub fn start_engine_with_flashinfer_gdn_for_accuracy( + model_path: &Path, + device_ordinal: usize, + max_batch: usize, + max_prefill_tokens: usize, +) -> Result<( + EngineHandle, + gdn_validation::GdnPrefillRuntimeEvidenceHandle, +)> { + anyhow::ensure!( + (1..=MAX_DECODE_BATCH).contains(&max_batch), + "Qwen3.5 max_batch must be in 1..={MAX_DECODE_BATCH}, got {max_batch}" + ); + let model_path = model_path + .to_str() + .ok_or_else(|| anyhow!("model path must be valid UTF-8"))?; + let model = weights::Qwen35Model::from_safetensors(model_path, device_ordinal, max_batch)?; + let evidence = model.flashinfer_gdn_runtime_evidence_handle()?; + let handle = scheduler::start_with_capacity(model, 42, max_batch, max_prefill_tokens)?; + Ok((handle, evidence)) +} + #[derive(Clone, Debug)] pub struct Qwen35LaunchOptions { /// CUDA device for single-GPU loads (ignored when `tp_size > 1`). diff --git a/pegainfer-qwen35/src/ops.rs b/pegainfer-qwen35/src/ops.rs index 3d288b074..fbb9fae13 100644 --- a/pegainfer-qwen35/src/ops.rs +++ b/pegainfer-qwen35/src/ops.rs @@ -22,5 +22,6 @@ pub(crate) use recurrent::conv1d_decode_batch_into; pub(crate) use recurrent::conv1d_prefill_batch_into; pub(crate) use recurrent::gated_delta_rule_decode_batch_into; pub use recurrent::gated_delta_rule_prefill_chunkwise_into; +pub(crate) use recurrent::gated_delta_rule_prefill_native_prepare_into; use crate::recurrent; diff --git a/pegainfer-qwen35/src/prefill.rs b/pegainfer-qwen35/src/prefill.rs index 6b65d9459..16f148d79 100644 --- a/pegainfer-qwen35/src/prefill.rs +++ b/pegainfer-qwen35/src/prefill.rs @@ -24,6 +24,8 @@ use pegainfer_core::kv_pool::KvState; use pegainfer_core::tensor::DeviceVec; use pegainfer_core::tensor::HiddenStates; +use super::flashinfer_gdn::FlashInferGdnChunkResources; +use super::flashinfer_gdn::GdnPrefillBackend; use super::prefill_buffers::GdrChunkwiseScratch35; use super::recurrent_state::RecurrentState; use super::weights::FullAttentionLayer; @@ -35,6 +37,11 @@ use crate::ffi; use crate::ops; use crate::ops::PrefillPagedPlan; +enum GdnPrefillChunkScratch { + Triton(Box), + FlashInfer(Box), +} + fn checked_prefill_end_pos( base_pos: usize, seq_len: usize, @@ -77,11 +84,13 @@ impl Qwen35Model { // per-pass GDR scratch (which grows with the pass length) at the budget // reserved at startup, so prompts longer than one chunk prefill without OOM. let mut hidden_batch: Option = None; + let gdn_backend = self.resolved_gdn_backend(); for chunk in token_ids.chunks(PREFILL_CHUNK_LEN) { // Free the previous chunk's hidden states before allocating the next // chunk's scratch so peak memory stays within one chunk's reservation. drop(hidden_batch.take()); - hidden_batch = Some(self.prefill_chunk_forward(chunk, kv_state, recurrent)?); + hidden_batch = + Some(self.prefill_chunk_forward(chunk, kv_state, recurrent, gdn_backend)?); } // `seq_len > 0` guarantees at least one chunk produced hidden states. let hidden_batch = hidden_batch.expect("prefill produced no chunk despite seq_len > 0"); @@ -140,9 +149,10 @@ impl Qwen35Model { token_ids: &[u32], kv_state: &mut KvState, recurrent: &mut RecurrentState, + gdn_backend: GdnPrefillBackend, ) -> Result { let seq_len = token_ids.len(); - debug_assert!( + anyhow::ensure!( seq_len > 0 && seq_len <= PREFILL_CHUNK_LEN, "prefill chunk length {seq_len} out of range 1..={PREFILL_CHUNK_LEN}" ); @@ -170,7 +180,20 @@ impl Qwen35Model { // Allocate the chunk scratch before advancing the KV state. It is the // largest, most allocation-prone buffer here, so failing first leaves // `kv_state` untouched and the request can be rejected cleanly. - let mut gdr_chunkwise_scratch = GdrChunkwiseScratch35::new(&self.ctx, c, seq_len)?; + let mut gdn_scratch = match gdn_backend { + GdnPrefillBackend::Triton => GdnPrefillChunkScratch::Triton(Box::new( + GdrChunkwiseScratch35::new(&self.ctx, c, seq_len)?, + )), + GdnPrefillBackend::FlashInfer => { + let backend = self.flashinfer_gdn()?; + GdnPrefillChunkScratch::FlashInfer(Box::new(FlashInferGdnChunkResources::new( + &self.ctx, + &self.config, + backend, + seq_len, + )?)) + } + }; // Advance paged KV state and build this chunk's prefill plan. kv_state.ensure_capacity(end_pos)?; @@ -196,7 +219,7 @@ impl Qwen35Model { layer_idx, layer, &hidden_batch, - &mut gdr_chunkwise_scratch, + &mut gdn_scratch, &mut linear_idx, &mut full_idx, kv_state, @@ -205,6 +228,10 @@ impl Qwen35Model { )?; } + if let GdnPrefillChunkScratch::FlashInfer(resources) = &gdn_scratch { + resources.ensure_prepare_inputs_finite(&self.ctx)?; + } + // Advance recurrent token count for the next chunk / decode step; the // paged KV position is tracked by `kv_state` (advanced above). recurrent.seq_len += seq_len; @@ -219,7 +246,7 @@ impl Qwen35Model { _layer_idx: usize, layer: &TransformerBlock35, hidden_batch: &HiddenStates, - gdr_chunkwise_scratch: &mut GdrChunkwiseScratch35, + gdn_scratch: &mut GdnPrefillChunkScratch, linear_idx: &mut usize, full_idx: &mut usize, kv_state: &KvState, @@ -259,7 +286,7 @@ impl Qwen35Model { &normed_batch, linear_idx, recurrent, - gdr_chunkwise_scratch, + gdn_scratch, seq_len, )?, }; @@ -441,7 +468,7 @@ impl Qwen35Model { normed_batch: &HiddenStates, linear_idx: &mut usize, recurrent: &mut RecurrentState, - gdr_chunkwise_scratch: &mut GdrChunkwiseScratch35, + gdn_scratch: &mut GdnPrefillChunkScratch, seq_len: usize, ) -> Result { let c = &self.config; @@ -466,34 +493,65 @@ impl Qwen35Model { c.linear_conv_kernel_dim, ); - let mut gdr_out_batch = HiddenStates::zeros(&self.ctx, z_dim, seq_len)?; - ops::gated_delta_rule_prefill_chunkwise_into( - &self.ctx, - &qkv_conv_batch, - &b_batch, - &a_batch, - &attn.dt_bias, - &attn.a_log, - &mut layer_state.state, - gdr_chunkwise_scratch, - &mut gdr_out_batch, - c.linear_num_key_heads, - c.linear_num_value_heads, - c.linear_key_head_dim, - c.linear_value_head_dim, - )?; - let mut normed_out_batch = HiddenStates::zeros(&self.ctx, z_dim, seq_len)?; - ops::rms_norm_gated_batch_into( - &self.ctx, - &gdr_out_batch, - &attn.norm_weight, - &z_batch, - &mut normed_out_batch, - c.linear_num_value_heads, - c.linear_value_head_dim, - c.rms_norm_eps, - ); + match gdn_scratch { + GdnPrefillChunkScratch::Triton(scratch) => { + let mut gdr_out_batch = HiddenStates::zeros(&self.ctx, z_dim, seq_len)?; + ops::gated_delta_rule_prefill_chunkwise_into( + &self.ctx, + &qkv_conv_batch, + &b_batch, + &a_batch, + &attn.dt_bias, + &attn.a_log, + &mut layer_state.state, + scratch, + &mut gdr_out_batch, + c.linear_num_key_heads, + c.linear_num_value_heads, + c.linear_key_head_dim, + c.linear_value_head_dim, + )?; + ops::rms_norm_gated_batch_into( + &self.ctx, + &gdr_out_batch, + &attn.norm_weight, + &z_batch, + &mut normed_out_batch, + c.linear_num_value_heads, + c.linear_value_head_dim, + c.rms_norm_eps, + ); + } + GdnPrefillChunkScratch::FlashInfer(resources) => { + ops::gated_delta_rule_prefill_native_prepare_into( + &self.ctx, + &qkv_conv_batch, + &b_batch, + &a_batch, + &attn.dt_bias, + &attn.a_log, + &mut resources.prepare, + )?; + resources.launch_in_place( + &self.ctx, + self.flashinfer_gdn()?, + &mut layer_state.state, + )?; + #[cfg(feature = "gdn-validation")] + self.gdn_validation_evidence.record_successful_launch(); + ops::rms_norm_gated_batch_into( + &self.ctx, + &resources.output, + &attn.norm_weight, + &z_batch, + &mut normed_out_batch, + c.linear_num_value_heads, + c.linear_value_head_dim, + c.rms_norm_eps, + ); + } + } *linear_idx += 1; @@ -514,35 +572,4 @@ impl Qwen35Model { } #[cfg(test)] -mod tests { - use super::checked_prefill_end_pos; - - #[test] - fn checked_prefill_end_pos_accepts_config_limit() { - assert_eq!( - checked_prefill_end_pos(0, 262_144, 262_144).unwrap(), - 262_144 - ); - assert_eq!( - checked_prefill_end_pos(262_143, 1, 262_144).unwrap(), - 262_144 - ); - } - - #[test] - fn checked_prefill_end_pos_rejects_past_config_limit() { - let err = checked_prefill_end_pos(0, 262_145, 262_144) - .unwrap_err() - .to_string(); - assert!(err.contains("beyond max_position_embeddings=262144")); - assert!(err.contains("requested end_pos=262145")); - } - - #[test] - fn checked_prefill_end_pos_rejects_overflow() { - let err = checked_prefill_end_pos(usize::MAX, 1, 262_144) - .unwrap_err() - .to_string(); - assert!(err.contains("prefill position overflow")); - } -} +mod tests; diff --git a/pegainfer-qwen35/src/prefill/tests.rs b/pegainfer-qwen35/src/prefill/tests.rs new file mode 100644 index 000000000..7da0c7644 --- /dev/null +++ b/pegainfer-qwen35/src/prefill/tests.rs @@ -0,0 +1,294 @@ +#[cfg(feature = "gdn-validation")] +use std::path::Path; + +#[cfg(feature = "gdn-validation")] +use anyhow::Result; + +#[cfg(feature = "gdn-validation")] +use super::GdnPrefillBackend; +use super::checked_prefill_end_pos; +#[cfg(feature = "gdn-validation")] +use crate::recurrent_state::RecurrentState; +#[cfg(feature = "gdn-validation")] +use crate::weights::Qwen35Model; + +// Stage 18 measured the real-model FP32 recurrent-state partition floor at +// mean=1.3146e-4 and p99=1.0670e-3. Stage 19 then calibrated the BF16 conv-state +// distribution on the real zero-state SM120 continuation gate. These bounds +// retain margin over those observed floors without turning token parity into +// the only continuation criterion. +#[cfg(feature = "gdn-validation")] +const RECURRENT_STATE_MEAN_TOL: f32 = 2.5e-4; +#[cfg(feature = "gdn-validation")] +const RECURRENT_STATE_P99_TOL: f32 = 2.0e-3; +#[cfg(feature = "gdn-validation")] +const CONV_STATE_MEAN_TOL: f32 = 1.5625e-2; +#[cfg(feature = "gdn-validation")] +const CONV_STATE_P99_TOL: f32 = 6.25e-2; +#[cfg(feature = "gdn-validation")] +const LOGIT_MEAN_TOL: f32 = 0.06; +#[cfg(feature = "gdn-validation")] +const LOGIT_P99_TOL: f32 = 0.20; +#[cfg(feature = "gdn-validation")] +const LOGIT_ARGMAX_REGRET_TOL: f32 = 0.20; + +#[cfg(feature = "gdn-validation")] +fn required_model_path() -> String { + let default = concat!(env!("CARGO_MANIFEST_DIR"), "/../models/Qwen3.5-4B"); + let path = std::env::var("PEGAINFER_TEST_MODEL_PATH").unwrap_or_else(|_| default.to_string()); + assert!( + Path::new(&path).join("config.json").is_file(), + "required chunk-continuation gate cannot read {path}/config.json; set PEGAINFER_TEST_MODEL_PATH" + ); + path +} + +#[cfg(feature = "gdn-validation")] +fn assert_distribution_close( + label: &str, + expected: &[f32], + actual: &[f32], + mean_tolerance: f32, + p99_tolerance: f32, +) { + assert_eq!(expected.len(), actual.len(), "{label} length mismatch"); + assert!(!expected.is_empty(), "{label} must not be empty"); + + let mut deltas = Vec::with_capacity(expected.len()); + for (index, (&left, &right)) in expected.iter().zip(actual).enumerate() { + assert!( + left.is_finite() && right.is_finite(), + "{label} contains a non-finite value at {index}: expected={left} actual={right}" + ); + deltas.push((left - right).abs()); + } + deltas.sort_by(f32::total_cmp); + + let mean = + (deltas.iter().map(|&value| f64::from(value)).sum::() / deltas.len() as f64) as f32; + let p50 = deltas[deltas.len().saturating_sub(1) * 50 / 100]; + let p99 = deltas[deltas.len().saturating_sub(1) * 99 / 100]; + let max = *deltas.last().expect("non-empty deltas"); + eprintln!( + "{label}: elements={} mean_abs={mean:.8} p50_abs={p50:.8} p99_abs={p99:.8} max_abs={max:.8} mean_tol={mean_tolerance} p99_tol={p99_tolerance}", + deltas.len() + ); + + assert!( + mean <= mean_tolerance, + "{label} mean_abs {mean} exceeds {mean_tolerance}" + ); + assert!( + p99 <= p99_tolerance, + "{label} p99_abs {p99} exceeds {p99_tolerance}" + ); +} + +#[cfg(feature = "gdn-validation")] +fn assert_recurrent_continuation( + model: &Qwen35Model, + unchunked: &RecurrentState, + chunked: &RecurrentState, +) -> Result<()> { + assert_eq!(unchunked.seq_len, 128); + assert_eq!(chunked.seq_len, 128); + assert_eq!( + unchunked.layers.len(), + chunked.layers.len(), + "linear recurrent layer count mismatch" + ); + + let ctx = model.device_ctx(); + for (layer, (expected, actual)) in unchunked.layers.iter().zip(&chunked.layers).enumerate() { + let expected_state = ctx.stream.clone_dtoh(&expected.state)?; + let actual_state = ctx.stream.clone_dtoh(&actual.state)?; + let expected_conv = expected.conv_state.to_host(ctx)?; + let actual_conv = actual.conv_state.to_host(ctx)?; + ctx.sync()?; + + assert_distribution_close( + &format!("real-model layer {layer} recurrent state"), + &expected_state, + &actual_state, + RECURRENT_STATE_MEAN_TOL, + RECURRENT_STATE_P99_TOL, + ); + assert_distribution_close( + &format!("real-model layer {layer} conv state"), + &expected_conv, + &actual_conv, + CONV_STATE_MEAN_TOL, + CONV_STATE_P99_TOL, + ); + } + Ok(()) +} + +#[cfg(feature = "gdn-validation")] +fn assert_logits_close(label: &str, expected: &[f32], actual: &[f32]) -> u32 { + assert_distribution_close(label, expected, actual, LOGIT_MEAN_TOL, LOGIT_P99_TOL); + + let expected_top = pegainfer_sample::token_logprob_from_row(expected, 0, 1) + .and_then(|summary| summary.top_logprobs.into_iter().next()) + .expect("baseline logits must contain a top token"); + let actual_top = pegainfer_sample::token_logprob_from_row(actual, 0, 1) + .and_then(|summary| summary.top_logprobs.into_iter().next()) + .expect("candidate logits must contain a top token"); + let actual_token_in_baseline = + pegainfer_sample::token_logprob_from_row(expected, actual_top.0, 0) + .expect("candidate token must be in the baseline vocabulary"); + let regret = expected_top.1 - actual_token_in_baseline.logprob; + + eprintln!( + "{label}: expected_token={} actual_token={} expected_logprob={:.6} actual_logprob={:.6} regret={regret:.6}", + expected_top.0, actual_top.0, expected_top.1, actual_top.1 + ); + assert!( + regret <= LOGIT_ARGMAX_REGRET_TOL, + "{label} candidate token {} has baseline regret {regret} > {LOGIT_ARGMAX_REGRET_TOL}", + actual_top.0 + ); + assert_eq!( + actual_top.0, expected_top.0, + "{label} greedy token parity failed" + ); + expected_top.0 +} + +#[cfg(feature = "gdn-validation")] +fn last_token_logits( + model: &Qwen35Model, + hidden: &pegainfer_core::tensor::HiddenStates, +) -> Result> { + let last = crate::ops::extract_vec(model.device_ctx(), hidden, hidden.seq_len - 1)?; + model + .batch_last_hidden_logits(&[last])? + .to_host(model.device_ctx()) +} + +#[cfg(feature = "gdn-validation")] +fn run_prefill_case( + model: &Qwen35Model, + tokens: &[u32], + backend: GdnPrefillBackend, + split_at: Option, +) -> Result<(pegainfer_core::kv_pool::KvState, RecurrentState, Vec)> { + let mut kv = model.alloc_kv(); + let mut recurrent = RecurrentState::new(model.device_ctx(), model.config())?; + let hidden = match split_at { + Some(split) => { + assert!(split > 0 && split < tokens.len()); + drop(model.prefill_chunk_forward( + &tokens[..split], + &mut kv, + &mut recurrent, + backend, + )?); + model.prefill_chunk_forward(&tokens[split..], &mut kv, &mut recurrent, backend)? + } + None => model.prefill_chunk_forward(tokens, &mut kv, &mut recurrent, backend)?, + }; + let logits = last_token_logits(model, &hidden)?; + Ok((kv, recurrent, logits)) +} + +#[cfg(feature = "gdn-validation")] +fn first_decode_logits( + model: &Qwen35Model, + token: u32, + kv: &mut pegainfer_core::kv_pool::KvState, + recurrent: &RecurrentState, +) -> Result> { + let mut graph = model.create_batch_decode_graph_state_with_capacity(1)?; + graph.copy_state_to_slot(model.device_ctx(), recurrent, 0)?; + let mut kv_refs = vec![kv]; + model.batch_decode_graph(&[token], &mut kv_refs, &mut graph)?; + graph.buffers.logits.to_host(model.device_ctx()) +} + +#[test] +fn checked_prefill_end_pos_accepts_config_limit() { + assert_eq!( + checked_prefill_end_pos(0, 262_144, 262_144).unwrap(), + 262_144 + ); + assert_eq!( + checked_prefill_end_pos(262_143, 1, 262_144).unwrap(), + 262_144 + ); +} + +#[test] +fn checked_prefill_end_pos_rejects_past_config_limit() { + let err = checked_prefill_end_pos(0, 262_145, 262_144) + .unwrap_err() + .to_string(); + assert!(err.contains("beyond max_position_embeddings=262144")); + assert!(err.contains("requested end_pos=262145")); +} + +#[test] +fn checked_prefill_end_pos_rejects_overflow() { + let err = checked_prefill_end_pos(usize::MAX, 1, 262_144) + .unwrap_err() + .to_string(); + assert!(err.contains("prefill position overflow")); +} + +#[cfg(feature = "gdn-validation")] +#[test] +#[ignore = "requires an SM120 GPU, Qwen3.5-4B weights, and a build-linked validated FlashInfer bundle"] +fn flashinfer_gdn_chunk_continuation_and_model_outputs_match() -> Result<()> { + let model_path = required_model_path(); + let model = Qwen35Model::from_safetensors(&model_path, 0, 1)?; + let backend = model.resolved_gdn_backend(); + assert_eq!(backend, GdnPrefillBackend::FlashInfer); + + let evidence_before = model.flashinfer_gdn_runtime_evidence()?; + assert_eq!(evidence_before.selected_backend, "flashinfer"); + assert_ne!(evidence_before.artifact_sha256, "unavailable"); + assert_eq!(evidence_before.artifact_sha256.len(), 64); + assert_eq!(evidence_before.successful_launches, 0); + + // These deterministic token ids are only model inputs. All hidden values, + // Q/K/V/gates, recurrent state, and logits come from the real 4B weights. + let tokens = (0..128) + .map(|index| 100 + (index * 17 % 1000) as u32) + .collect::>(); + let (mut unchunked_kv, unchunked_state, unchunked_prefill_logits) = + run_prefill_case(&model, &tokens, backend, None)?; + let (mut chunked_kv, chunked_state, chunked_prefill_logits) = + run_prefill_case(&model, &tokens, backend, Some(64))?; + + assert_recurrent_continuation(&model, &unchunked_state, &chunked_state)?; + let decode_token = assert_logits_close( + "real-model last-token logits", + &unchunked_prefill_logits, + &chunked_prefill_logits, + ); + + let unchunked_decode = + first_decode_logits(&model, decode_token, &mut unchunked_kv, &unchunked_state)?; + let chunked_decode = + first_decode_logits(&model, decode_token, &mut chunked_kv, &chunked_state)?; + assert_logits_close( + "real-model first-decode logits", + &unchunked_decode, + &chunked_decode, + ); + + let evidence_after = model.flashinfer_gdn_runtime_evidence()?; + assert_eq!(evidence_after.selected_backend, "flashinfer"); + assert_eq!( + evidence_after.artifact_sha256, + evidence_before.artifact_sha256 + ); + let linear_layers = + model.config().num_hidden_layers - model.config().num_full_attention_layers(); + assert_eq!( + evidence_after.successful_launches - evidence_before.successful_launches, + (3 * linear_layers) as u64, + "chunk continuation gate did not execute one unchunked pass and two resumed model chunks" + ); + Ok(()) +} diff --git a/pegainfer-qwen35/src/prefill_buffers.rs b/pegainfer-qwen35/src/prefill_buffers.rs index 92b38eb8a..dd352bd83 100644 --- a/pegainfer-qwen35/src/prefill_buffers.rs +++ b/pegainfer-qwen35/src/prefill_buffers.rs @@ -8,6 +8,66 @@ use pegainfer_core::tensor::HiddenStates; use super::config::Config35; +/// Outputs of the native, non-expanded GDN prepare kernel. +/// +/// This buffer is intentionally separate from `GdrChunkwiseScratch35`: the +/// production Triton path below still requires value-head-expanded Q/K, while +/// the FlashInfer candidate consumes native Hq/Hk tensors directly. +pub(crate) struct GdnPrepareScratch35 { + /// Normalized native Q, bf16 token-major `[T,Hq,D]`. + pub(crate) q: HiddenStates, + /// Normalized native K, bf16 token-major `[T,Hk,D]`. + pub(crate) k: HiddenStates, + /// Raw V, bf16 token-major `[T,Hv,D]`. + pub(crate) v: HiddenStates, + /// Per-token decay multiplier, fp32 `[T,Hv]` (not log/cumulative alpha). + pub(crate) alpha: CudaSlice, + /// Per-token beta, fp32 `[T,Hv]`. + pub(crate) beta: CudaSlice, + /// Async validation result: zero means all consumed inputs were finite. + pub(crate) non_finite_status: CudaSlice, +} + +impl GdnPrepareScratch35 { + pub(crate) fn new(ctx: &DeviceContext, config: &Config35, seq_len: usize) -> Result { + anyhow::ensure!( + config.linear_num_key_heads == 16 + && config.linear_num_value_heads == 32 + && config.linear_key_head_dim == 128 + && config.linear_value_head_dim == 128, + "native GDN prepare requires Hq/Hk/Hv/D=16/16/32/128" + ); + Self::for_tokens(ctx, seq_len) + } + + pub(crate) fn for_tokens(ctx: &DeviceContext, seq_len: usize) -> Result { + anyhow::ensure!(seq_len > 0, "native GDN prepare requires T>=1"); + + const H_Q: usize = 16; + const H_K: usize = 16; + const H_V: usize = 32; + const HEAD_DIM: usize = 128; + + Ok(Self { + q: HiddenStates::zeros(ctx, H_Q * HEAD_DIM, seq_len)?, + k: HiddenStates::zeros(ctx, H_K * HEAD_DIM, seq_len)?, + v: HiddenStates::zeros(ctx, H_V * HEAD_DIM, seq_len)?, + alpha: ctx + .stream + .alloc_zeros(seq_len * H_V) + .map_err(|e| anyhow::anyhow!("Alloc native GDN alpha failed: {e}"))?, + beta: ctx + .stream + .alloc_zeros(seq_len * H_V) + .map_err(|e| anyhow::anyhow!("Alloc native GDN beta failed: {e}"))?, + non_finite_status: ctx + .stream + .alloc_zeros(1) + .map_err(|e| anyhow::anyhow!("Alloc native GDN status failed: {e}"))?, + }) + } +} + /// Scratch buffers for a single Qwen3.5 linear-attention chunk-wise GDR prefill call. /// /// The first implementation target is intentionally narrow: @@ -111,6 +171,30 @@ impl GdrChunkwiseScratch35 { seq_len.div_ceil(Self::CHUNK_SIZE) } + /// Device bytes owned by the Triton GDN operator for one prefill chunk. + /// + /// This intentionally excludes model-wide hidden/MLP/full-attention + /// temporaries and the recurrent state, which are common to both GDN + /// backends. The allocation list mirrors [`Self::from_dims`]. + fn operator_scratch_bytes_from_dims( + num_value_heads: usize, + key_dim: usize, + value_dim: usize, + seq_len: usize, + ) -> usize { + let kv_hidden_dim = num_value_heads * key_dim; + let vv_hidden_dim = num_value_heads * value_dim; + let num_chunks = seq_len.div_ceil(Self::CHUNK_SIZE); + + let f32_elems = seq_len * num_value_heads * 2 + + seq_len * num_value_heads * Self::CHUNK_SIZE + + num_chunks * num_value_heads * value_dim * key_dim; + let bf16_elems = seq_len * num_value_heads * Self::CHUNK_SIZE + + kv_hidden_dim * seq_len * 3 + + vv_hidden_dim * seq_len * 3; + f32_elems * size_of::() + bf16_elems * size_of::() + } + /// Estimate peak GPU memory (bytes) for prefill scratch at a given seq_len. /// /// Accounts for: @@ -124,28 +208,11 @@ impl GdrChunkwiseScratch35 { let num_vh = config.linear_num_value_heads; let key_dim = config.linear_key_head_dim; let val_dim = config.linear_value_head_dim; - let chunk_sz = Self::CHUNK_SIZE; - let num_chunks = max_seq_len.div_ceil(chunk_sz); let seq = max_seq_len; - let kv_hidden = num_vh * key_dim; - let vv_hidden = num_vh * val_dim; - // 1. GDR scratch (bf16 = 2 bytes, f32 = 4 bytes) - let gdr_bytes = { - let f32_elems = seq * num_vh // g_cumsum - + seq * num_vh // beta - + seq * num_vh * chunk_sz // a_tril - + num_chunks * num_vh * val_dim * key_dim; // chunk_state - let bf16_elems = seq * num_vh * chunk_sz // a_inv - + kv_hidden * seq // q_expanded - + kv_hidden * seq // k_expanded - + vv_hidden * seq // v_raw - + kv_hidden * seq // w - + vv_hidden * seq // u - + vv_hidden * seq; // v_new - f32_elems * 4 + bf16_elems * 2 - }; + let gdr_bytes = + Self::operator_scratch_bytes_from_dims(num_vh, key_dim, val_dim, max_seq_len); // 2. Per-layer transient peak (all bf16 = 2 bytes). // Attention and MLP temps don't coexist — MLP runs after attention. diff --git a/pegainfer-qwen35/src/recurrent.rs b/pegainfer-qwen35/src/recurrent.rs index 3fc76f10b..eb6054f01 100644 --- a/pegainfer-qwen35/src/recurrent.rs +++ b/pegainfer-qwen35/src/recurrent.rs @@ -1,3 +1,4 @@ +use anyhow::Context; use anyhow::Result; use cudarc::driver::CudaSlice; use cudarc::driver::DevicePtr; @@ -10,6 +11,7 @@ use crate::config::GDN_AOT_KEY_HEAD_DIM; use crate::config::GDN_AOT_VALUE_HEAD_DIM; use crate::config::LINEAR_CONV_MAX_KERNEL_DIM; use crate::ffi; +use crate::prefill_buffers::GdnPrepareScratch35; use crate::prefill_buffers::GdrChunkwiseScratch35; #[cfg(test)] @@ -178,6 +180,120 @@ pub(crate) fn conv1d_prefill_batch_into( } } +/// Prepare native Q/K/V plus per-token alpha/beta for the FlashInfer GDN +/// candidate. +/// +/// The preparation kernel reports non-finite inputs through a sticky device +/// status word owned by the chunk. The caller validates it once after the +/// layer loop, avoiding one D2H synchronization per layer while still refusing +/// to return a candidate result containing invalid inputs. +#[allow(clippy::too_many_arguments)] +pub(crate) fn gated_delta_rule_prefill_native_prepare_into( + ctx: &DeviceContext, + qkv: &HiddenStates, + b_proj: &HiddenStates, + a_proj: &HiddenStates, + dt_bias: &DeviceVec, + a_log: &CudaSlice, + scratch: &mut GdnPrepareScratch35, +) -> Result<()> { + const H_Q: usize = 16; + const H_K: usize = 16; + const H_V: usize = 32; + const HEAD_DIM: usize = 128; + anyhow::ensure!(qkv.seq_len > 0, "native GDN prepare requires T>=1"); + let expected_qkv = (H_Q + H_K + H_V) * HEAD_DIM; + anyhow::ensure!( + qkv.hidden_dim == expected_qkv, + "native GDN qkv hidden dim mismatch: expected {expected_qkv}, got {}", + qkv.hidden_dim + ); + anyhow::ensure!( + b_proj.hidden_dim == H_V && b_proj.seq_len == qkv.seq_len, + "native GDN b projection must be [T,Hv]=[{},{}]", + qkv.seq_len, + H_V + ); + anyhow::ensure!( + a_proj.hidden_dim == H_V && a_proj.seq_len == qkv.seq_len, + "native GDN a projection must be [T,Hv]=[{},{}]", + qkv.seq_len, + H_V + ); + anyhow::ensure!( + dt_bias.len == H_V, + "native GDN dt_bias length must be {H_V}, got {}", + dt_bias.len + ); + anyhow::ensure!( + a_log.len() == H_V, + "native GDN A_log length must be {H_V}, got {}", + a_log.len() + ); + anyhow::ensure!( + scratch.q.hidden_dim == H_Q * HEAD_DIM && scratch.q.seq_len == qkv.seq_len, + "native GDN Q output shape mismatch" + ); + anyhow::ensure!( + scratch.k.hidden_dim == H_K * HEAD_DIM && scratch.k.seq_len == qkv.seq_len, + "native GDN K output shape mismatch" + ); + anyhow::ensure!( + scratch.v.hidden_dim == H_V * HEAD_DIM && scratch.v.seq_len == qkv.seq_len, + "native GDN V output shape mismatch" + ); + anyhow::ensure!( + scratch.alpha.len() == qkv.seq_len * H_V, + "native GDN alpha output length mismatch" + ); + anyhow::ensure!( + scratch.beta.len() == qkv.seq_len * H_V, + "native GDN beta output length mismatch" + ); + anyhow::ensure!( + scratch.non_finite_status.len() == 1, + "native GDN status output length mismatch" + ); + let tokens: i32 = qkv + .seq_len + .try_into() + .context("native GDN prepare T exceeds i32")?; + + { + let (qkv_ptr, _gqkv) = qkv.data.device_ptr(&ctx.stream); + let (b_ptr, _gb) = b_proj.data.device_ptr(&ctx.stream); + let (a_ptr, _ga) = a_proj.data.device_ptr(&ctx.stream); + let (dt_ptr, _gdt) = dt_bias.data.device_ptr(&ctx.stream); + let (alog_ptr, _gal) = a_log.device_ptr(&ctx.stream); + let (q_out, _gqo) = scratch.q.data.device_ptr_mut(&ctx.stream); + let (k_out, _gko) = scratch.k.data.device_ptr_mut(&ctx.stream); + let (v_out, _gvo) = scratch.v.data.device_ptr_mut(&ctx.stream); + let (alpha_out, _gaout) = scratch.alpha.device_ptr_mut(&ctx.stream); + let (beta_out, _gbout) = scratch.beta.device_ptr_mut(&ctx.stream); + let (status_out, _gsout) = scratch.non_finite_status.device_ptr_mut(&ctx.stream); + + let result = unsafe { + ffi::gated_delta_rule_prefill_native_prepare_cuda( + qkv_ptr as *const ffi::Half, + b_ptr as *const ffi::Half, + a_ptr as *const ffi::Half, + dt_ptr as *const ffi::Half, + alog_ptr as *const f32, + q_out as *mut ffi::Half, + k_out as *mut ffi::Half, + v_out as *mut ffi::Half, + alpha_out as *mut f32, + beta_out as *mut f32, + status_out as *mut u32, + tokens, + ctx.stream.cu_stream(), + ) + }; + result.result()?; + } + Ok(()) +} + #[allow(clippy::too_many_arguments)] fn gated_delta_rule_prefill_chunk_prepare_into( ctx: &DeviceContext, @@ -548,435 +664,4 @@ pub fn gated_delta_rule_prefill_chunkwise_into( } #[cfg(test)] -mod tests { - use anyhow::Result; - use cudarc::driver::DevicePtrMut; - use half::bf16; - use pegainfer_core::tensor::DeviceContext; - use pegainfer_core::tensor::DeviceVec; - use pegainfer_core::tensor::HiddenStates; - - use super::conv1d_prefill_batch_into; - use super::gated_delta_rule_decode_batch_into; - use super::gated_delta_rule_decode_vec_into; - use super::gated_delta_rule_prefill_chunkwise_into; - use crate::prefill_buffers::GdrChunkwiseScratch35; - - fn bf16_vec(data: &[f32]) -> Vec { - data.iter().map(|&x| bf16::from_f32(x)).collect() - } - - #[test] - fn conv1d_prefill_handoff_matches_single_prefill() -> Result<()> { - let ctx = DeviceContext::new()?; - let num_channels = 1024usize; - let kernel_size = 4usize; - let total_seq = 18usize; - let prefix_seq = 5usize; - - let x_host = bf16_vec( - &(0..num_channels * total_seq) - .map(|i| ((i % 71) as f32 - 35.0) * 0.03125) - .collect::>(), - ); - let w_host = bf16_vec( - &(0..num_channels * kernel_size) - .map(|i| ((i % 19) as f32 - 9.0) * 0.0625) - .collect::>(), - ); - - let x_all = HiddenStates { - data: ctx.stream.clone_htod(&x_host)?, - hidden_dim: num_channels, - seq_len: total_seq, - }; - let conv_weight = DeviceVec::from_host(&ctx, &w_host)?; - let state_len = num_channels * (kernel_size - 1); - let zero_state = vec![bf16::ZERO; state_len]; - - let mut state_all = DeviceVec::from_host(&ctx, &zero_state)?; - let mut out_all = HiddenStates::zeros(&ctx, num_channels, total_seq)?; - conv1d_prefill_batch_into( - &ctx, - &x_all, - &conv_weight, - &mut state_all, - &mut out_all, - kernel_size, - ); - - let x_prefix = HiddenStates { - data: ctx - .stream - .clone_htod(&x_host[..num_channels * prefix_seq])?, - hidden_dim: num_channels, - seq_len: prefix_seq, - }; - let mut state_split = DeviceVec::from_host(&ctx, &zero_state)?; - let mut out_prefix = HiddenStates::zeros(&ctx, num_channels, prefix_seq)?; - conv1d_prefill_batch_into( - &ctx, - &x_prefix, - &conv_weight, - &mut state_split, - &mut out_prefix, - kernel_size, - ); - - for step in prefix_seq..total_seq { - let x_step = HiddenStates { - data: ctx - .stream - .clone_htod(&x_host[num_channels * step..num_channels * (step + 1)])?, - hidden_dim: num_channels, - seq_len: 1, - }; - let mut out_step = HiddenStates::zeros(&ctx, num_channels, 1)?; - conv1d_prefill_batch_into( - &ctx, - &x_step, - &conv_weight, - &mut state_split, - &mut out_step, - kernel_size, - ); - } - - let out_all_host = ctx.stream.clone_dtoh(&out_all.data)?; - let state_all_host = state_all.to_host(&ctx)?; - let state_split_host = state_split.to_host(&ctx)?; - ctx.sync()?; - - let out_all_host: Vec = out_all_host.iter().map(|x| x.to_f32()).collect(); - let expected_last = &out_all_host[num_channels * (total_seq - 1)..num_channels * total_seq]; - - let x_last = HiddenStates { - data: ctx - .stream - .clone_htod(&x_host[num_channels * (total_seq - 1)..num_channels * total_seq])?, - hidden_dim: num_channels, - seq_len: 1, - }; - let mut state_last = DeviceVec::from_host(&ctx, &zero_state)?; - let x_before_last = HiddenStates { - data: ctx - .stream - .clone_htod(&x_host[..num_channels * (total_seq - 1)])?, - hidden_dim: num_channels, - seq_len: total_seq - 1, - }; - let mut scratch_before_last = HiddenStates::zeros(&ctx, num_channels, total_seq - 1)?; - conv1d_prefill_batch_into( - &ctx, - &x_before_last, - &conv_weight, - &mut state_last, - &mut scratch_before_last, - kernel_size, - ); - let mut out_last = HiddenStates::zeros(&ctx, num_channels, 1)?; - conv1d_prefill_batch_into( - &ctx, - &x_last, - &conv_weight, - &mut state_last, - &mut out_last, - kernel_size, - ); - let out_last_host = ctx.stream.clone_dtoh(&out_last.data)?; - ctx.sync()?; - let out_last_host: Vec = out_last_host.iter().map(|x| x.to_f32()).collect(); - - let max_out_diff = expected_last - .iter() - .zip(out_last_host.iter()) - .map(|(a, b)| (a - b).abs()) - .fold(0.0_f32, f32::max); - let max_state_diff = state_all_host - .iter() - .zip(state_split_host.iter()) - .map(|(a, b)| (a - b).abs()) - .fold(0.0_f32, f32::max); - - assert!(max_out_diff < 0.02, "output diff {max_out_diff}"); - assert!(max_state_diff < 0.02, "state diff {max_state_diff}"); - Ok(()) - } - - #[test] - fn gdr_decode_batch_matches_single_slot_reference() -> Result<()> { - let ctx = DeviceContext::new()?; - let batch_size = 3usize; - let num_key_heads = 16usize; - let num_value_heads = 48usize; - let key_dim = 128usize; - let val_dim = 128usize; - - let qkv_dim = 2 * num_key_heads * key_dim + num_value_heads * val_dim; - let out_dim = num_value_heads * val_dim; - let state_len = num_value_heads * key_dim * val_dim; - - let qkv_host = bf16_vec( - &(0..batch_size * qkv_dim) - .map(|i| ((i % 89) as f32 - 44.0) * 0.007_812_5) - .collect::>(), - ); - let b_host = bf16_vec( - &(0..batch_size * num_value_heads) - .map(|i| ((i % 11) as f32 - 5.0) * 0.03125) - .collect::>(), - ); - let a_host = bf16_vec( - &(0..batch_size * num_value_heads) - .map(|i| ((i % 13) as f32 - 6.0) * 0.03125) - .collect::>(), - ); - let dt_host = bf16_vec( - &(0..num_value_heads) - .map(|i| ((i % 7) as f32 - 3.0) * 0.0625) - .collect::>(), - ); - let alog_host: Vec = (0..num_value_heads) - .map(|i| ((i % 5) as f32 - 2.0) * 0.125) - .collect(); - - let qkv_batch = HiddenStates { - data: ctx.stream.clone_htod(&qkv_host)?, - hidden_dim: qkv_dim, - seq_len: batch_size, - }; - let b_batch = HiddenStates { - data: ctx.stream.clone_htod(&b_host)?, - hidden_dim: num_value_heads, - seq_len: batch_size, - }; - let a_batch = HiddenStates { - data: ctx.stream.clone_htod(&a_host)?, - hidden_dim: num_value_heads, - seq_len: batch_size, - }; - let dt_bias = DeviceVec::from_host(&ctx, &dt_host)?; - let a_log = ctx.stream.clone_htod(&alog_host)?; - - let mut batch_states: Vec> = (0..batch_size) - .map(|_| ctx.stream.alloc_zeros(state_len)) - .collect::, _>>()?; - let mut state_ptrs = Vec::with_capacity(batch_size); - for state in &mut batch_states { - let (ptr, _guard) = state.device_ptr_mut(&ctx.stream); - state_ptrs.push(ptr); - } - let state_ptrs_d = ctx.stream.clone_htod(&state_ptrs)?; - - let mut out_batch = HiddenStates::zeros(&ctx, out_dim, batch_size)?; - gated_delta_rule_decode_batch_into( - &ctx, - &qkv_batch, - &b_batch, - &a_batch, - &dt_bias, - &a_log, - &state_ptrs_d, - &mut out_batch, - batch_size, - num_key_heads, - num_value_heads, - key_dim, - val_dim, - ); - - let mut out_ref_rows: Vec = Vec::with_capacity(batch_size * out_dim); - let mut ref_states = Vec::with_capacity(batch_size); - for row in 0..batch_size { - let qkv_row = - DeviceVec::from_host(&ctx, &qkv_host[row * qkv_dim..(row + 1) * qkv_dim])?; - let b_row = DeviceVec::from_host( - &ctx, - &b_host[row * num_value_heads..(row + 1) * num_value_heads], - )?; - let a_row = DeviceVec::from_host( - &ctx, - &a_host[row * num_value_heads..(row + 1) * num_value_heads], - )?; - let mut state_ref: cudarc::driver::CudaSlice = - ctx.stream.alloc_zeros(state_len)?; - let mut out_row = DeviceVec::zeros(&ctx, out_dim)?; - gated_delta_rule_decode_vec_into( - &ctx, - &qkv_row, - &b_row, - &a_row, - &dt_bias, - &a_log, - &mut state_ref, - &mut out_row, - num_key_heads, - num_value_heads, - key_dim, - val_dim, - ); - out_ref_rows.extend_from_slice(&out_row.to_host(&ctx)?); - ref_states.push(state_ref); - } - - let out_batch_host = ctx.stream.clone_dtoh(&out_batch.data)?; - ctx.sync()?; - let out_batch_host: Vec = out_batch_host.iter().map(|x| x.to_f32()).collect(); - let max_out_diff = out_batch_host - .iter() - .zip(out_ref_rows.iter()) - .map(|(a, b)| (a - b).abs()) - .fold(0.0_f32, f32::max); - - let mut max_state_diff = 0.0_f32; - for (batch_state, ref_state) in batch_states.iter().zip(ref_states.iter()) { - let batch_state_host = ctx.stream.clone_dtoh(batch_state)?; - let ref_state_host = ctx.stream.clone_dtoh(ref_state)?; - ctx.sync()?; - max_state_diff = max_state_diff.max( - batch_state_host - .iter() - .zip(ref_state_host.iter()) - .map(|(a, b)| (a - b).abs()) - .fold(0.0_f32, f32::max), - ); - } - - assert!(max_out_diff < 0.05, "output diff {max_out_diff}"); - assert!(max_state_diff < 0.05, "state diff {max_state_diff}"); - Ok(()) - } - - #[test] - fn gdn_chunkwise_prefill_matches_stepwise_decode_at_48_value_heads() -> Result<()> { - let ctx = DeviceContext::new()?; - let num_key_heads = 16usize; - let num_value_heads = 48usize; - let key_dim = 128usize; - let val_dim = 128usize; - let seq_len = 96usize; - - let qkv_dim = 2 * num_key_heads * key_dim + num_value_heads * val_dim; - let out_dim = num_value_heads * val_dim; - let state_len = num_value_heads * key_dim * val_dim; - - let qkv_host = bf16_vec( - &(0..seq_len * qkv_dim) - .map(|i| ((i % 73) as f32 - 36.0) * 0.01) - .collect::>(), - ); - let b_host = bf16_vec( - &(0..seq_len * num_value_heads) - .map(|i| ((i % 13) as f32 - 6.0) * 0.05) - .collect::>(), - ); - let a_host = bf16_vec( - &(0..seq_len * num_value_heads) - .map(|i| ((i % 17) as f32 - 8.0) * 0.05) - .collect::>(), - ); - let dt_host = bf16_vec( - &(0..num_value_heads) - .map(|i| ((i % 7) as f32 - 3.0) * 0.1) - .collect::>(), - ); - let alog_host: Vec = (0..num_value_heads) - .map(|i| ((i % 5) as f32 - 2.0) * 0.2) - .collect(); - - let dt_bias = DeviceVec::from_host(&ctx, &dt_host)?; - let a_log = ctx.stream.clone_htod(&alog_host)?; - - let qkv_all = HiddenStates { - data: ctx.stream.clone_htod(&qkv_host)?, - hidden_dim: qkv_dim, - seq_len, - }; - let b_all = HiddenStates { - data: ctx.stream.clone_htod(&b_host)?, - hidden_dim: num_value_heads, - seq_len, - }; - let a_all = HiddenStates { - data: ctx.stream.clone_htod(&a_host)?, - hidden_dim: num_value_heads, - seq_len, - }; - let mut state_chunk: cudarc::driver::CudaSlice = ctx.stream.alloc_zeros(state_len)?; - let mut scratch = - GdrChunkwiseScratch35::from_dims(&ctx, num_value_heads, key_dim, val_dim, seq_len)?; - let mut out_chunk = HiddenStates::zeros(&ctx, out_dim, seq_len)?; - gated_delta_rule_prefill_chunkwise_into( - &ctx, - &qkv_all, - &b_all, - &a_all, - &dt_bias, - &a_log, - &mut state_chunk, - &mut scratch, - &mut out_chunk, - num_key_heads, - num_value_heads, - key_dim, - val_dim, - )?; - - let mut state_step: cudarc::driver::CudaSlice = ctx.stream.alloc_zeros(state_len)?; - let mut out_step_rows: Vec = Vec::with_capacity(seq_len * out_dim); - for t in 0..seq_len { - let qkv_t = DeviceVec::from_host(&ctx, &qkv_host[t * qkv_dim..(t + 1) * qkv_dim])?; - let b_t = DeviceVec::from_host( - &ctx, - &b_host[t * num_value_heads..(t + 1) * num_value_heads], - )?; - let a_t = DeviceVec::from_host( - &ctx, - &a_host[t * num_value_heads..(t + 1) * num_value_heads], - )?; - let mut out_t = DeviceVec::from_host(&ctx, &vec![bf16::ZERO; out_dim])?; - gated_delta_rule_decode_vec_into( - &ctx, - &qkv_t, - &b_t, - &a_t, - &dt_bias, - &a_log, - &mut state_step, - &mut out_t, - num_key_heads, - num_value_heads, - key_dim, - val_dim, - ); - let row = out_t.to_host(&ctx)?; - out_step_rows.extend_from_slice(&row); - } - - let out_chunk_host = ctx.stream.clone_dtoh(&out_chunk.data)?; - let state_chunk_host = ctx.stream.clone_dtoh(&state_chunk)?; - let state_step_host = ctx.stream.clone_dtoh(&state_step)?; - ctx.sync()?; - let out_chunk_host: Vec = out_chunk_host.iter().map(|x| x.to_f32()).collect(); - - let max_out_diff = out_chunk_host - .iter() - .zip(out_step_rows.iter()) - .map(|(a, b)| (a - b).abs()) - .fold(0.0_f32, f32::max); - let max_state_diff = state_chunk_host - .iter() - .zip(state_step_host.iter()) - .map(|(a, b)| (a - b).abs()) - .fold(0.0_f32, f32::max); - - assert!( - out_chunk_host.iter().all(|x| x.is_finite()) - && state_chunk_host.iter().all(|x| x.is_finite()), - "chunkwise outputs must be finite" - ); - assert!(max_out_diff < 0.05, "output diff {max_out_diff}"); - assert!(max_state_diff < 0.05, "state diff {max_state_diff}"); - Ok(()) - } -} +mod tests; diff --git a/pegainfer-qwen35/src/recurrent/tests.rs b/pegainfer-qwen35/src/recurrent/tests.rs new file mode 100644 index 000000000..72a02b149 --- /dev/null +++ b/pegainfer-qwen35/src/recurrent/tests.rs @@ -0,0 +1,789 @@ +use anyhow::Result; +use cudarc::driver::DevicePtrMut; +use half::bf16; +use pegainfer_core::tensor::DeviceContext; +use pegainfer_core::tensor::DeviceVec; +use pegainfer_core::tensor::HiddenStates; + +use super::conv1d_prefill_batch_into; +use super::gated_delta_rule_decode_batch_into; +use super::gated_delta_rule_decode_vec_into; +use super::gated_delta_rule_prefill_chunkwise_into; +use super::gated_delta_rule_prefill_native_prepare_into; +use crate::prefill_buffers::GdnPrepareScratch35; +use crate::prefill_buffers::GdrChunkwiseScratch35; + +fn bf16_vec(data: &[f32]) -> Vec { + data.iter().map(|&x| bf16::from_f32(x)).collect() +} + +fn assert_f32_close_with_stats( + label: &str, + expected: &[f32], + actual: &[f32], + atol: f32, + rtol: f32, +) { + assert_eq!(expected.len(), actual.len(), "{label} length mismatch"); + let mut deltas = Vec::with_capacity(expected.len()); + let mut max_relative = 0.0_f32; + let mut violation_count = 0usize; + let mut first_violation = None; + for (index, (&expected, &actual)) in expected.iter().zip(actual).enumerate() { + let delta = (expected - actual).abs(); + let relative = delta / expected.abs().max(actual.abs()).max(1.0e-12); + deltas.push(delta); + max_relative = max_relative.max(relative); + let violation = !expected.is_finite() + || !actual.is_finite() + || delta > atol + rtol * expected.abs().max(actual.abs()); + if violation { + violation_count += 1; + if first_violation.is_none() { + first_violation = Some((index, expected, actual, delta)); + } + } + } + deltas.sort_by(f32::total_cmp); + let max = deltas.last().copied().unwrap_or(0.0); + let mean = if deltas.is_empty() { + 0.0 + } else { + deltas.iter().sum::() / deltas.len() as f32 + }; + let p99 = deltas + .get(deltas.len().saturating_sub(1) * 99 / 100) + .copied() + .unwrap_or(0.0); + eprintln!( + "{label}: elements={} violations={violation_count} max_abs={max:.8} mean_abs={mean:.8} p99_abs={p99:.8} max_rel={max_relative:.8} atol={atol} rtol={rtol}", + deltas.len() + ); + assert!( + first_violation.is_none(), + "{label} first violation {:?}; violations={violation_count}/{} max_abs={max} mean_abs={mean} p99_abs={p99} max_rel={max_relative}", + first_violation, + expected.len(), + ); +} + +fn assert_bf16_bits_equal(label: &str, expected: &[bf16], actual: &[bf16]) { + assert_eq!(expected.len(), actual.len(), "{label} length mismatch"); + let first_mismatch = expected + .iter() + .zip(actual) + .position(|(expected, actual)| expected.to_bits() != actual.to_bits()); + assert!( + first_mismatch.is_none(), + "{label} first bitwise mismatch at {:?}: expected={:?} actual={:?}", + first_mismatch, + first_mismatch.map(|index| expected[index].to_f32()), + first_mismatch.map(|index| actual[index].to_f32()), + ); + eprintln!("{label}: elements={} bitwise_mismatches=0", expected.len()); +} + +fn softplus(value: f32) -> f32 { + if value > 20.0 { + value + } else if value < -20.0 { + value.exp() + } else { + value.exp().ln_1p() + } +} + +fn sigmoid(value: f32) -> f32 { + let exp = if value < 0.0 { + value.exp() + } else { + (-value).exp() + }; + if value >= 0.0 { + 1.0 / (1.0 + exp) + } else { + exp / (1.0 + exp) + } +} + +#[test] +#[ignore = "requires a CUDA GPU"] +fn native_prepare_hv32_dynamic_t_and_non_finite_inputs() -> Result<()> { + let ctx = DeviceContext::new()?; + let h_q = 16usize; + let h_k = 16usize; + let h_v = 32usize; + let d = 128usize; + let qkv_dim = (h_q + h_k + h_v) * d; + let dt_host = bf16_vec( + &(0..h_v) + .map(|head| (head as f32 - h_v as f32 / 2.0) / 64.0) + .collect::>(), + ); + let a_log_host = (0..h_v) + .map(|head| -2.5 + head as f32 / h_v as f32) + .collect::>(); + let dt_bias = DeviceVec::from_host(&ctx, &dt_host)?; + let a_log = ctx.stream.clone_htod(&a_log_host)?; + + for tokens in [1usize, 63, 64, 65, 128, 2048] { + let qkv_host = bf16_vec( + &(0..tokens * qkv_dim) + .map(|index| { + let signed = ((index * 37 + 11) % 251) as i32 - 125; + signed as f32 / 31.0 + }) + .collect::>(), + ); + let b_host = bf16_vec( + &(0..tokens * h_v) + .map(|index| ((index * 13 % 41) as f32 - 20.0) / 7.0) + .collect::>(), + ); + let a_host = bf16_vec( + &(0..tokens * h_v) + .map(|index| ((index * 17 % 47) as f32 - 23.0) / 9.0) + .collect::>(), + ); + let qkv = HiddenStates { + data: ctx.stream.clone_htod(&qkv_host)?, + hidden_dim: qkv_dim, + seq_len: tokens, + }; + let b = HiddenStates { + data: ctx.stream.clone_htod(&b_host)?, + hidden_dim: h_v, + seq_len: tokens, + }; + let a = HiddenStates { + data: ctx.stream.clone_htod(&a_host)?, + hidden_dim: h_v, + seq_len: tokens, + }; + let mut prepared = GdnPrepareScratch35::for_tokens(&ctx, tokens)?; + gated_delta_rule_prefill_native_prepare_into( + &ctx, + &qkv, + &b, + &a, + &dt_bias, + &a_log, + &mut prepared, + )?; + + let status = ctx.stream.clone_dtoh(&prepared.non_finite_status)?; + let q_actual = ctx.stream.clone_dtoh(&prepared.q.data)?; + let k_actual = ctx.stream.clone_dtoh(&prepared.k.data)?; + let v_actual = ctx.stream.clone_dtoh(&prepared.v.data)?; + let alpha_actual = ctx.stream.clone_dtoh(&prepared.alpha)?; + let beta_actual = ctx.stream.clone_dtoh(&prepared.beta)?; + ctx.sync()?; + assert_eq!(status, [0], "finite Hv32 T={tokens} fixture was rejected"); + + let mut q_expected = Vec::with_capacity(tokens * h_q * d); + let mut k_expected = Vec::with_capacity(tokens * h_k * d); + let mut v_expected = Vec::with_capacity(tokens * h_v * d); + for token in 0..tokens { + let token_qkv = token * qkv_dim; + for head in 0..h_q { + let input = token_qkv + head * d; + let output = (token * h_q + head) * d; + let sum_sq = qkv_host[input..input + d] + .iter() + .map(|value| value.to_f32().powi(2)) + .sum::(); + let inv_norm = (sum_sq + 1.0e-12).sqrt().recip(); + for lane in 0..d { + q_expected.push(qkv_host[input + lane].to_f32() * inv_norm); + } + debug_assert_eq!(q_expected.len(), output + d); + } + for head in 0..h_k { + let input = token_qkv + h_q * d + head * d; + let output = (token * h_k + head) * d; + let sum_sq = qkv_host[input..input + d] + .iter() + .map(|value| value.to_f32().powi(2)) + .sum::(); + let inv_norm = (sum_sq + 1.0e-12).sqrt().recip(); + for lane in 0..d { + k_expected.push(qkv_host[input + lane].to_f32() * inv_norm); + } + debug_assert_eq!(k_expected.len(), output + d); + } + let v_input = token_qkv + (h_q + h_k) * d; + v_expected.extend_from_slice(&qkv_host[v_input..v_input + h_v * d]); + } + let q_actual_f32 = q_actual + .iter() + .map(|value| value.to_f32()) + .collect::>(); + let k_actual_f32 = k_actual + .iter() + .map(|value| value.to_f32()) + .collect::>(); + assert_f32_close_with_stats( + &format!("native prepare Q [T={tokens},H={h_q},D={d},bf16]"), + &q_expected, + &q_actual_f32, + 1.0 / 256.0, + 0.0, + ); + assert_f32_close_with_stats( + &format!("native prepare K [T={tokens},H={h_k},D={d},bf16]"), + &k_expected, + &k_actual_f32, + 1.0 / 256.0, + 0.0, + ); + assert_bf16_bits_equal( + &format!("native prepare V [T={tokens},H={h_v},D={d},bf16]"), + &v_expected, + &v_actual, + ); + let mut alpha_expected = Vec::with_capacity(tokens * h_v); + let mut beta_expected = Vec::with_capacity(tokens * h_v); + for index in 0..tokens * h_v { + let head = index % h_v; + let a_value = a_host[index].to_f32(); + let b_value = b_host[index].to_f32(); + let expected_alpha = + (-a_log_host[head].exp() * softplus(a_value + dt_host[head].to_f32())).exp(); + let expected_beta = sigmoid(b_value); + alpha_expected.push(expected_alpha); + beta_expected.push(expected_beta); + } + assert_f32_close_with_stats( + &format!("native prepare alpha [T={tokens},H={h_v},f32]"), + &alpha_expected, + &alpha_actual, + 2.0e-6, + 2.0e-6, + ); + assert_f32_close_with_stats( + &format!("native prepare beta [T={tokens},H={h_v},f32]"), + &beta_expected, + &beta_actual, + 2.0e-6, + 2.0e-6, + ); + } + + for non_finite_source in ["q", "v", "gate"] { + let mut qkv_host = vec![bf16::from_f32(0.25); qkv_dim]; + let b_host = vec![bf16::from_f32(-0.5); h_v]; + let mut a_host = vec![bf16::from_f32(0.5); h_v]; + match non_finite_source { + "q" => qkv_host[0] = bf16::from_bits(0x7fc0), + "v" => qkv_host[(h_q + h_k) * d + 7] = bf16::from_bits(0x7fc0), + "gate" => a_host[0] = bf16::from_bits(0x7fc0), + _ => unreachable!(), + } + let qkv = HiddenStates { + data: ctx.stream.clone_htod(&qkv_host)?, + hidden_dim: qkv_dim, + seq_len: 1, + }; + let b = HiddenStates { + data: ctx.stream.clone_htod(&b_host)?, + hidden_dim: h_v, + seq_len: 1, + }; + let a = HiddenStates { + data: ctx.stream.clone_htod(&a_host)?, + hidden_dim: h_v, + seq_len: 1, + }; + let mut prepared = GdnPrepareScratch35::for_tokens(&ctx, 1)?; + gated_delta_rule_prefill_native_prepare_into( + &ctx, + &qkv, + &b, + &a, + &dt_bias, + &a_log, + &mut prepared, + )?; + let status = ctx.stream.clone_dtoh(&prepared.non_finite_status)?; + ctx.sync()?; + assert_eq!( + status, + [1], + "non-finite {non_finite_source} input was not reported" + ); + } + + let finite_qkv_host = vec![bf16::from_f32(0.25); qkv_dim]; + let mut non_finite_qkv_host = finite_qkv_host.clone(); + non_finite_qkv_host[0] = bf16::from_bits(0x7fc0); + let gate_b_host = vec![bf16::from_f32(-0.5); h_v]; + let gate_a_host = vec![bf16::from_f32(0.5); h_v]; + let make_hidden = |values: &[bf16], hidden_dim: usize| -> Result { + Ok(HiddenStates { + data: ctx.stream.clone_htod(values)?, + hidden_dim, + seq_len: 1, + }) + }; + let non_finite_qkv = make_hidden(&non_finite_qkv_host, qkv_dim)?; + let finite_qkv = make_hidden(&finite_qkv_host, qkv_dim)?; + let gate_b = make_hidden(&gate_b_host, h_v)?; + let gate_a = make_hidden(&gate_a_host, h_v)?; + let mut sticky = GdnPrepareScratch35::for_tokens(&ctx, 1)?; + gated_delta_rule_prefill_native_prepare_into( + &ctx, + &non_finite_qkv, + &gate_b, + &gate_a, + &dt_bias, + &a_log, + &mut sticky, + )?; + gated_delta_rule_prefill_native_prepare_into( + &ctx, + &finite_qkv, + &gate_b, + &gate_a, + &dt_bias, + &a_log, + &mut sticky, + )?; + let sticky_status = ctx.stream.clone_dtoh(&sticky.non_finite_status)?; + ctx.sync()?; + assert_eq!( + sticky_status, + [1], + "a later finite layer cleared the chunk-owned non-finite status" + ); + + let mut fresh_chunk = GdnPrepareScratch35::for_tokens(&ctx, 1)?; + gated_delta_rule_prefill_native_prepare_into( + &ctx, + &finite_qkv, + &gate_b, + &gate_a, + &dt_bias, + &a_log, + &mut fresh_chunk, + )?; + let fresh_status = ctx.stream.clone_dtoh(&fresh_chunk.non_finite_status)?; + ctx.sync()?; + assert_eq!( + fresh_status, + [0], + "a new chunk did not start with a clear non-finite status" + ); + Ok(()) +} + +#[test] +fn conv1d_prefill_handoff_matches_single_prefill() -> Result<()> { + let ctx = DeviceContext::new()?; + let num_channels = 1024usize; + let kernel_size = 4usize; + let total_seq = 18usize; + let prefix_seq = 5usize; + + let x_host = bf16_vec( + &(0..num_channels * total_seq) + .map(|i| ((i % 71) as f32 - 35.0) * 0.03125) + .collect::>(), + ); + let w_host = bf16_vec( + &(0..num_channels * kernel_size) + .map(|i| ((i % 19) as f32 - 9.0) * 0.0625) + .collect::>(), + ); + + let x_all = HiddenStates { + data: ctx.stream.clone_htod(&x_host)?, + hidden_dim: num_channels, + seq_len: total_seq, + }; + let conv_weight = DeviceVec::from_host(&ctx, &w_host)?; + let state_len = num_channels * (kernel_size - 1); + let zero_state = vec![bf16::ZERO; state_len]; + + let mut state_all = DeviceVec::from_host(&ctx, &zero_state)?; + let mut out_all = HiddenStates::zeros(&ctx, num_channels, total_seq)?; + conv1d_prefill_batch_into( + &ctx, + &x_all, + &conv_weight, + &mut state_all, + &mut out_all, + kernel_size, + ); + + let x_prefix = HiddenStates { + data: ctx + .stream + .clone_htod(&x_host[..num_channels * prefix_seq])?, + hidden_dim: num_channels, + seq_len: prefix_seq, + }; + let mut state_split = DeviceVec::from_host(&ctx, &zero_state)?; + let mut out_prefix = HiddenStates::zeros(&ctx, num_channels, prefix_seq)?; + conv1d_prefill_batch_into( + &ctx, + &x_prefix, + &conv_weight, + &mut state_split, + &mut out_prefix, + kernel_size, + ); + + for step in prefix_seq..total_seq { + let x_step = HiddenStates { + data: ctx + .stream + .clone_htod(&x_host[num_channels * step..num_channels * (step + 1)])?, + hidden_dim: num_channels, + seq_len: 1, + }; + let mut out_step = HiddenStates::zeros(&ctx, num_channels, 1)?; + conv1d_prefill_batch_into( + &ctx, + &x_step, + &conv_weight, + &mut state_split, + &mut out_step, + kernel_size, + ); + } + + let out_all_host = ctx.stream.clone_dtoh(&out_all.data)?; + let state_all_host = state_all.to_host(&ctx)?; + let state_split_host = state_split.to_host(&ctx)?; + ctx.sync()?; + + let out_all_host: Vec = out_all_host.iter().map(|x| x.to_f32()).collect(); + let expected_last = &out_all_host[num_channels * (total_seq - 1)..num_channels * total_seq]; + + let x_last = HiddenStates { + data: ctx + .stream + .clone_htod(&x_host[num_channels * (total_seq - 1)..num_channels * total_seq])?, + hidden_dim: num_channels, + seq_len: 1, + }; + let mut state_last = DeviceVec::from_host(&ctx, &zero_state)?; + let x_before_last = HiddenStates { + data: ctx + .stream + .clone_htod(&x_host[..num_channels * (total_seq - 1)])?, + hidden_dim: num_channels, + seq_len: total_seq - 1, + }; + let mut scratch_before_last = HiddenStates::zeros(&ctx, num_channels, total_seq - 1)?; + conv1d_prefill_batch_into( + &ctx, + &x_before_last, + &conv_weight, + &mut state_last, + &mut scratch_before_last, + kernel_size, + ); + let mut out_last = HiddenStates::zeros(&ctx, num_channels, 1)?; + conv1d_prefill_batch_into( + &ctx, + &x_last, + &conv_weight, + &mut state_last, + &mut out_last, + kernel_size, + ); + let out_last_host = ctx.stream.clone_dtoh(&out_last.data)?; + ctx.sync()?; + let out_last_host: Vec = out_last_host.iter().map(|x| x.to_f32()).collect(); + + let max_out_diff = expected_last + .iter() + .zip(out_last_host.iter()) + .map(|(a, b)| (a - b).abs()) + .fold(0.0_f32, f32::max); + let max_state_diff = state_all_host + .iter() + .zip(state_split_host.iter()) + .map(|(a, b)| (a - b).abs()) + .fold(0.0_f32, f32::max); + + assert!(max_out_diff < 0.02, "output diff {max_out_diff}"); + assert!(max_state_diff < 0.02, "state diff {max_state_diff}"); + Ok(()) +} + +#[test] +fn gdr_decode_batch_matches_single_slot_reference() -> Result<()> { + let ctx = DeviceContext::new()?; + let batch_size = 3usize; + let num_key_heads = 16usize; + let num_value_heads = 48usize; + let key_dim = 128usize; + let val_dim = 128usize; + + let qkv_dim = 2 * num_key_heads * key_dim + num_value_heads * val_dim; + let out_dim = num_value_heads * val_dim; + let state_len = num_value_heads * key_dim * val_dim; + + let qkv_host = bf16_vec( + &(0..batch_size * qkv_dim) + .map(|i| ((i % 89) as f32 - 44.0) * 0.007_812_5) + .collect::>(), + ); + let b_host = bf16_vec( + &(0..batch_size * num_value_heads) + .map(|i| ((i % 11) as f32 - 5.0) * 0.03125) + .collect::>(), + ); + let a_host = bf16_vec( + &(0..batch_size * num_value_heads) + .map(|i| ((i % 13) as f32 - 6.0) * 0.03125) + .collect::>(), + ); + let dt_host = bf16_vec( + &(0..num_value_heads) + .map(|i| ((i % 7) as f32 - 3.0) * 0.0625) + .collect::>(), + ); + let alog_host: Vec = (0..num_value_heads) + .map(|i| ((i % 5) as f32 - 2.0) * 0.125) + .collect(); + + let qkv_batch = HiddenStates { + data: ctx.stream.clone_htod(&qkv_host)?, + hidden_dim: qkv_dim, + seq_len: batch_size, + }; + let b_batch = HiddenStates { + data: ctx.stream.clone_htod(&b_host)?, + hidden_dim: num_value_heads, + seq_len: batch_size, + }; + let a_batch = HiddenStates { + data: ctx.stream.clone_htod(&a_host)?, + hidden_dim: num_value_heads, + seq_len: batch_size, + }; + let dt_bias = DeviceVec::from_host(&ctx, &dt_host)?; + let a_log = ctx.stream.clone_htod(&alog_host)?; + + let mut batch_states: Vec> = (0..batch_size) + .map(|_| ctx.stream.alloc_zeros(state_len)) + .collect::, _>>()?; + let mut state_ptrs = Vec::with_capacity(batch_size); + for state in &mut batch_states { + let (ptr, _guard) = state.device_ptr_mut(&ctx.stream); + state_ptrs.push(ptr); + } + let state_ptrs_d = ctx.stream.clone_htod(&state_ptrs)?; + + let mut out_batch = HiddenStates::zeros(&ctx, out_dim, batch_size)?; + gated_delta_rule_decode_batch_into( + &ctx, + &qkv_batch, + &b_batch, + &a_batch, + &dt_bias, + &a_log, + &state_ptrs_d, + &mut out_batch, + batch_size, + num_key_heads, + num_value_heads, + key_dim, + val_dim, + ); + + let mut out_ref_rows: Vec = Vec::with_capacity(batch_size * out_dim); + let mut ref_states = Vec::with_capacity(batch_size); + for row in 0..batch_size { + let qkv_row = DeviceVec::from_host(&ctx, &qkv_host[row * qkv_dim..(row + 1) * qkv_dim])?; + let b_row = DeviceVec::from_host( + &ctx, + &b_host[row * num_value_heads..(row + 1) * num_value_heads], + )?; + let a_row = DeviceVec::from_host( + &ctx, + &a_host[row * num_value_heads..(row + 1) * num_value_heads], + )?; + let mut state_ref: cudarc::driver::CudaSlice = ctx.stream.alloc_zeros(state_len)?; + let mut out_row = DeviceVec::zeros(&ctx, out_dim)?; + gated_delta_rule_decode_vec_into( + &ctx, + &qkv_row, + &b_row, + &a_row, + &dt_bias, + &a_log, + &mut state_ref, + &mut out_row, + num_key_heads, + num_value_heads, + key_dim, + val_dim, + ); + out_ref_rows.extend_from_slice(&out_row.to_host(&ctx)?); + ref_states.push(state_ref); + } + + let out_batch_host = ctx.stream.clone_dtoh(&out_batch.data)?; + ctx.sync()?; + let out_batch_host: Vec = out_batch_host.iter().map(|x| x.to_f32()).collect(); + let max_out_diff = out_batch_host + .iter() + .zip(out_ref_rows.iter()) + .map(|(a, b)| (a - b).abs()) + .fold(0.0_f32, f32::max); + + let mut max_state_diff = 0.0_f32; + for (batch_state, ref_state) in batch_states.iter().zip(ref_states.iter()) { + let batch_state_host = ctx.stream.clone_dtoh(batch_state)?; + let ref_state_host = ctx.stream.clone_dtoh(ref_state)?; + ctx.sync()?; + max_state_diff = max_state_diff.max( + batch_state_host + .iter() + .zip(ref_state_host.iter()) + .map(|(a, b)| (a - b).abs()) + .fold(0.0_f32, f32::max), + ); + } + + assert!(max_out_diff < 0.05, "output diff {max_out_diff}"); + assert!(max_state_diff < 0.05, "state diff {max_state_diff}"); + Ok(()) +} + +#[test] +fn gdn_chunkwise_prefill_matches_stepwise_decode_at_48_value_heads() -> Result<()> { + let ctx = DeviceContext::new()?; + let num_key_heads = 16usize; + let num_value_heads = 48usize; + let key_dim = 128usize; + let val_dim = 128usize; + let seq_len = 96usize; + + let qkv_dim = 2 * num_key_heads * key_dim + num_value_heads * val_dim; + let out_dim = num_value_heads * val_dim; + let state_len = num_value_heads * key_dim * val_dim; + + let qkv_host = bf16_vec( + &(0..seq_len * qkv_dim) + .map(|i| ((i % 73) as f32 - 36.0) * 0.01) + .collect::>(), + ); + let b_host = bf16_vec( + &(0..seq_len * num_value_heads) + .map(|i| ((i % 13) as f32 - 6.0) * 0.05) + .collect::>(), + ); + let a_host = bf16_vec( + &(0..seq_len * num_value_heads) + .map(|i| ((i % 17) as f32 - 8.0) * 0.05) + .collect::>(), + ); + let dt_host = bf16_vec( + &(0..num_value_heads) + .map(|i| ((i % 7) as f32 - 3.0) * 0.1) + .collect::>(), + ); + let alog_host: Vec = (0..num_value_heads) + .map(|i| ((i % 5) as f32 - 2.0) * 0.2) + .collect(); + + let dt_bias = DeviceVec::from_host(&ctx, &dt_host)?; + let a_log = ctx.stream.clone_htod(&alog_host)?; + + let qkv_all = HiddenStates { + data: ctx.stream.clone_htod(&qkv_host)?, + hidden_dim: qkv_dim, + seq_len, + }; + let b_all = HiddenStates { + data: ctx.stream.clone_htod(&b_host)?, + hidden_dim: num_value_heads, + seq_len, + }; + let a_all = HiddenStates { + data: ctx.stream.clone_htod(&a_host)?, + hidden_dim: num_value_heads, + seq_len, + }; + let mut state_chunk: cudarc::driver::CudaSlice = ctx.stream.alloc_zeros(state_len)?; + let mut scratch = + GdrChunkwiseScratch35::from_dims(&ctx, num_value_heads, key_dim, val_dim, seq_len)?; + let mut out_chunk = HiddenStates::zeros(&ctx, out_dim, seq_len)?; + gated_delta_rule_prefill_chunkwise_into( + &ctx, + &qkv_all, + &b_all, + &a_all, + &dt_bias, + &a_log, + &mut state_chunk, + &mut scratch, + &mut out_chunk, + num_key_heads, + num_value_heads, + key_dim, + val_dim, + )?; + + let mut state_step: cudarc::driver::CudaSlice = ctx.stream.alloc_zeros(state_len)?; + let mut out_step_rows: Vec = Vec::with_capacity(seq_len * out_dim); + for t in 0..seq_len { + let qkv_t = DeviceVec::from_host(&ctx, &qkv_host[t * qkv_dim..(t + 1) * qkv_dim])?; + let b_t = DeviceVec::from_host( + &ctx, + &b_host[t * num_value_heads..(t + 1) * num_value_heads], + )?; + let a_t = DeviceVec::from_host( + &ctx, + &a_host[t * num_value_heads..(t + 1) * num_value_heads], + )?; + let mut out_t = DeviceVec::from_host(&ctx, &vec![bf16::ZERO; out_dim])?; + gated_delta_rule_decode_vec_into( + &ctx, + &qkv_t, + &b_t, + &a_t, + &dt_bias, + &a_log, + &mut state_step, + &mut out_t, + num_key_heads, + num_value_heads, + key_dim, + val_dim, + ); + let row = out_t.to_host(&ctx)?; + out_step_rows.extend_from_slice(&row); + } + + let out_chunk_host = ctx.stream.clone_dtoh(&out_chunk.data)?; + let state_chunk_host = ctx.stream.clone_dtoh(&state_chunk)?; + let state_step_host = ctx.stream.clone_dtoh(&state_step)?; + ctx.sync()?; + let out_chunk_host: Vec = out_chunk_host.iter().map(|x| x.to_f32()).collect(); + + let max_out_diff = out_chunk_host + .iter() + .zip(out_step_rows.iter()) + .map(|(a, b)| (a - b).abs()) + .fold(0.0_f32, f32::max); + let max_state_diff = state_chunk_host + .iter() + .zip(state_step_host.iter()) + .map(|(a, b)| (a - b).abs()) + .fold(0.0_f32, f32::max); + + assert!( + out_chunk_host.iter().all(|x| x.is_finite()) + && state_chunk_host.iter().all(|x| x.is_finite()), + "chunkwise outputs must be finite" + ); + assert!(max_out_diff < 0.05, "output diff {max_out_diff}"); + assert!(max_state_diff < 0.05, "state diff {max_state_diff}"); + Ok(()) +} diff --git a/pegainfer-qwen35/src/scheduler.rs b/pegainfer-qwen35/src/scheduler.rs index c2552c23c..b17029b64 100644 --- a/pegainfer-qwen35/src/scheduler.rs +++ b/pegainfer-qwen35/src/scheduler.rs @@ -825,6 +825,8 @@ impl SingleGpuBackend { } self.graph_state.slot_states[compaction.moved_to].seq_len = self.graph_state.slot_states[compaction.moved_from].seq_len; + #[cfg(feature = "gdn-validation")] + self.graph_state.record_slot_compaction(); match &mut active[compaction.moved_to].backend_state { ActiveBackendState::Single { graph_slot_idx, .. } => { diff --git a/pegainfer-qwen35/src/weights.rs b/pegainfer-qwen35/src/weights.rs index 70cb6bb80..aabede0b1 100644 --- a/pegainfer-qwen35/src/weights.rs +++ b/pegainfer-qwen35/src/weights.rs @@ -107,6 +107,11 @@ impl Default for ModelRuntimeConfig { /// Qwen3.5 model (text-only). pub struct Qwen35Model { pub(super) ctx: DeviceContext, + /// Opaque kernels-owned AOT operation. `None` is an explicit capability + /// fallback (non-SM120 or non-Hv32), never a corrupt-artifact fallback. + pub(super) flashinfer_gdn: Option, + #[cfg(feature = "gdn-validation")] + pub(super) gdn_validation_evidence: super::gdn_validation::GdnValidationEvidenceHandle, pub(super) config: Config35, pub(super) geometry: LocalGeometry, pub(super) embed_tokens: DeviceMatrix, @@ -535,8 +540,40 @@ impl Qwen35Model { num_pages, )?; + // The first production specialization is deliberately single-GPU. + // TP remains an explicit capability fallback to the existing Triton path. + let flashinfer_gdn = if geometry.world_size() == 1 { + pegainfer_kernels::ops::Qwen35GdnAot::load_for_production( + &ctx, + super::flashinfer_gdn::model_geometry(&config), + )? + } else { + None + }; + if let Some(backend) = &flashinfer_gdn { + info!( + "Qwen3.5 GDN production backend: FlashInfer AOT object {}", + backend.artifact_sha256() + ); + } else if geometry.world_size() > 1 { + info!( + "Qwen3.5 GDN production backend: Triton (explicit capability fallback: TP world_size={})", + geometry.world_size() + ); + } else { + let (major, minor) = ctx.ctx.compute_capability()?; + info!( + "Qwen3.5 GDN production backend: Triton (explicit capability fallback: sm_{}{}, geometry={:?})", + major, + minor, + super::flashinfer_gdn::model_geometry(&config) + ); + } Ok(Self { ctx, + flashinfer_gdn, + #[cfg(feature = "gdn-validation")] + gdn_validation_evidence: Default::default(), config, geometry, embed_tokens, @@ -726,13 +763,16 @@ impl Qwen35Model { "requested graph capacity {max_batch} exceeds loaded capacity {}", self.reserved_decode_slots ); - super::batch_decode_graph::BatchDecodeGraphState::with_capacity( + let graph = super::batch_decode_graph::BatchDecodeGraphState::with_capacity( &self.ctx, &self.config, self.geometry, &self.kv_pool, max_batch, - ) + )?; + #[cfg(feature = "gdn-validation")] + let graph = graph.with_validation_evidence(self.gdn_validation_evidence.clone()); + Ok(graph) } pub(crate) fn create_batch_decode_buffers_with_capacity( diff --git a/pegainfer-qwen35/tests/chunked_prefill.rs b/pegainfer-qwen35/tests/chunked_prefill.rs deleted file mode 100644 index dee13c8b2..000000000 --- a/pegainfer-qwen35/tests/chunked_prefill.rs +++ /dev/null @@ -1,128 +0,0 @@ -//! Qwen3.5 scheduler-level chunked prefill regression tests. -//! -//! These tests exercise resumed prefill (`base_pos > 0`) through the real -//! scheduler path. A small `max_prefill_tokens` budget forces one request's -//! prompt to be prefilling across multiple scheduler steps; the same prompt is -//! also run with an effectively unchunked budget and the generated greedy token -//! ids must match. - -use std::path::Path; - -use pegainfer_frontend::engine::EngineHandle; -use pegainfer_frontend::engine::EngineLoadOptions; -use pegainfer_frontend::engine::FinishReason; -use pegainfer_frontend::engine::GenerateRequest; -use pegainfer_frontend::engine::TokenEvent; -use pegainfer_frontend::engine::TokenSink; -use pegainfer_frontend::sampler::SamplingParams; - -mod common; - -const CHUNK_BUDGET: usize = 16; -const BASELINE_PREFILL_BUDGET: usize = 1 << 20; -const MAX_BATCH: usize = 2; -const GENERATED_TOKENS: usize = 8; - -fn start_engine(model_path: &str, max_prefill_tokens: usize) -> EngineHandle { - pegainfer_qwen35::start_engine( - Path::new(model_path), - EngineLoadOptions { - enable_cuda_graph: true, - device_ordinals: vec![0], - seed: 42, - ..EngineLoadOptions::default() - }, - MAX_BATCH, - max_prefill_tokens, - ) - .expect("failed to start Qwen3.5 engine") -} - -fn generate(handle: &EngineHandle, prompt_tokens: Vec) -> (Vec, FinishReason) { - let (token_tx, mut rx) = TokenSink::standalone(); - handle - .submit(GenerateRequest { - trace_parent: None, - request_id: None, - queued_at_unix_s: None, - data_parallel_rank: None, - prompt_tokens, - params: SamplingParams { - ignore_eos: true, - ..SamplingParams::default() - }, - max_tokens: GENERATED_TOKENS, - lora_adapter: None, - kv_transfer_params: None, - token_tx, - logprobs: 0, - echo: false, - }) - .expect("submit failed"); - - let mut tokens = Vec::new(); - loop { - match rx.blocking_recv().map(|(_, event)| event) { - Some(TokenEvent::Token { id, .. }) => tokens.push(id), - Some( - TokenEvent::Scheduled { .. } - | TokenEvent::PromptTokens { .. } - | TokenEvent::KvTransfer { .. }, - ) => {} - Some(TokenEvent::Finished { finish_reason, .. }) => return (tokens, finish_reason), - Some(TokenEvent::Error { message, .. }) => panic!("generation failed: {message}"), - Some(TokenEvent::Rejected { message, .. }) => panic!("generation rejected: {message}"), - None => panic!("scheduler channel closed without Finished"), - } - } -} - -#[test] -fn chunked_prefill_matches_unchunked_prefill_for_resumed_paged_kv() { - let Some(model_path) = common::model_path_or_skip( - "chunked_prefill_matches_unchunked_prefill_for_resumed_paged_kv", - ) else { - return; - }; - let tokenizer = common::load_tokenizer(&model_path); - let prompt = concat!( - "Write a concise technical explanation of paged KV cache updates, ", - "chunked prefill scheduling, and deterministic greedy decoding. ", - "Mention request state ownership, recurrent state, and why resumed ", - "prefill must append K/V instead of overwriting earlier pages. ", - "Then summarize the behavior in three short sentences. ", - "Repeat the explanation with different wording so the prompt is long ", - "enough to cross several small prefill chunks." - ); - let prompt_tokens = tokenizer.encode(prompt, false).expect("encode failed"); - assert!( - prompt_tokens.len() > CHUNK_BUDGET * 2, - "test prompt must force resumed prefill: prompt_len={} chunk_budget={CHUNK_BUDGET}", - prompt_tokens.len() - ); - - let (baseline_tokens, baseline_finish) = { - let handle = start_engine(&model_path, BASELINE_PREFILL_BUDGET); - generate(&handle, prompt_tokens.clone()) - }; - assert_eq!( - baseline_finish, - FinishReason::Length, - "ignore_eos should force baseline generation to the requested length" - ); - - let (chunked_tokens, chunked_finish) = { - let handle = start_engine(&model_path, CHUNK_BUDGET); - generate(&handle, prompt_tokens) - }; - assert_eq!( - chunked_finish, - FinishReason::Length, - "ignore_eos should force chunked generation to the requested length" - ); - - assert_eq!( - chunked_tokens, baseline_tokens, - "chunked prefill must match effectively unchunked prefill; a mismatch suggests resumed direct-paged K/V writes used the wrong base_pos and corrupted earlier cache positions" - ); -} diff --git a/pegainfer-qwen35/tests/e2e_scheduler.rs b/pegainfer-qwen35/tests/e2e_scheduler.rs index 9886399bb..3d9e5719b 100644 --- a/pegainfer-qwen35/tests/e2e_scheduler.rs +++ b/pegainfer-qwen35/tests/e2e_scheduler.rs @@ -580,8 +580,9 @@ fn run_full_scheduler_e2e( } } - // ── 4b. Mixed concurrent logprobs requests ───────────────────────── - info!("=== Phase 4b: Mixed concurrent logprobs ==="); + // ── 4a. Mixed concurrent logprobs requests + info!("=== Phase 4a: Mixed concurrent logprobs ==="); + { let mixed = [ ("mixed_no_logprobs", CASES[0].prompt, 0usize), @@ -664,6 +665,50 @@ fn run_full_scheduler_e2e( info!("All Qwen3.5 scheduler tests passed for {label}!"); } +fn run_graph_lifecycle_boundary(handle: &EngineHandle, tokenizer: &DynTokenizer) { + let run_batch = |cases: &[(&str, &str, usize)]| { + let mut receivers = Vec::with_capacity(cases.len()); + for &(name, prompt, max_tokens) in cases { + let prompt_tokens = tokenizer.encode(prompt, false).expect("encode failed"); + let (token_tx, token_rx) = TokenSink::standalone(); + handle + .submit(GenerateRequest { + trace_parent: None, + request_id: Some(name.to_string()), + queued_at_unix_s: None, + data_parallel_rank: None, + prompt_tokens, + params: SamplingParams { + ignore_eos: true, + ..SamplingParams::default() + }, + max_tokens, + lora_adapter: None, + kv_transfer_params: None, + token_tx, + logprobs: 0, + echo: false, + }) + .expect("submit graph-lifecycle request"); + receivers.push((name, max_tokens, token_rx)); + } + for (name, max_tokens, mut receiver) in receivers { + let result = collect_generation(&mut receiver, name, 0); + assert_eq!(result.finish_reason, FinishReason::Length); + assert_eq!(result.tokens.len(), max_tokens); + } + }; + + // The first row retires while two longer rows remain, forcing compaction. + // Multiple decode steps also force replay after the first capture. + run_batch(&[ + ("compact-short", "A short request", 8), + ("compact-long-a", "A longer request about CUDA graphs", 24), + ("compact-long-b", "Another longer request about state", 24), + ]); + // A second wave must copy into a previously occupied stable slot. + run_batch(&[("reuse-slot", "Reuse the graph slot", 3)]); +} #[test] fn test_e2e_qwen35_scheduler() { let Some(model_path) = common::model_path_or_skip("test_e2e_qwen35_scheduler") else { @@ -690,6 +735,88 @@ fn test_e2e_qwen35_scheduler() { run_full_scheduler_e2e(&handle, &tokenizer, max_context_tokens, "TP1"); } +#[cfg(feature = "gdn-validation")] +#[test] +#[ignore = "requires an SM120 GPU, Qwen3.5-4B weights, and the validated Hv32 FlashInfer artifact"] +fn test_e2e_qwen35_scheduler_flashinfer_gdn() { + let model_path = std::env::var("PEGAINFER_TEST_MODEL_PATH").expect( + "required FlashInfer scheduler gate needs PEGAINFER_TEST_MODEL_PATH set to the pinned Qwen3.5-4B snapshot", + ); + assert!( + Path::new(&model_path).join("config.json").is_file(), + "required FlashInfer scheduler gate cannot read {model_path}/config.json; set PEGAINFER_TEST_MODEL_PATH" + ); + info!("Loading Qwen3.5 model for FlashInfer scheduler test..."); + let start = Instant::now(); + let tokenizer = common::load_tokenizer(&model_path); + let (handle, evidence) = + pegainfer_qwen35::runtime::start_engine_with_flashinfer_gdn_for_accuracy( + Path::new(&model_path), + 0, + 8, + pegainfer_qwen35::DEFAULT_MAX_PREFILL_TOKENS, + ) + .expect("Failed to start FlashInfer Qwen3.5 scheduler"); + let initial = evidence.snapshot(); + assert_eq!(initial.selected_backend, "flashinfer"); + assert_ne!(initial.artifact_sha256, "unavailable"); + assert_eq!(initial.artifact_sha256.len(), 64); + assert_eq!(initial.successful_launches, 0); + assert_eq!(initial.graph_captures, 0); + assert_eq!(initial.graph_replays, 0); + assert_eq!(initial.graph_eager_fallbacks, 0); + assert_eq!(initial.state_slot_copies, 0); + assert_eq!(initial.state_slot_reuses, 0); + assert_eq!(initial.slot_compactions, 0); + info!( + "FlashInfer identity: object_sha256={}", + initial.artifact_sha256 + ); + info!("FlashInfer scheduler loaded in {:.2?}", start.elapsed()); + + run_graph_lifecycle_boundary(&handle, &tokenizer); + let final_evidence = evidence.snapshot(); + assert_eq!(final_evidence.selected_backend, "flashinfer"); + assert_eq!(final_evidence.artifact_sha256, initial.artifact_sha256); + assert!( + final_evidence.successful_launches > 0, + "scheduler e2e completed without a successful FlashInfer GDN launch" + ); + assert!( + final_evidence.graph_captures >= 1, + "scheduler E2E did not capture any CUDA decode graph" + ); + assert!( + final_evidence.graph_replays >= 1, + "scheduler E2E did not replay a captured CUDA decode graph" + ); + assert_eq!( + final_evidence.graph_eager_fallbacks, 0, + "scheduler E2E silently used the eager decode fallback" + ); + assert!( + final_evidence.state_slot_copies >= 1, + "scheduler E2E did not copy prefill recurrent state into a graph slot" + ); + assert!( + final_evidence.state_slot_reuses >= 1, + "scheduler E2E did not reuse a stable graph slot" + ); + assert!( + final_evidence.slot_compactions >= 1, + "scheduler E2E did not exercise graph-slot compaction" + ); + info!( + "FlashInfer scheduler evidence: launches={} graph_captures={} graph_replays={} state_slot_copies={} state_slot_reuses={} slot_compactions={}", + final_evidence.successful_launches, + final_evidence.graph_captures, + final_evidence.graph_replays, + final_evidence.state_slot_copies, + final_evidence.state_slot_reuses, + final_evidence.slot_compactions, + ); +} + #[test] fn test_e2e_qwen35_shared_sm_last_decoder() { pegainfer_core::logging::init_default(); diff --git a/pegainfer-qwen35/tests/hf_golden_gate.rs b/pegainfer-qwen35/tests/hf_golden_gate.rs index 558866a6f..06e31f75b 100644 --- a/pegainfer-qwen35/tests/hf_golden_gate.rs +++ b/pegainfer-qwen35/tests/hf_golden_gate.rs @@ -94,6 +94,18 @@ const BUCKET_STRADDLES: [usize; 2] = [5, 3]; const SLOT_COMPACTION_BATCH: usize = 5; const SLOT_COMPACTION_DROP_INDEX: usize = 1; +fn required_model_path() -> String { + let path = std::env::var("PEGAINFER_TEST_MODEL_PATH").expect( + "required Qwen3.5 production gate needs PEGAINFER_TEST_MODEL_PATH set to the pinned Qwen3.5-4B snapshot", + ); + let config = Path::new(&path).join("config.json"); + assert!( + config.is_file(), + "required Qwen3.5 production gate cannot read {}; set PEGAINFER_TEST_MODEL_PATH to the pinned Qwen3.5-4B snapshot", + config.display() + ); + path +} fn sha256_file(path: impl AsRef) -> Option { let bytes = std::fs::read(path).ok()?; let mut digest = Sha256::new(); @@ -834,6 +846,55 @@ fn pega_logprobs_match_hf_long_golden_within_qwen35_tolerance() { ); } +#[cfg(feature = "gdn-validation")] +#[test] +#[ignore = "requires an SM120 GPU, Qwen3.5-4B weights, and the validated Hv32 FlashInfer artifact"] +fn production_flashinfer_gdn_matches_hf_short_golden() { + let model_path = required_model_path(); + assert_eq!( + fixture_size_name(&model_path), + Some("4b"), + "FlashInfer production HF gate is scoped to the Qwen3.5-4B Hv32 geometry" + ); + let golden = Golden::load_for(&model_path, false) + .expect("required Qwen3.5-4B HF fixture is missing or unrecognized"); + assert!( + check_fixture_metadata(&model_path, &golden), + "required Qwen3.5 production gate could not prove the pinned model revision" + ); + report_fixture_shape(&golden); + let all = (0..golden.num_seqs).collect::>(); + let mut production = build_executor(&model_path); + let production_before = production + .flashinfer_gdn_runtime_evidence() + .expect("SM120/Hv32 production Auto dispatch must expose FlashInfer evidence"); + assert_eq!(production_before.selected_backend, "flashinfer"); + assert_ne!(production_before.artifact_sha256, "unavailable"); + assert_eq!(production_before.artifact_sha256.len(), 64); + assert_eq!(production_before.successful_launches, 0); + let (production_stats, _) = run(&golden, &mut production, &all, false); + report_and_assert("production Auto sequential bs=1 graph", &production_stats); + let production_after = production + .flashinfer_gdn_runtime_evidence() + .expect("production Auto dispatch lost FlashInfer identity"); + assert_eq!( + production_after.artifact_sha256, + production_before.artifact_sha256 + ); + assert_eq!(production_after.selected_backend, "flashinfer"); + assert!( + production_after.successful_launches > production_before.successful_launches, + "production Auto HF replay completed without a FlashInfer launch" + ); + eprintln!( + "qwen35 hf_golden_gate [production Auto]: selected_backend={} object_sha256={} successful_launches={} -> {}", + production_after.selected_backend, + production_after.artifact_sha256, + production_before.successful_launches, + production_after.successful_launches, + ); +} + #[test] #[ignore = "requires two CUDA devices, NCCL, and Qwen3.5 weights"] fn pega_logprobs_match_hf_golden_within_qwen35_tolerance_tp2() { diff --git a/pegainfer-qwen35/tools/run_gdn_production_gates.sh b/pegainfer-qwen35/tools/run_gdn_production_gates.sh new file mode 100755 index 000000000..130b8eded --- /dev/null +++ b/pegainfer-qwen35/tools/run_gdn_production_gates.sh @@ -0,0 +1,188 @@ +#!/usr/bin/env bash +set -euo pipefail + +repo_root="$(git rev-parse --show-toplevel)" +cd "$repo_root" + +require_env() { + local name="$1" + if [[ -z "${!name:-}" ]]; then + echo "required environment variable is missing: $name" >&2 + exit 2 + fi +} + +require_env PEGAINFER_QWEN35_GDN_AOT_BUNDLE +require_env PEGAINFER_TEST_MODEL_PATH +require_env PEGAINFER_TEST_MODEL_REVISION +require_env PEGAINFER_TRITON_PYTHON +require_env PEGAINFER_CUDA_SM +require_env CARGO_TARGET_DIR + +expected_revision="851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a" +expected_config_sha="ddc63e1c717afa86c865bb5e01313d89d72bb53b97ad4a8a03ba8510c0621670" +bundle="$PEGAINFER_QWEN35_GDN_AOT_BUNDLE" +model="$PEGAINFER_TEST_MODEL_PATH" +python="$PEGAINFER_TRITON_PYTHON" +log_root="${PEGAINFER_GDN_GATE_LOG_DIR:-$repo_root/target/gdn-production-gate-logs}" + +mkdir -p "$log_root" +echo "GDN gate log root: $log_root" + +for command in \ + git nvidia-smi nvcc rustc cargo protoc cc c++ clang cmake ninja pkg-config \ + rg sha256sum awk sed tee timeout; do + if ! command -v "$command" >/dev/null 2>&1; then + echo "required GDN gate command is missing: $command" >&2 + exit 2 + fi +done + +if [[ "$PEGAINFER_CUDA_SM" != "120" ]]; then + echo "GDN production gates require PEGAINFER_CUDA_SM=120, got $PEGAINFER_CUDA_SM" >&2 + exit 2 +fi + +if [[ -n "${PEGAINFER_GDN_EXPECT_BRANCH:-}" ]]; then + actual_branch="$(git branch --show-current)" + if [[ "$actual_branch" != "$PEGAINFER_GDN_EXPECT_BRANCH" ]]; then + echo "GDN gate branch mismatch: expected $PEGAINFER_GDN_EXPECT_BRANCH, got $actual_branch" >&2 + exit 2 + fi +fi + +if [[ -n "$(git status --short --untracked-files=no)" ]]; then + echo "GDN production gates require a clean tracked working tree" >&2 + exit 2 +fi + +gpu_compute_cap="$(nvidia-smi --query-gpu=compute_cap --format=csv,noheader | sed -n '1p' | tr -d '[:space:]')" +if [[ "$gpu_compute_cap" != "12.0" ]]; then + echo "GDN production gates require compute capability 12.0, got $gpu_compute_cap" >&2 + exit 2 +fi + +if [[ "$PEGAINFER_TEST_MODEL_REVISION" != "$expected_revision" ]]; then + echo "model revision mismatch: expected $expected_revision, got $PEGAINFER_TEST_MODEL_REVISION" >&2 + exit 2 +fi + +for required in \ + "$model/config.json" \ + "$bundle/manifest.json" \ + "$bundle/kernel.o" \ + "$python"; do + if [[ ! -f "$required" ]]; then + echo "required GDN gate input is missing: $required" >&2 + exit 2 + fi +done +if [[ ! -x "$python" ]]; then + echo "GDN gate Python is not executable: $python" >&2 + exit 2 +fi + +actual_config_sha="$(sha256sum "$model/config.json" | awk '{print $1}')" +if [[ "$actual_config_sha" != "$expected_config_sha" ]]; then + echo "model config SHA mismatch: expected $expected_config_sha, got $actual_config_sha" >&2 + exit 2 +fi + +"$python" pegainfer-kernels/tools/flashinfer_gdn/artifact_contract.py \ + validate-candidate "$bundle" \ + --flashinfer-dir pegainfer-kernels/third_party/flashinfer + +commit_sha="$(git rev-parse HEAD)" +submodule_sha="$(git -C pegainfer-kernels/third_party/flashinfer rev-parse HEAD)" +manifest_sha="$(sha256sum "$bundle/manifest.json" | awk '{print $1}')" +object_sha="$(sha256sum "$bundle/kernel.o" | awk '{print $1}')" +{ + echo "commit_sha=$commit_sha" + echo "branch=$(git branch --show-current)" + echo "expected_branch=${PEGAINFER_GDN_EXPECT_BRANCH:-not-enforced}" + echo "flashinfer_submodule_sha=$submodule_sha" + echo "model_revision=$PEGAINFER_TEST_MODEL_REVISION" + echo "model_config_sha256=$actual_config_sha" + echo "manifest_sha256=$manifest_sha" + echo "object_sha256=$object_sha" + echo "gpu_compute_cap=$gpu_compute_cap" + nvidia-smi --query-gpu=name,driver_version,memory.total --format=csv,noheader + nvcc --version + rustc --version + cargo --version + protoc --version + clang --version | sed -n '1p' + cmake --version | sed -n '1p' + ninja --version + "$python" --version +} | tee "$log_root/provenance.log" + +timeout 90m cargo build --release --locked \ + -p pegainfer-server \ + --no-default-features \ + --features qwen35 \ + --bin pegainfer 2>&1 | tee "$log_root/production-build.log" + +timeout 60m cargo test --release --locked \ + -p pegainfer-kernels --features qwen35 --lib --no-run \ + 2>&1 | tee "$log_root/kernels-tests-build.log" +timeout 60m cargo test --release --locked \ + -p pegainfer-qwen35 --features qwen35 --lib --tests --no-run \ + 2>&1 | tee "$log_root/qwen35-default-tests-build.log" +timeout 60m cargo test --release --locked \ + -p pegainfer-qwen35 --features qwen35,gdn-validation --lib --tests --no-run \ + 2>&1 | tee "$log_root/qwen35-validation-tests-build.log" + +run_exact_gate() { + local label="$1" + local exact_name="$2" + shift 2 + local list_log="$log_root/$label-list.log" + local run_log="$log_root/$label.log" + + timeout 60m cargo test --release --locked "$@" "$exact_name" \ + -- --ignored --exact --list >"$list_log" 2>&1 + local listed + listed="$(rg -c "^${exact_name}: test$" "$list_log" || true)" + if [[ "$listed" != "1" ]]; then + echo "$label exact filter matched $listed tests, expected 1" >&2 + sed -n '1,160p' "$list_log" >&2 + exit 3 + fi + + timeout 60m cargo test --release --locked "$@" "$exact_name" \ + -- --ignored --exact --nocapture 2>&1 | tee "$run_log" + local passed + passed="$(rg -c "^test ${exact_name} \.\.\. ok$" "$run_log" || true)" + if [[ "$passed" != "1" ]]; then + echo "$label executed-pass count was $passed, expected 1" >&2 + exit 3 + fi +} + +run_exact_gate \ + gate1-real-aot-boundary \ + ops::qwen35::tests::sm120_stable_abi_alias_and_separate_state_are_bitwise_identical \ + -p pegainfer-kernels --features qwen35 --lib + +run_exact_gate \ + gate2-native-prepare-cpu-oracle \ + recurrent::tests::native_prepare_hv32_dynamic_t_and_non_finite_inputs \ + -p pegainfer-qwen35 --features qwen35,gdn-validation --lib + +run_exact_gate \ + gate3-production-hf-golden \ + production_flashinfer_gdn_matches_hf_short_golden \ + -p pegainfer-qwen35 --features qwen35,gdn-validation --test hf_golden_gate + +run_exact_gate \ + gate4-chunk-continuation \ + prefill::tests::flashinfer_gdn_chunk_continuation_and_model_outputs_match \ + -p pegainfer-qwen35 --features qwen35,gdn-validation --lib + +run_exact_gate \ + gate5-scheduler-cuda-graph \ + test_e2e_qwen35_scheduler_flashinfer_gdn \ + -p pegainfer-qwen35 --features qwen35,gdn-validation --test e2e_scheduler + +echo "all five Qwen3.5 GDN production gates passed for $commit_sha object $object_sha"