From d77faca885e376c446b093a98b2eaa1603185232 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Fri, 17 Apr 2026 23:46:25 +0200 Subject: [PATCH] feat: add Attention cuBLAS element-wise cubin kernels Add 6 new CUDA kernel functions to attention_backward_kernel.cu for the batched cuBLAS attention path: attn_sdp_fwd, attn_sdp_bwd (sigmoid-gated feature attention), attn_layer_norm_fwd, attn_layer_norm_bwd (per-sample LayerNorm with residual), attn_bias_add, and attn_bias_grad_reduce. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../attention_backward_kernel.cu | 214 ++++++++++++++++++ 1 file changed, 214 insertions(+) diff --git a/crates/ml/src/cuda_pipeline/attention_backward_kernel.cu b/crates/ml/src/cuda_pipeline/attention_backward_kernel.cu index b14d9fa15..45640cd32 100644 --- a/crates/ml/src/cuda_pipeline/attention_backward_kernel.cu +++ b/crates/ml/src/cuda_pipeline/attention_backward_kernel.cu @@ -584,3 +584,217 @@ extern "C" __global__ void attn_adam_kernel( /* Zero gradient for next step */ grads[tid] = 0.0f; } + +/* ===================================================================== + * ATTENTION ELEMENT-WISE CUBIN KERNELS + * + * Standalone element-wise kernels for batched cuBLAS attention path. + * All tensors use column-major layout [D, B] (feature-major stride). + * ===================================================================== */ + +/* ===================================================================== + * attn_sdp_fwd — Scaled dot-product attention forward (feature-level) + * + * Feature attention on h_s2[D, B] where D=256, B=16384, num_heads=4, + * head_dim=64. NOT sequence attention — each sample is independent. + * Per head: score = dot(Q_h, K_h) / sqrt(head_dim), gate = sigmoid(score), + * attn_out_h = V_h * gate. + * + * Layout: Q, K, V, attn_out are [D, B] (column-major). + * save_scores is [num_heads, B]. + * Grid: (B), Block: (num_heads or power-of-two >= num_heads). + * ===================================================================== */ +extern "C" __global__ +void attn_sdp_fwd( + const float* __restrict__ Q, + const float* __restrict__ K, + const float* __restrict__ V, + float* __restrict__ attn_out, + float* __restrict__ save_scores, + int B, int D, int num_heads +) { + int sample = blockIdx.x; + if (sample >= B) return; + int head_dim = D / num_heads; + float scale = 1.0f / sqrtf((float)head_dim); + + for (int h = threadIdx.x; h < num_heads; h += blockDim.x) { + int off = h * head_dim; + float score = 0.0f; + for (int d = 0; d < head_dim; d++) { + int idx = (off + d) * B + sample; + score += Q[idx] * K[idx]; + } + score *= scale; + float gate = 1.0f / (1.0f + expf(-score)); + save_scores[h * B + sample] = gate; + for (int d = 0; d < head_dim; d++) { + int idx = (off + d) * B + sample; + attn_out[idx] = V[idx] * gate; + } + } +} + +/* ===================================================================== + * attn_sdp_bwd — Scaled dot-product attention backward + * + * Backpropagates through sigmoid-gated feature attention. + * Computes d_Q, d_K, d_V given d_attn_out and saved sigmoid scores. + * + * Layout: all [D, B] except scores which is [num_heads, B]. + * Grid: (B), Block: (num_heads or power-of-two >= num_heads). + * ===================================================================== */ +extern "C" __global__ +void attn_sdp_bwd( + const float* __restrict__ d_attn_out, + const float* __restrict__ Q, + const float* __restrict__ K, + const float* __restrict__ V, + const float* __restrict__ scores, + float* __restrict__ d_Q, + float* __restrict__ d_K, + float* __restrict__ d_V, + int B, int D, int num_heads +) { + int sample = blockIdx.x; + if (sample >= B) return; + int head_dim = D / num_heads; + float scale = 1.0f / sqrtf((float)head_dim); + + for (int h = threadIdx.x; h < num_heads; h += blockDim.x) { + int off = h * head_dim; + float gate = scores[h * B + sample]; + float dgate = 0.0f; + for (int d = 0; d < head_dim; d++) { + int idx = (off + d) * B + sample; + d_V[idx] = d_attn_out[idx] * gate; + dgate += d_attn_out[idx] * V[idx]; + } + float dscore = dgate * gate * (1.0f - gate) * scale; + for (int d = 0; d < head_dim; d++) { + int idx = (off + d) * B + sample; + d_Q[idx] = dscore * K[idx]; + d_K[idx] = dscore * Q[idx]; + } + } +} + +/* ===================================================================== + * attn_layer_norm_fwd — Per-sample LayerNorm with residual connection + * + * Computes: out = gamma * ((x + residual - mean) / std) + beta + * Saves mean and rstd per sample for backward pass. + * + * Layout: x, residual, out are [D, B]. gamma, beta are [D]. + * save_mean, save_rstd are [B]. + * Grid: (B), Block: (1) — serial over D per sample. + * ===================================================================== */ +extern "C" __global__ +void attn_layer_norm_fwd( + const float* __restrict__ x, + const float* __restrict__ residual, + const float* __restrict__ gamma, + const float* __restrict__ beta, + float* __restrict__ out, + float* __restrict__ save_mean, + float* __restrict__ save_rstd, + int D, int B +) { + int sample = blockIdx.x; + if (sample >= B) return; + float mean = 0.0f; + for (int d = 0; d < D; d++) mean += x[d * B + sample] + residual[d * B + sample]; + mean /= (float)D; + float var = 0.0f; + for (int d = 0; d < D; d++) { + float v = x[d * B + sample] + residual[d * B + sample] - mean; + var += v * v; + } + var /= (float)D; + float rstd = 1.0f / sqrtf(var + 1e-5f); + save_mean[sample] = mean; + save_rstd[sample] = rstd; + for (int d = 0; d < D; d++) { + float normed = (x[d * B + sample] + residual[d * B + sample] - mean) * rstd; + out[d * B + sample] = gamma[d] * normed + beta[d]; + } +} + +/* ===================================================================== + * attn_layer_norm_bwd — LayerNorm backward with residual connection + * + * Backpropagates through per-sample LayerNorm + residual. + * Computes d_x (same gradient as d_residual due to residual add), + * d_gamma, and d_beta accumulated over the batch dimension. + * + * Layout: d_out, x, residual, d_x are [D, B]. gamma, d_gamma, d_beta + * are [D]. save_mean, save_rstd are [B]. + * Grid: (ceil(D/blockDim.x)), Block: (256). + * ===================================================================== */ +extern "C" __global__ +void attn_layer_norm_bwd( + const float* __restrict__ d_out, + const float* __restrict__ x, + const float* __restrict__ residual, + const float* __restrict__ gamma, + const float* __restrict__ save_mean, + const float* __restrict__ save_rstd, + float* __restrict__ d_x, + float* __restrict__ d_gamma, + float* __restrict__ d_beta, + int D, int B +) { + int d = blockIdx.x * blockDim.x + threadIdx.x; + if (d >= D) return; + float dg = 0.0f, db = 0.0f; + for (int b = 0; b < B; b++) { + float mean = save_mean[b]; + float rstd = save_rstd[b]; + float normed = (x[d * B + b] + residual[d * B + b] - mean) * rstd; + dg += d_out[d * B + b] * normed; + db += d_out[d * B + b]; + d_x[d * B + b] = d_out[d * B + b] * gamma[d] * rstd; + } + d_gamma[d] = dg; + d_beta[d] = db; +} + +/* ===================================================================== + * attn_bias_add — In-place bias addition for attention projections + * + * Adds bias vector (length out_dim) to a flat [out_dim, B] activation. + * idx % out_dim selects the bias element for each output position. + * + * Grid: (ceil(N/blockDim.x)), Block: (256). N = out_dim * B. + * ===================================================================== */ +extern "C" __global__ +void attn_bias_add( + float* __restrict__ x, + const float* __restrict__ bias, + int out_dim, int N +) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < N) { x[idx] += bias[idx % out_dim]; } +} + +/* ===================================================================== + * attn_bias_grad_reduce — Bias gradient reduction for attention projections + * + * Reduces d_pre [B, out_dim] (row-major) over the batch dimension + * and scales by inv_batch to produce d_bias [out_dim]. + * + * Grid: (ceil(out_dim/blockDim.x)), Block: (256). + * ===================================================================== */ +extern "C" __global__ +void attn_bias_grad_reduce( + const float* __restrict__ d_pre, + float* __restrict__ d_bias, + int out_dim, int B, float inv_batch +) { + int j = blockIdx.x * blockDim.x + threadIdx.x; + if (j < out_dim) { + float sum = 0.0f; + for (int b = 0; b < B; b++) { sum += d_pre[j + b * out_dim]; } + d_bias[j] = sum * inv_batch; + } +}