diff --git a/recipe/cua_s1/native.md b/recipe/cua_s1/native.md index 0e352007..83322b25 100644 --- a/recipe/cua_s1/native.md +++ b/recipe/cua_s1/native.md @@ -28,7 +28,7 @@ each exact prompt length warms the GEMM plans and captures the forward pass; later requests replay it with freshly uploaded token ids. At most eight lengths are cached. Growing the scratch allocation clears the captures before freeing their buffers. Capture adds first-use latency; leave the variable unset to use -the eager control. Rebuild both the worker and CUDA library together (ABI 4). +the eager control. Rebuild both the worker and CUDA library together (ABI 6). If capture fails, the worker returns the completed eager result and disables Graph capture/replay for its remaining lifetime, logging the failure to stderr. diff --git a/recipe/open_jev/native.md b/recipe/open_jev/native.md index 7537d48f..574a43fa 100644 --- a/recipe/open_jev/native.md +++ b/recipe/open_jev/native.md @@ -44,9 +44,7 @@ or an incomplete export. The saved limit defaults to 4096 tokens per candidate; The CUDA kernels require compute capability 8.0 or newer. The current build target below is Ada (`89`); pass your GPU's compute capability explicitly. -The CUDA shared library and both Rust workers must be rebuilt together because -the gated-attention entry point updates the library ABI to version 4 alongside -the shared CUDA Graph entry points. +The CUDA shared library and both Rust workers must be rebuilt together (ABI 6). ```sh src/backends/cuda/qwen3_5/build.sh target/release 89 diff --git a/src/backends/cuda/contract.md b/src/backends/cuda/contract.md index 41301af1..e7c01b2d 100644 --- a/src/backends/cuda/contract.md +++ b/src/backends/cuda/contract.md @@ -49,7 +49,7 @@ change. Schema: ```json { "name": "qwen3_5", - "abi_version": 4, + "abi_version": 6, "status": "validated", "sources": ["common.cuh", "mma.cuh", "ops.h", "norm.cu", "elementwise.cu", "attention.cu", "gdn_prefill.cu", "gemm.cu", "runtime.cu"], @@ -136,7 +136,7 @@ reference for the shape. The contract fixes four points: whose value it does not know. The manifest's `abi_version` is **that number**, not a version of this document, and the checker reads the macro out of the declared sources and requires the two to agree. `#19`'s `ops.h` says - `CS1_ABI_VERSION 4` today, so a `qwen3_5` manifest declares 4. + `CS1_ABI_VERSION 6` today, so a `qwen3_5` manifest declares 6. Two backends may declare different values. They are independent libraries, and the repository layout says CUDA and Metal implementations need not share diff --git a/src/backends/cuda/qwen3_5/README.md b/src/backends/cuda/qwen3_5/README.md index 27a9f9b3..ff98b3c0 100644 --- a/src/backends/cuda/qwen3_5/README.md +++ b/src/backends/cuda/qwen3_5/README.md @@ -11,7 +11,29 @@ The norm, elementwise and q/k preparation kernels round to bfloat16 where Transf `cs1_attention_gated` fuses the sigmoid gate into the attention epilogue, preserving the BF16 rounding of both attention and sigmoid before multiplication. The native workers use this entry point; the separate operations remain available for kernel -comparisons. Rebuild the library and workers together for ABI version 4, which -includes the CUDA Graph entry points and gated attention. +comparisons. Rebuild the library and workers together for ABI version 6, which +includes the CUDA Graph entry points, gated attention, the vision operations and +the continuation operations below. + +Three operations continue a sequence after a shared prefix, for prefix reuse +([#85](https://github.com/ThinkFlowLab/system1-omni/issues/85)); the native +workers do not call them yet: + +- `cs1_gdn_conv_history` reads the conv inputs of the three positions before its + first token and can write those of its last three. Every output equals the + unsplit conv's. +- `cs1_gdn_prefill_state` starts the chunked gated delta rule from a float32 state + `[H, 128, 128]` and can write the final state. When every split falls on a multiple + of 64 tokens, it reproduces the unsplit prefill bit for bit; elsewhere the chunks + fall differently, within the float64-reference tolerance of the unsplit kernel. +- `cs1_attention_gated_cached` runs the queries of the last positions against keys + and values that also cover the positions before them. Key tiles start at position + 0 either way, so each output row matches the unsplit call. + +`cs1_copy_rows` queues a pitched device-to-device copy, for keeping cached values +in their own rows. The existing `cs1_gdn_conv`, `cs1_gdn_prefill` and +`cs1_attention_gated` are these operations without history, state or cached +positions. `tests/qwen3_5/kernels.rs` checks each against the unsplit call at +prefix lengths around and inside 64-token chunks. Gated DeltaNet preparation stores converted TF32 operands in three-byte component planes, preserves the original four-term TF32 accumulation, and writes U/W fragments directly as bfloat16. Dynamic shared memory is 72 KiB per block. The [H200 comparison](../../../../benchmarks/gdn/README.md) records complete GDN call latency, numerical checks, and the small end-to-end change measured with the Open-Jev worker from PR #55. diff --git a/src/backends/cuda/qwen3_5/attention.cu b/src/backends/cuda/qwen3_5/attention.cu index 20ffceed..162943db 100644 --- a/src/backends/cuda/qwen3_5/attention.cu +++ b/src/backends/cuda/qwen3_5/attention.cu @@ -5,7 +5,9 @@ // four warps of 16 rows each, and walks the keys up to its last query in tiles of // 32, keeping the output and the online softmax in registers. The probabilities are // rounded to bfloat16 for the P*V product, as in flash attention; the running sums -// stay float32. +// stay float32. The queries can be the last Tq of Tk positions (cached keys before +// them); key tiles always start at position 0, so a query sees the same tiles in the +// same order wherever the queries start, and tiles past its own position change nothing. #include "common.cuh" #include "mma.cuh" #include "ops.h" @@ -81,7 +83,7 @@ constexpr int SMEM_BYTES = (BM + 2 * BN) * LDS * 2; template __global__ void __launch_bounds__(THREADS) flash_kernel(const bf16* __restrict__ q, const bf16* __restrict__ k, const bf16* __restrict__ v, int ldv, - const bf16* __restrict__ gate, bf16* __restrict__ out, int T, int Hq, int Hk, + const bf16* __restrict__ gate, bf16* __restrict__ out, int Tq, int Tk, int Hq, int Hk, float scale_log2) { extern __shared__ __align__(16) unsigned char smem[]; bf16* qs = reinterpret_cast(smem); @@ -92,10 +94,11 @@ __global__ void __launch_bounds__(THREADS) const int tid = threadIdx.x, warp = tid / 32, lane = tid % 32; const int g = lane / 4, t = lane % 4; const int row0 = q0 + warp * 16; // this warp's first query + const int offset = Tk - Tq; // the position of query 0 for (int c = tid; c < BM * (D / 8); c += THREADS) { const int r = c / (D / 8), col = (c % (D / 8)) * 8, row = q0 + r; - cp_async16(qs + r * LDS + col, q + ((size_t)min(row, T - 1) * Hq + h) * D + col, row < T); + cp_async16(qs + r * LDS + col, q + ((size_t)min(row, Tq - 1) * Hq + h) * D + col, row < Tq); } cp_async_commit(); @@ -104,23 +107,24 @@ __global__ void __launch_bounds__(THREADS) for (int n = 0; n < D / 8; n++) o[n][0] = o[n][1] = o[n][2] = o[n][3] = 0.f; float m[2] = {-INFINITY, -INFINITY}, l[2] = {0.f, 0.f}; - const int kv_end = min(T, q0 + BM); + const int kv_end = min(Tk, offset + q0 + BM); for (int k0 = 0; k0 < kv_end; k0 += BN) { for (int c = tid; c < BN * (D / 8); c += THREADS) { const int r = c / (D / 8), col = (c % (D / 8)) * 8, s = k0 + r; - cp_async16(ks + r * LDS + col, k + ((size_t)min(s, T - 1) * Hk + hk) * D + col, s < T); + cp_async16(ks + r * LDS + col, k + ((size_t)min(s, Tk - 1) * Hk + hk) * D + col, s < Tk); } cp_async_commit(); for (int c = tid; c < BN * (D / 8); c += THREADS) { const int r = c / (D / 8), col = (c % (D / 8)) * 8, s = k0 + r; - cp_async16(vs + r * LDS + col, v + (size_t)min(s, T - 1) * ldv + (size_t)hk * D + col, s < T); + cp_async16(vs + r * LDS + col, v + (size_t)min(s, Tk - 1) * ldv + (size_t)hk * D + col, s < Tk); } cp_async_commit(); cp_async_wait<1>(); // Q and K __syncthreads(); - // keys past every query of this warp contribute nothing - const bool active = k0 <= row0 + 15; + // keys past every query of this warp contribute nothing, and rows past the + // last query are never stored + const bool active = row0 < Tq && k0 <= offset + row0 + 15; float sc[BN / 8][4]; #pragma unroll for (int n = 0; n < BN / 8; n++) sc[n][0] = sc[n][1] = sc[n][2] = sc[n][3] = 0.f; @@ -146,8 +150,8 @@ __global__ void __launch_bounds__(THREADS) for (int n = 0; n < BN / 8; n++) { #pragma unroll for (int e = 0; e < 4; e++) { - const int key = k0 + n * 8 + 2 * t + (e & 1), row = row0 + g + (e >> 1) * 8; - sc[n][e] = (key <= row && key < T) ? sc[n][e] * scale_log2 : -INFINITY; + const int key = k0 + n * 8 + 2 * t + (e & 1), pos = offset + row0 + g + (e >> 1) * 8; + sc[n][e] = (key <= pos && key < Tk) ? sc[n][e] * scale_log2 : -INFINITY; mx[e >> 1] = fmaxf(mx[e >> 1], sc[n][e]); } } @@ -213,7 +217,7 @@ __global__ void __launch_bounds__(THREADS) #pragma unroll for (int r = 0; r < 2; r++) { const int row = row0 + g + r * 8; - if (row >= T) continue; + if (row >= Tq) continue; bf16* dst = out + ((size_t)row * Hq + h) * D + 2 * t; #pragma unroll for (int n = 0; n < D / 8; n++) { @@ -232,20 +236,20 @@ __global__ void __launch_bounds__(THREADS) template int launch(const void* q, const void* k, const void* v, int ldv, const void* gate, void* out, - int T, int Hq, int Hk, int Dh, float scale, void* stream) { - if (Dh != D || Hk <= 0 || Hq <= 0 || Hq % Hk != 0 || ldv % 8 != 0 || ldv < Hk * Dh || T < 0) + int Tq, int Tk, int Hq, int Hk, int Dh, float scale, void* stream) { + if (Dh != D || Hk <= 0 || Hq <= 0 || Hq % Hk != 0 || ldv % 8 != 0 || ldv < Hk * Dh || Tq < 0 || Tk < Tq) return cudaErrorInvalidValue; - if (T == 0) return cudaSuccess; + if (Tq == 0) return cudaSuccess; if (Gated && gate == nullptr) return cudaErrorInvalidValue; // Once per specialization (for the device current at the first call). static const cudaError_t configured = cudaFuncSetAttribute( flash_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_BYTES); if (configured != cudaSuccess) return configured; constexpr float LOG2E = 1.4426950408889634f; - flash_kernel<<<<(stream)>>>( static_cast(q), static_cast(k), static_cast(v), ldv, - static_cast(gate), static_cast(out), T, Hq, Hk, scale * LOG2E); + static_cast(gate), static_cast(out), Tq, Tk, Hq, Hk, scale * LOG2E); return cudaGetLastError(); } @@ -274,10 +278,16 @@ extern "C" int cs1_attn_prep(const void* qg, const void* kr, int ld, const void* extern "C" int cs1_attention(const void* q, const void* k, const void* v, int ldv, void* out, int T, int Hq, int Hk, int Dh, float scale, void* stream) { - return flash::launch(q, k, v, ldv, nullptr, out, T, Hq, Hk, Dh, scale, stream); + return flash::launch(q, k, v, ldv, nullptr, out, T, T, Hq, Hk, Dh, scale, stream); } extern "C" int cs1_attention_gated(const void* q, const void* k, const void* v, int ldv, const void* gate, void* out, int T, int Hq, int Hk, int Dh, float scale, void* stream) { - return flash::launch(q, k, v, ldv, gate, out, T, Hq, Hk, Dh, scale, stream); + return flash::launch(q, k, v, ldv, gate, out, T, T, Hq, Hk, Dh, scale, stream); +} + +extern "C" int cs1_attention_gated_cached(const void* q, const void* k, const void* v, int ldv, const void* gate, + void* out, int Tq, int Tk, int Hq, int Hk, int Dh, float scale, + void* stream) { + return flash::launch(q, k, v, ldv, gate, out, Tq, Tk, Hq, Hk, Dh, scale, stream); } diff --git a/src/backends/cuda/qwen3_5/elementwise.cu b/src/backends/cuda/qwen3_5/elementwise.cu index 97ae9c5a..2d9881ed 100644 --- a/src/backends/cuda/qwen3_5/elementwise.cu +++ b/src/backends/cuda/qwen3_5/elementwise.cu @@ -16,10 +16,13 @@ __global__ void embed_kernel(const int32_t* __restrict__ ids, const Pack8* __res } // F.conv1d in bfloat16 (float32 accumulation, rounded), then SiLU (rounded again), -// written to three contiguous outputs. +// written to three contiguous outputs. Positions before the first row come from +// history [3, channels] when there is one; without it they are skipped, as the zero +// padding of a sequence's start adds nothing. __global__ void gdn_conv_kernel(const bf16* __restrict__ qkv, int ld, const bf16* __restrict__ w, - bf16* __restrict__ q, bf16* __restrict__ k, bf16* __restrict__ v, int T, - int key_dim, int value_dim) { + const bf16* __restrict__ history, bf16* __restrict__ q, + bf16* __restrict__ k, bf16* __restrict__ v, int T, int key_dim, + int value_dim) { const int channels = 2 * key_dim + value_dim; const size_t idx = (size_t)blockIdx.x * blockDim.x + threadIdx.x; if (idx >= (size_t)T * channels) return; @@ -28,7 +31,10 @@ __global__ void gdn_conv_kernel(const bf16* __restrict__ qkv, int ld, const bf16 #pragma unroll for (int j = 0; j < 4; j++) { const int s = t - 3 + j; - if (s >= 0) acc = fmaf(f32(w[c * 4 + j]), f32(qkv[(size_t)s * ld + c]), acc); + if (s >= 0) + acc = fmaf(f32(w[c * 4 + j]), f32(qkv[(size_t)s * ld + c]), acc); + else if (history) + acc = fmaf(f32(w[c * 4 + j]), f32(history[(size_t)(3 + s) * channels + c]), acc); } const bf16 y = to_bf16(silu(round_bf16(acc))); if (c < key_dim) @@ -39,6 +45,16 @@ __global__ void gdn_conv_kernel(const bf16* __restrict__ qkv, int ld, const bf16 v[(size_t)t * value_dim + c - 2 * key_dim] = y; } +// The conv inputs of the last three positions, from qkv or, before its first row, from +// history (zeros without one). +__global__ void gdn_conv_history_kernel(const bf16* __restrict__ qkv, int ld, const bf16* __restrict__ history, + bf16* __restrict__ out, int T, int channels) { + const int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= 3 * channels) return; + const int r = idx / channels, c = idx % channels, s = T - 3 + r; + out[idx] = s >= 0 ? qkv[(size_t)s * ld + c] : history ? history[(3 + s) * channels + c] : to_bf16(0.f); +} + // beta = sigmoid(b) in bfloat16; g = -exp(A_log) * softplus(a + dt_bias) in float32 // (F.softplus with threshold 20). __global__ void gdn_gates_kernel(const bf16* __restrict__ b, const bf16* __restrict__ a, int ld, @@ -101,17 +117,30 @@ extern "C" int cs1_embed(const int32_t* ids, const void* table, void* out, int T return cudaGetLastError(); } -extern "C" int cs1_gdn_conv(const void* qkv, int ld, const void* w, void* q, void* k, void* v, int T, int key_dim, - int value_dim, void* stream) { +extern "C" int cs1_gdn_conv_history(const void* qkv, int ld, const void* w, const void* history, void* history_out, + void* q, void* k, void* v, int T, int key_dim, int value_dim, void* stream) { if (T < 0 || key_dim < 0 || value_dim < 0 || ld < 2 * key_dim + value_dim) return cudaErrorInvalidValue; - const size_t n = (size_t)T * (2 * key_dim + value_dim); - if (n == 0) return cudaSuccess; - gdn_conv_kernel<<(stream)>>>( - static_cast(qkv), ld, static_cast(w), static_cast(q), - static_cast(k), static_cast(v), T, key_dim, value_dim); + const int channels = 2 * key_dim + value_dim; + const size_t n = (size_t)T * channels; + const cudaStream_t st = static_cast(stream); + if (n > 0) { + gdn_conv_kernel<<>>( + static_cast(qkv), ld, static_cast(w), static_cast(history), + static_cast(q), static_cast(k), static_cast(v), T, key_dim, value_dim); + } + if (history_out && channels > 0) { + gdn_conv_history_kernel<<>>( + static_cast(qkv), ld, static_cast(history), static_cast(history_out), + T, channels); + } return cudaGetLastError(); } +extern "C" int cs1_gdn_conv(const void* qkv, int ld, const void* w, void* q, void* k, void* v, int T, int key_dim, + int value_dim, void* stream) { + return cs1_gdn_conv_history(qkv, ld, w, nullptr, nullptr, q, k, v, T, key_dim, value_dim, stream); +} + extern "C" int cs1_gdn_gates(const void* b, const void* a, int ld, const void* A_log, const void* dt_bias, void* beta, float* g, int T, int H, void* stream) { if (T < 0 || H < 0 || ld < H) return cudaErrorInvalidValue; diff --git a/src/backends/cuda/qwen3_5/gdn_prefill.cu b/src/backends/cuda/qwen3_5/gdn_prefill.cu index 676ce446..67624dbb 100644 --- a/src/backends/cuda/qwen3_5/gdn_prefill.cu +++ b/src/backends/cuda/qwen3_5/gdn_prefill.cu @@ -8,7 +8,8 @@ // (bfloat16 with float32 accumulation). Results are stored as bfloat16. // 2. gdn_chunk_state, per (head, 32 value columns), over the chunks in order: keeps // the state S in float32 registers, stores it as bfloat16 before each chunk, and -// computes v_new = u - w S and S = decay S + kd^T v_new with mma.sync. +// computes v_new = u - w S and S = decay S + kd^T v_new with mma.sync. S starts +// from zero or from a given float32 state, and the final one can be written out. // 3. gdn_chunk_out, per (chunk, head): o = qd S + P v_new with mma.sync. // Transformers computes all of this in float32. Keeping the intermediate results in // bfloat16, as flash-linear-attention does, makes this kernel less precise than that @@ -337,7 +338,8 @@ constexpr int SS_LD = BVS + 8; // bfloat16 row stride of the S copy and v_n constexpr int STAGE = C * WS_LD; // elements of one staged w or kd constexpr size_t SMEM2_BYTES = (4 * STAGE + K * SS_LD + C * SS_LD) * 2; -__global__ void __launch_bounds__(ST_THREADS) gdn_chunk_state(Work ws, int NC) { +__global__ void __launch_bounds__(ST_THREADS) + gdn_chunk_state(Work ws, int NC, const float* s0, float* s_out) { // s0 and s_out may alias extern __shared__ __align__(128) unsigned char sm[]; bf16* wbuf = reinterpret_cast(sm); // [2][C][WS_LD] bf16* kbuf = wbuf + 2 * STAGE; // [2][C][WS_LD] @@ -356,12 +358,22 @@ __global__ void __launch_bounds__(ST_THREADS) gdn_chunk_state(Work ws, int NC) { cs1::cp_async_commit(); }; - // S rows warp * 32 + mt * 16 + {g, g + 8}, columns nt * 8 + {2t, 2t + 1} + // S rows warp * 32 + mt * 16 + {g, g + 8}, columns nt * 8 + {2t, 2t + 1}; the given + // state is float [H, K, V], read and written by the block that owns its columns float st[2][4][4]; + const size_t s_at = (size_t)h * K * V + vb0; #pragma unroll for (int mt = 0; mt < 2; mt++) #pragma unroll - for (int nt = 0; nt < 4; nt++) st[mt][nt][0] = st[mt][nt][1] = st[mt][nt][2] = st[mt][nt][3] = 0.f; + for (int nt = 0; nt < 4; nt++) +#pragma unroll + for (int r = 0; r < 2; r++) { + const int row = warp * 32 + mt * 16 + g + r * 8, col = nt * 8 + 2 * t; + const float2 x = s0 ? *reinterpret_cast(s0 + s_at + (size_t)row * V + col) + : make_float2(0.f, 0.f); + st[mt][nt][2 * r] = x.x; + st[mt][nt][2 * r + 1] = x.y; + } load(0, 0); for (int c = 0; c < NC; c++) { @@ -438,6 +450,18 @@ __global__ void __launch_bounds__(ST_THREADS) gdn_chunk_state(Work ws, int NC) { } } } + if (s_out) { +#pragma unroll + for (int mt = 0; mt < 2; mt++) +#pragma unroll + for (int nt = 0; nt < 4; nt++) +#pragma unroll + for (int r = 0; r < 2; r++) { + const int row = warp * 32 + mt * 16 + g + r * 8, col = nt * 8 + 2 * t; + *reinterpret_cast(s_out + s_at + (size_t)row * V + col) = + make_float2(st[mt][nt][2 * r], st[mt][nt][2 * r + 1]); + } + } } // ---- kernel 3: the output, per chunk ---- @@ -535,10 +559,19 @@ extern "C" { size_t cs1_gdn_workspace_floats(int T, int H) { return (Layout(T, H).total + 3) / 4; } -int cs1_gdn_prefill(const void* q, const void* k, const void* v, const float* g, const void* beta, - void* o, float* workspace, int T, int H, int HK, float scale, void* stream) { - if (T < 0 || HK <= 0 || H % HK != 0) return cudaErrorInvalidValue; - if (T == 0) return cudaSuccess; +int cs1_gdn_prefill_state(const void* q, const void* k, const void* v, const float* g, const void* beta, + void* o, float* workspace, const float* initial_state, float* final_state, int T, + int H, int HK, float scale, void* stream) { + if (T < 0 || H < 0 || HK <= 0 || H % HK != 0) return cudaErrorInvalidValue; + if ((reinterpret_cast(initial_state) | reinterpret_cast(final_state)) & 7) + return cudaErrorInvalidValue; + if (T == 0) { + if (!final_state || final_state == initial_state) return cudaSuccess; + const size_t bytes = (size_t)H * K * V * sizeof(float); + cudaStream_t st = static_cast(stream); + return initial_state ? cudaMemcpyAsync(final_state, initial_state, bytes, cudaMemcpyDeviceToDevice, st) + : cudaMemsetAsync(final_state, 0, bytes, st); + } // once per process (for the device current at the first call) static const cudaError_t configured = [] { cudaError_t e = cudaFuncSetAttribute(gdn_chunk_prep, cudaFuncAttributeMaxDynamicSharedMemorySize, @@ -558,9 +591,14 @@ int cs1_gdn_prefill(const void* q, const void* k, const void* v, const float* g, gdn_chunk_prep<<>>( static_cast(q), static_cast(k), static_cast(v), g, static_cast(beta), ws, T, H, HK, scale); - gdn_chunk_state<<>>(ws, NC); + gdn_chunk_state<<>>(ws, NC, initial_state, final_state); gdn_chunk_out<<>>(ws, static_cast(o), T, H); return cudaGetLastError(); } +int cs1_gdn_prefill(const void* q, const void* k, const void* v, const float* g, const void* beta, + void* o, float* workspace, int T, int H, int HK, float scale, void* stream) { + return cs1_gdn_prefill_state(q, k, v, g, beta, o, workspace, nullptr, nullptr, T, H, HK, scale, stream); +} + } // extern "C" diff --git a/src/backends/cuda/qwen3_5/ops.h b/src/backends/cuda/qwen3_5/ops.h index 37a08a77..440a01c5 100644 --- a/src/backends/cuda/qwen3_5/ops.h +++ b/src/backends/cuda/qwen3_5/ops.h @@ -13,7 +13,7 @@ #include // Bumped whenever the required interface below changes. -#define CS1_ABI_VERSION 5 +#define CS1_ABI_VERSION 6 #ifdef __cplusplus extern "C" { @@ -36,6 +36,10 @@ int cs1_graph_destroy(void* exec); // Copy and wait for the copy. int cs1_upload(void* dst, const void* src, size_t bytes, void* stream); int cs1_download(void* dst, const void* src, size_t bytes, void* stream); +// Queue a device-to-device copy of `rows` rows of `row_bytes` bytes, with row +// pitches in bytes; nothing waits for it. +int cs1_copy_rows(void* dst, size_t dst_pitch, const void* src, size_t src_pitch, size_t row_bytes, + int rows, void* stream); // ---- operations ---- @@ -59,6 +63,14 @@ int cs1_gated_rms_norm(const void* x, const void* z, int ldz, const void* w, voi int cs1_gdn_conv(const void* qkv, int ld, const void* w, void* q, void* k, void* v, int T, int key_dim, int value_dim, void* stream); +// The same conv continuing a sequence: history [3, key_dim*2 + value_dim] holds the conv +// inputs of the three positions before qkv's first row, oldest first (null: zeros, as at +// the start of a sequence). If history_out is not null, it receives the inputs of the +// last three positions in the same layout, taken from history where T < 3. history_out +// must not overlap history or qkv. Each output equals the unsplit conv's bit for bit. +int cs1_gdn_conv_history(const void* qkv, int ld, const void* w, const void* history, void* history_out, + void* q, void* k, void* v, int T, int key_dim, int value_dim, void* stream); + // beta = sigmoid(b) (bfloat16) and g = -exp(A_log) * softplus(a + dt_bias) (float32), [T, H]; // b and a are [T, H] in rows of ld. int cs1_gdn_gates(const void* b, const void* a, int ld, const void* A_log, const void* dt_bias, @@ -70,6 +82,16 @@ size_t cs1_gdn_workspace_floats(int T, int H); int cs1_gdn_prefill(const void* q, const void* k, const void* v, const float* g, const void* beta, void* o, float* workspace, int T, int H, int HK, float scale, void* stream); +// The same prefill continuing a sequence from initial_state (float [H, 128, 128], key by +// value per head; null: zeros) and, if final_state is not null, writing the state after +// the last token there in the same layout. Both are 8-byte aligned, and either the same +// buffer or not overlapping. With T = 0, final_state receives initial_state. When every +// split falls on a multiple of 64 tokens, the outputs match the unsplit prefill bit for +// bit; elsewhere the chunks fall differently. +int cs1_gdn_prefill_state(const void* q, const void* k, const void* v, const float* g, const void* beta, + void* o, float* workspace, const float* initial_state, float* final_state, int T, + int H, int HK, float scale, void* stream); + // Attention inputs: q and gate from qg [T, Hq, 2*Dh], k from kr [T, Hk, Dh], both in rows // of ld; per-head zero-centred RMSNorm, then rotary embedding on the first 2*half dims // using cos/sin [T, half] (bfloat16). Writes q [T, Hq, Dh], gate [T, Hq*Dh], k [T, Hk, Dh]. @@ -88,6 +110,14 @@ int cs1_attention(const void* q, const void* k, const void* v, int ldv, void* ou int cs1_attention_gated(const void* q, const void* k, const void* v, int ldv, const void* gate, void* out, int T, int Hq, int Hk, int Dh, float scale, void* stream); +// Gated attention for the last Tq of Tk positions: k [Tk, Hk, Dh] and v (Tk rows of ldv) +// cover all positions, while q, gate and out [Tq, Hq, Dh] are the queries at positions +// Tk - Tq onwards, each attending to the keys up to its own position. Keys are visited +// in the same order as in cs1_attention_gated, so each output row matches the unsplit call. +int cs1_attention_gated_cached(const void* q, const void* k, const void* v, int ldv, const void* gate, + void* out, int Tq, int Tk, int Hq, int Hk, int Dh, float scale, + void* stream); + // x = x * sigmoid(gate), n elements. int cs1_sigmoid_gate(void* x, const void* gate, size_t n, void* stream); diff --git a/src/backends/cuda/qwen3_5/runtime.cu b/src/backends/cuda/qwen3_5/runtime.cu index 697c9db1..1a17b616 100644 --- a/src/backends/cuda/qwen3_5/runtime.cu +++ b/src/backends/cuda/qwen3_5/runtime.cu @@ -36,6 +36,14 @@ int cs1_download(void* dst, const void* src, size_t bytes, void* stream) { return e != cudaSuccess ? e : cudaStreamSynchronize(st); } +int cs1_copy_rows(void* dst, size_t dst_pitch, const void* src, size_t src_pitch, size_t row_bytes, + int rows, void* stream) { + if (rows < 0) return cudaErrorInvalidValue; + if (rows == 0 || row_bytes == 0) return cudaSuccess; + return cudaMemcpy2DAsync(dst, dst_pitch, src, src_pitch, row_bytes, rows, cudaMemcpyDeviceToDevice, + static_cast(stream)); +} + int cs1_graph_begin(void* stream) { const cudaError_t e = cudaStreamBeginCapture(static_cast(stream), cudaStreamCaptureModeThreadLocal); if (e != cudaSuccess) (void)cudaGetLastError(); diff --git a/src/models/qwen3_5/native/src/cuda.rs b/src/models/qwen3_5/native/src/cuda.rs index 3de05686..de307955 100644 --- a/src/models/qwen3_5/native/src/cuda.rs +++ b/src/models/qwen3_5/native/src/cuda.rs @@ -9,7 +9,7 @@ use std::sync::OnceLock; use anyhow::{Context, Result, bail, ensure}; /// `CS1_ABI_VERSION` in ops.h. -const ABI_VERSION: u32 = 5; +const ABI_VERSION: u32 = 6; pub const LIBRARY: &str = "libqwen3_5_cuda.so"; /// A `cudaStream_t`. @@ -73,6 +73,10 @@ api! { cs1_graph_destroy(exec: *mut c_void) -> c_int; cs1_upload(dst: *mut c_void, src: *const c_void, bytes: usize, stream: Stream) -> c_int; cs1_download(dst: *mut c_void, src: *const c_void, bytes: usize, stream: Stream) -> c_int; + cs1_copy_rows( + dst: *mut c_void, dst_pitch: usize, src: *const c_void, src_pitch: usize, row_bytes: usize, rows: c_int, + stream: Stream, + ) -> c_int; cs1_embed(ids: *const i32, table: *const c_void, out: *mut c_void, t: c_int, d: c_int, stream: Stream) -> c_int; cs1_rms_norm( x: *const c_void, w: *const c_void, out: *mut c_void, rows: c_int, d: c_int, eps: f32, stream: Stream, @@ -89,6 +93,10 @@ api! { qkv: *const c_void, ld: c_int, w: *const c_void, q: *mut c_void, k: *mut c_void, v: *mut c_void, t: c_int, key_dim: c_int, value_dim: c_int, stream: Stream, ) -> c_int; + cs1_gdn_conv_history( + qkv: *const c_void, ld: c_int, w: *const c_void, history: *const c_void, history_out: *mut c_void, + q: *mut c_void, k: *mut c_void, v: *mut c_void, t: c_int, key_dim: c_int, value_dim: c_int, stream: Stream, + ) -> c_int; cs1_gdn_gates( b: *const c_void, a: *const c_void, ld: c_int, a_log: *const c_void, dt_bias: *const c_void, beta: *mut c_void, g: *mut f32, t: c_int, h: c_int, stream: Stream, @@ -98,6 +106,11 @@ api! { q: *const c_void, k: *const c_void, v: *const c_void, g: *const f32, beta: *const c_void, o: *mut c_void, workspace: *mut f32, t: c_int, h: c_int, hk: c_int, scale: f32, stream: Stream, ) -> c_int; + cs1_gdn_prefill_state( + q: *const c_void, k: *const c_void, v: *const c_void, g: *const f32, beta: *const c_void, o: *mut c_void, + workspace: *mut f32, initial_state: *const f32, final_state: *mut f32, t: c_int, h: c_int, hk: c_int, + scale: f32, stream: Stream, + ) -> c_int; cs1_attn_prep( qg: *const c_void, kr: *const c_void, ld: c_int, qw: *const c_void, kw: *const c_void, cos: *const c_void, sin: *const c_void, q: *mut c_void, gate: *mut c_void, k: *mut c_void, t: c_int, hq: c_int, hk: c_int, @@ -111,6 +124,10 @@ api! { q: *const c_void, k: *const c_void, v: *const c_void, ldv: c_int, gate: *const c_void, out: *mut c_void, t: c_int, hq: c_int, hk: c_int, dh: c_int, scale: f32, stream: Stream, ) -> c_int; + cs1_attention_gated_cached( + q: *const c_void, k: *const c_void, v: *const c_void, ldv: c_int, gate: *const c_void, + out: *mut c_void, tq: c_int, tk: c_int, hq: c_int, hk: c_int, dh: c_int, scale: f32, stream: Stream, + ) -> c_int; cs1_sigmoid_gate(x: *mut c_void, gate: *const c_void, n: usize, stream: Stream) -> c_int; cs1_silu_mul(gate_up: *const c_void, ld: c_int, out: *mut c_void, t: c_int, i: c_int, stream: Stream) -> c_int; cs1_gemm_create(workspace_bytes: usize) -> *mut c_void; diff --git a/tests/qwen3_5/kernels.rs b/tests/qwen3_5/kernels.rs index 99c9b8ea..ce9f1a88 100644 --- a/tests/qwen3_5/kernels.rs +++ b/tests/qwen3_5/kernels.rs @@ -5,6 +5,7 @@ //! CUA_S1_CUDA_LIB=$PWD/target/release/libqwen3_5_cuda.so \ //! cargo test --release -p omni-qwen3-5-native --test kernels -- --ignored +use std::ffi::c_void; use std::path::PathBuf; use half::bf16; @@ -48,6 +49,18 @@ fn f32_to_device(v: &[f32], st: Stream) -> DeviceBuffer { buf } +fn f32_from_device(buf: &DeviceBuffer, n: usize, st: Stream) -> Vec { + let mut bytes = vec![0u8; n * 4]; + // SAFETY: the buffer holds n float32 values. + unsafe { cuda::download(&mut bytes, buf.at(0), st).unwrap() }; + let (words, _) = bytes.as_chunks::<4>(); + words.iter().map(|&b| f32::from_le_bytes(b)).collect() +} + +fn bits(values: &[f32]) -> Vec { + values.iter().map(|v| v.to_bits()).collect() +} + fn from_device(buf: &DeviceBuffer, n: usize, st: Stream) -> Vec { let mut bytes = vec![0u8; n * 2]; // SAFETY: the buffer holds n bfloat16 values. @@ -471,7 +484,8 @@ fn flash_attention_matches_float64_reference() { } /// Transformers' torch_recurrent_gated_delta_rule in float64, one token at a time, -/// with the L2 norms of q and k and q scaled by K^-1/2. +/// with the L2 norms of q and k and q scaled by K^-1/2. Also returns the state +/// [h, K, V] after each count of tokens in `states_at`. #[allow(clippy::too_many_arguments)] fn gated_delta_reference( q: &[bf16], @@ -483,8 +497,10 @@ fn gated_delta_reference( h: usize, hk: usize, d: usize, -) -> Vec { + states_at: &[usize], +) -> (Vec, Vec>) { let mut out = vec![0f64; t * h * d]; + let mut states = vec![vec![0f64; h * d * d]; states_at.len()]; for head in 0..h { let kh = head / (h / hk); let mut s = vec![0f64; d * d]; // [K][V] @@ -516,9 +532,14 @@ fn gated_delta_reference( for j in 0..d { out[(tok * h + head) * d + j] = (0..d).map(|i| qv[i] * s[i * d + j]).sum(); } + for (at, state) in states_at.iter().zip(&mut states) { + if *at == tok + 1 { + state[head * d * d..][..d * d].copy_from_slice(&s); + } + } } } - out + (out, states) } #[test] @@ -567,7 +588,7 @@ fn gated_delta_rule_matches_recurrent_reference() { .iter() .map(|x| bf16::from_f32(x.to_f32() + 0.5)) .collect(); - let want = gated_delta_reference(&q, &k, &v, &g, &beta, t, h, hk, d); + let (want, _) = gated_delta_reference(&q, &k, &v, &g, &beta, t, h, hk, d, &[]); let (qd, kd, vd, gd, bd) = ( to_device(&q, st), to_device(&k, st), @@ -617,3 +638,477 @@ fn gated_delta_rule_matches_recurrent_reference() { assert!(worst <= 2e-2 * scale, "t = {t}: {worst} vs scale {scale}"); } } + +/// Prefix lengths around and inside 64-token chunks, and the branch lengths after them. +const PREFIXES: [usize; 12] = [1, 2, 3, 31, 63, 64, 65, 127, 128, 129, 1000, 1024]; +const BRANCHES: [usize; 5] = [1, 2, 64, 65, 200]; + +#[test] +#[ignore = "needs a GPU and CUA_S1_CUDA_LIB"] +fn conv_with_history_matches_unsplit_conv() { + let st = setup(); + // 4B/9B (16 key heads, 32 value heads) and 27B (48 value heads); each GDN input + // row also carries z, b and a after the conv channels. + for (key_dim, value_dim, heads) in [(2048usize, 4096usize, 32usize), (2048, 6144, 48)] { + let channels = 2 * key_dim + value_dim; + let ld = channels + value_dim + 2 * heads; + let total = 1300usize; + let input = random(total * ld, 31, 4.0); + let weights = random(channels * 4, 32, 0.5); + let (qkv, w) = (to_device(&input, st), to_device(&weights, st)); + // q|k|v of rows [at, at + t), each token's channels in order + let conv = |at: usize, t: usize, history: *const c_void, history_out: *mut c_void| { + let q = DeviceBuffer::new(t * key_dim * 2).unwrap(); + let k = DeviceBuffer::new(t * key_dim * 2).unwrap(); + let v = DeviceBuffer::new(t * value_dim * 2).unwrap(); + // SAFETY: qkv holds rows [at, at + t) of width ld, the outputs t rows each, + // and history/history_out are null or 3 rows of the conv channels. + unsafe { + check( + (api().cs1_gdn_conv_history)( + qkv.at(at * ld * 2), + ld as i32, + w.at(0), + history, + history_out, + q.at(0), + k.at(0), + v.at(0), + t as i32, + key_dim as i32, + value_dim as i32, + st, + ), + "conv with history", + ) + .unwrap(); + } + let (q, k, v) = ( + from_device(&q, t * key_dim, st), + from_device(&k, t * key_dim, st), + from_device(&v, t * value_dim, st), + ); + (0..t) + .flat_map(|r| { + [ + &q[r * key_dim..][..key_dim], + &k[r * key_dim..][..key_dim], + &v[r * value_dim..][..value_dim], + ] + .concat() + }) + .collect::>() + }; + // the conv inputs of the three positions before `end`, zeros before the start + let source = &input; + let history_before = |end: usize| -> Vec { + (0..3) + .flat_map(|r| { + let s = end as isize - 3 + r; + (0..channels).map(move |c| { + if s < 0 { + 0.0 + } else { + source[s as usize * ld + c].to_f32() + } + }) + }) + .collect() + }; + let full = conv(0, total, std::ptr::null(), std::ptr::null_mut()); + let rows = |a: usize, b: usize| bits(&full[a * channels..b * channels]); + let null = std::ptr::null(); + for p in PREFIXES { + let history = DeviceBuffer::new(3 * channels * 2).unwrap(); + assert_eq!(bits(&conv(0, p, null, history.at(0))), rows(0, p)); + assert_eq!( + from_device(&history, 3 * channels, st), + history_before(p), + "history after {p}" + ); + for b in BRANCHES { + let branch = conv(p, b, history.at(0), std::ptr::null_mut()); + assert_eq!(bits(&branch), rows(p, p + b), "split {p} + {b}"); + } + } + // request prefix, question prefix, then a branch; also chains shorter than the kernel + for (p1, p2, b) in [ + (1usize, 2usize, 1usize), + (2, 3, 2), + (900, 960, 65), + (64, 65, 1), + ] { + let (h1, h2) = ( + DeviceBuffer::new(3 * channels * 2).unwrap(), + DeviceBuffer::new(3 * channels * 2).unwrap(), + ); + conv(0, p1, null, h1.at(0)); + assert_eq!(bits(&conv(p1, p2 - p1, h1.at(0), h2.at(0))), rows(p1, p2)); + assert_eq!(from_device(&h2, 3 * channels, st), history_before(p2)); + assert_eq!( + bits(&conv(p2, b, h2.at(0), std::ptr::null_mut())), + rows(p2, p2 + b), + "chain {p1}, {p2}, {b}" + ); + } + // no tokens: history_out is the history itself, shifted by nothing + let (h1, h2) = ( + DeviceBuffer::new(3 * channels * 2).unwrap(), + DeviceBuffer::new(3 * channels * 2).unwrap(), + ); + conv(0, 5, null, h1.at(0)); + conv(5, 0, h1.at(0), h2.at(0)); + assert_eq!(from_device(&h2, 3 * channels, st), history_before(5)); + } +} + +#[test] +#[ignore = "needs a GPU and CUA_S1_CUDA_LIB"] +fn gated_delta_state_continues_unsplit_prefill() { + let st = setup(); + let d = 128usize; + // 4B/9B (32 value heads over 16 key heads) and 27B (48 over 16) + for (h, hk) in [(32usize, 16usize), (48, 16)] { + let total = 1300usize; + let k = random(total * hk * d, 42, 1.0); + let q: Vec = k + .iter() + .zip(random(total * hk * d, 41, 1.0)) + .map(|(k, n)| bf16::from_f32(0.8 * k.to_f32() + 0.2 * n.to_f32())) + .collect(); + let v = random(total * h * d, 43, 1.0); + // slow decays, so that the state carries far across the split + let g: Vec = random(total * h, 44, 1.0) + .iter() + .map(|x| (x.to_f32() - 1.0) * 0.01) + .collect(); + let beta: Vec = random(total * h, 45, 0.5) + .iter() + .map(|x| bf16::from_f32(x.to_f32() + 0.5)) + .collect(); + // reference states after each prefix and after each question prefix below + let marks: Vec = PREFIXES.iter().copied().chain([192, 1100]).collect(); + let (want, ref_states) = + gated_delta_reference(&q, &k, &v, &g, &beta, total, h, hk, d, &marks); + let ref_state = |at: usize| &ref_states[marks.iter().position(|&m| m == at).unwrap()]; + let (qd, kd, vd, gd, bd) = ( + to_device(&q, st), + to_device(&k, st), + to_device(&v, st), + f32_to_device(&g, st), + to_device(&beta, st), + ); + // SAFETY: pure function of its arguments. + let floats = unsafe { (api().cs1_gdn_workspace_floats)(total as i32, h as i32) }; + let ws = DeviceBuffer::new(floats * 4).unwrap(); + let state_floats = h * d * d; + let run = |at: usize, t: usize, initial: *const f32, last: *mut f32| { + let o = DeviceBuffer::new((t * h * d).max(1) * 2).unwrap(); + // SAFETY: the inputs hold rows [at, at + t), the workspace fits the longest + // run, and the states are null or [h, 128, 128] floats. + unsafe { + check( + (api().cs1_gdn_prefill_state)( + qd.at(at * hk * d * 2), + kd.at(at * hk * d * 2), + vd.at(at * h * d * 2), + gd.at(at * h * 4).cast::(), + bd.at(at * h * 2), + o.at(0), + ws.at(0).cast::(), + initial, + last, + t as i32, + h as i32, + hk as i32, + (d as f32).powf(-0.5), + st, + ), + "gdn prefill with state", + ) + .unwrap(); + } + from_device(&o, t * h * d, st) + }; + let null = std::ptr::null(); + let full = run(0, total, null, std::ptr::null_mut()); + let rows = |a: usize, b: usize| &full[a * h * d..b * h * d]; + let scale_of = |a: usize, b: usize| { + want[a * h * d..b * h * d] + .iter() + .fold(0f64, |m, x| m.max(x.abs())) + }; + let worst = |got: &[f32], want: &[f64]| { + got.iter() + .zip(want) + .map(|(a, b)| (*a as f64 - b).abs()) + .fold(0f64, f64::max) + }; + for (i, p) in PREFIXES.into_iter().enumerate() { + let state = DeviceBuffer::new(state_floats * 4).unwrap(); + let prefix = run(0, p, null, state.at(0).cast()); + assert_eq!( + bits(&prefix), + bits(rows(0, p)), + "writing the state changed the output, p = {p}" + ); + let s = f32_from_device(&state, state_floats, st); + let s_scale = ref_states[i].iter().fold(0f64, |m, x| m.max(x.abs())); + let s_worst = worst(&s, ref_state(p)); + eprintln!( + "gated delta state after {p} ({h} heads): largest difference {s_worst:.2e}, largest |reference| {s_scale:.2e}" + ); + assert!( + s_worst <= 2e-2 * s_scale, + "state after {p}: {s_worst} vs {s_scale}" + ); + for b in BRANCHES { + let branch = run(p, b, state.at(0).cast(), std::ptr::null_mut()); + assert!(branch.iter().all(|x| x.is_finite())); + if p % 64 == 0 { + assert_eq!( + bits(&branch), + bits(rows(p, p + b)), + "aligned split {p} + {b}" + ); + continue; + } + let scale = scale_of(p, p + b); + let to_ref = worst(&branch, &want[p * h * d..(p + b) * h * d]); + let to_unsplit = branch + .iter() + .zip(rows(p, p + b)) + .map(|(a, b)| (a - b).abs() as f64) + .fold(0f64, f64::max); + eprintln!( + "gated delta split {p} + {b} ({h} heads): to float64 {to_ref:.2e}, to unsplit {to_unsplit:.2e}, largest |reference| {scale:.2e}" + ); + assert!( + to_ref <= 2e-2 * scale, + "split {p} + {b}: {to_ref} vs {scale}" + ); + } + } + // request prefix -> question prefix -> branch, also with the question state + // updated in place + for (p1, p2, b) in [ + (64usize, 128usize, 65usize), + (100, 192, 65), + (1000, 1100, 100), + ] { + let (s1, s2, s1b) = ( + DeviceBuffer::new(state_floats * 4).unwrap(), + DeviceBuffer::new(state_floats * 4).unwrap(), + DeviceBuffer::new(state_floats * 4).unwrap(), + ); + run(0, p1, null, s1.at(0).cast()); + run(0, p1, null, s1b.at(0).cast()); + let question = run(p1, p2 - p1, s1.at(0).cast(), s2.at(0).cast()); + let in_place = run(p1, p2 - p1, s1b.at(0).cast(), s1b.at(0).cast()); + assert_eq!(bits(&question), bits(&in_place)); + assert_eq!( + bits(&f32_from_device(&s2, state_floats, st)), + bits(&f32_from_device(&s1b, state_floats, st)), + "in-place state {p1}, {p2}" + ); + let branch = run(p2, b, s2.at(0).cast(), std::ptr::null_mut()); + if p1 % 64 == 0 && p2 % 64 == 0 { + assert_eq!(bits(&question), bits(rows(p1, p2))); + assert_eq!( + bits(&branch), + bits(rows(p2, p2 + b)), + "aligned chain {p1}, {p2}" + ); + } else { + let q_ref = worst(&question, &want[p1 * h * d..p2 * h * d]); + assert!( + q_ref <= 2e-2 * scale_of(p1, p2), + "question {p1}, {p2}: {q_ref}" + ); + let s2_ref = ref_state(p2); + let s2_worst = worst(&f32_from_device(&s2, state_floats, st), s2_ref); + let s2_scale = s2_ref.iter().fold(0f64, |m, x| m.max(x.abs())); + assert!(s2_worst <= 2e-2 * s2_scale, "state {p1}, {p2}: {s2_worst}"); + let scale = scale_of(p2, p2 + b); + let to_ref = worst(&branch, &want[p2 * h * d..(p2 + b) * h * d]); + eprintln!( + "gated delta chain {p1}, {p2} + {b} ({h} heads): question to float64 {q_ref:.2e}, state {s2_worst:.2e}, branch {to_ref:.2e}" + ); + assert!(to_ref <= 2e-2 * scale, "chain {p1}, {p2} + {b}"); + } + } + // no tokens: the final state is the initial one, or zeros without one + let (s1, s2) = ( + DeviceBuffer::new(state_floats * 4).unwrap(), + DeviceBuffer::new(state_floats * 4).unwrap(), + ); + run(0, 200, null, s1.at(0).cast()); + run(200, 0, s1.at(0).cast(), s2.at(0).cast()); + assert_eq!( + bits(&f32_from_device(&s2, state_floats, st)), + bits(&f32_from_device(&s1, state_floats, st)) + ); + run(200, 0, null, s2.at(0).cast()); + assert!( + f32_from_device(&s2, state_floats, st) + .iter() + .all(|x| x.to_bits() == 0) + ); + // a misaligned state is rejected before any work is queued + for (initial, last) in [ + (s1.at(4).cast::().cast_const(), std::ptr::null_mut()), + (std::ptr::null(), s2.at(4).cast::()), + ] { + // SAFETY: the call returns before launching anything. + let code = unsafe { + (api().cs1_gdn_prefill_state)( + std::ptr::null(), + std::ptr::null(), + std::ptr::null(), + std::ptr::null(), + std::ptr::null(), + std::ptr::null_mut(), + std::ptr::null_mut(), + initial, + last, + 1, + h as i32, + hk as i32, + 1.0, + st, + ) + }; + assert_ne!(code, 0); + } + } +} + +#[test] +#[ignore = "needs a GPU and CUA_S1_CUDA_LIB"] +fn cached_attention_matches_unsplit_attention() { + let st = setup(); + let dh = 256usize; + // 4B/9B (16 query heads) and 27B (24) over 4 KV heads + for (hq, hk) in [(16usize, 4usize), (24, 4)] { + let total = 1300usize; + let ldv = hk * dh + 16; // a strided V buffer, as read from the projection output + let q = to_device(&random(total * hq * dh, 51, 2.0), st); + let k = to_device(&random(total * hk * dh, 52, 2.0), st); + let v = to_device(&random(total * ldv, 53, 1.0), st); + let gate = to_device(&random(total * hq * dh, 54, 12.0), st); + let row = hq * dh; + let attend = |at: usize, tq: usize, tk: usize| { + let out = DeviceBuffer::new((tq * row).max(1) * 2).unwrap(); + // SAFETY: q and gate hold rows [at, at + tq), k and v the first tk rows. + unsafe { + check( + (api().cs1_attention_gated_cached)( + q.at(at * row * 2), + k.at(0), + v.at(0), + ldv as i32, + gate.at(at * row * 2), + out.at(0), + tq as i32, + tk as i32, + hq as i32, + hk as i32, + dh as i32, + 0.0625, + st, + ), + "cached attention", + ) + .unwrap(); + } + from_device(&out, tq * row, st) + }; + // with every position as a query, the cached call runs the same code as the + // unsplit one; the check below only guards the two entry points + let full = attend(0, total, total); + let unsplit = DeviceBuffer::new(total * row * 2).unwrap(); + // SAFETY: as above, with total rows everywhere. + unsafe { + check( + (api().cs1_attention_gated)( + q.at(0), + k.at(0), + v.at(0), + ldv as i32, + gate.at(0), + unsplit.at(0), + total as i32, + hq as i32, + hk as i32, + dh as i32, + 0.0625, + st, + ), + "gated attention", + ) + .unwrap(); + } + assert_eq!(bits(&full), bits(&from_device(&unsplit, total * row, st))); + for p in PREFIXES { + for b in BRANCHES { + let got = attend(p, b, p + b); + assert!(got.iter().all(|x| x.is_finite())); + assert_eq!( + bits(&got), + bits(&full[p * row..(p + b) * row]), + "split {p} + {b}, {hq} heads" + ); + } + } + // no queries is valid; more queries than keys is not + // SAFETY: both calls return before launching anything. + unsafe { + let null = std::ptr::null(); + let out = std::ptr::null_mut(); + let call = |tq: i32, tk: i32| { + (api().cs1_attention_gated_cached)( + null, null, null, 1024, null, out, tq, tk, 16, 4, 256, 0.0625, st, + ) + }; + assert_eq!(call(0, 5), 0); + assert_ne!(call(2, 1), 0); + } + } +} + +#[test] +#[ignore = "needs a GPU and CUA_S1_CUDA_LIB"] +fn copy_rows_copies_pitched_rows() { + let st = setup(); + let (rows, width, src_pitch, dst_pitch) = (37usize, 1024usize, 10240usize, 2048usize); + let src = random(rows * src_pitch / 2, 61, 1.0); + let source = to_device(&src, st); + let target = to_device(&vec![bf16::from_f32(7.0); rows * dst_pitch / 2], st); + // SAFETY: both buffers hold `rows` rows of their pitch, each wider than `width` bytes. + unsafe { + check( + (api().cs1_copy_rows)( + target.at(0), + dst_pitch, + source.at(0), + src_pitch, + width, + rows as i32, + st, + ), + "copy rows", + ) + .unwrap(); + } + let got = from_device(&target, rows * dst_pitch / 2, st); + for r in 0..rows { + for i in 0..dst_pitch / 2 { + let want = if i < width / 2 { + src[r * src_pitch / 2 + i].to_f32() + } else { + 7.0 + }; + assert_eq!(got[r * dst_pitch / 2 + i], want, "row {r}, element {i}"); + } + } +}