Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion recipe/cua_s1/native.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
4 changes: 1 addition & 3 deletions recipe/open_jev/native.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
4 changes: 2 additions & 2 deletions src/backends/cuda/contract.md
Original file line number Diff line number Diff line change
Expand Up @@ -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"],
Expand Down Expand Up @@ -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
Expand Down
26 changes: 24 additions & 2 deletions src/backends/cuda/qwen3_5/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
46 changes: 28 additions & 18 deletions src/backends/cuda/qwen3_5/attention.cu
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -81,7 +83,7 @@ constexpr int SMEM_BYTES = (BM + 2 * BN) * LDS * 2;
template <bool Gated>
__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<bf16*>(smem);
Expand All @@ -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();

Expand All @@ -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;
Expand All @@ -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]);
}
}
Expand Down Expand Up @@ -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++) {
Expand All @@ -232,20 +236,20 @@ __global__ void __launch_bounds__(THREADS)

template <bool Gated>
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<Gated>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_BYTES);
if (configured != cudaSuccess) return configured;
constexpr float LOG2E = 1.4426950408889634f;
flash_kernel<Gated><<<dim3((T + BM - 1) / BM, Hq), THREADS, SMEM_BYTES,
flash_kernel<Gated><<<dim3((Tq + BM - 1) / BM, Hq), THREADS, SMEM_BYTES,
static_cast<cudaStream_t>(stream)>>>(
static_cast<const bf16*>(q), static_cast<const bf16*>(k), static_cast<const bf16*>(v), ldv,
static_cast<const bf16*>(gate), static_cast<bf16*>(out), T, Hq, Hk, scale * LOG2E);
static_cast<const bf16*>(gate), static_cast<bf16*>(out), Tq, Tk, Hq, Hk, scale * LOG2E);
return cudaGetLastError();
}

Expand Down Expand Up @@ -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<false>(q, k, v, ldv, nullptr, out, T, Hq, Hk, Dh, scale, stream);
return flash::launch<false>(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<true>(q, k, v, ldv, gate, out, T, Hq, Hk, Dh, scale, stream);
return flash::launch<true>(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<true>(q, k, v, ldv, gate, out, Tq, Tk, Hq, Hk, Dh, scale, stream);
}
51 changes: 40 additions & 11 deletions src/backends/cuda/qwen3_5/elementwise.cu
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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)
Expand All @@ -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,
Expand Down Expand Up @@ -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<<<blocks(n), THREADS, 0, static_cast<cudaStream_t>(stream)>>>(
static_cast<const bf16*>(qkv), ld, static_cast<const bf16*>(w), static_cast<bf16*>(q),
static_cast<bf16*>(k), static_cast<bf16*>(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<cudaStream_t>(stream);
if (n > 0) {
gdn_conv_kernel<<<blocks(n), THREADS, 0, st>>>(
static_cast<const bf16*>(qkv), ld, static_cast<const bf16*>(w), static_cast<const bf16*>(history),
static_cast<bf16*>(q), static_cast<bf16*>(k), static_cast<bf16*>(v), T, key_dim, value_dim);
}
if (history_out && channels > 0) {
gdn_conv_history_kernel<<<blocks((size_t)3 * channels), THREADS, 0, st>>>(
static_cast<const bf16*>(qkv), ld, static_cast<const bf16*>(history), static_cast<bf16*>(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;
Expand Down
Loading
Loading