// --- Shared memory for dS accumulator [D, D] in fp32 // Also used for loading workspace tiles extern __shared__ __align__(128) unsignedchar shared_mem[]; float* dS_smem = reinterpret_cast<float*>(shared_mem); // [D*D] // After dS: scratch for loading tiles constexprint dS_size = D * D; float* scratch_f32 = dS_smem + dS_size;
// Initialize dS from ds_init or zero // ds_init (dfinal_state) has layout [V, K] (transposed state), but dS_smem is [K, V] // So we transpose on load: dS_smem[k*D + v] = src[v*D + k] { int state_linear = seq_idx * H + head_idx; if (ds_init_ptr != nullptr) { BF16 const* src = ds_init_ptr + int64_t(state_linear) * D * D; for (int i = tid; i < dS_size; i += NumThreads) { int k = i / D; int v_idx = i % D; dS_smem[i] = bf16_to_f32(src[v_idx * D + k]); } } else { for (int i = tid; i < dS_size; i += NumThreads) { dS_smem[i] = 0.0f; } } } __syncthreads();
// Iterate chunks from back to front for (int t = t_tiles - 1; t >= 0; --t) { int ws_idx = head_idx * total_tiles + tile_base + t; int actual_len = min(CHUNK, seq_len - t * CHUNK); constexprint CD = CHUNK * D;
// --- Load all needed data for this chunk into registers --- // We use a simple approach: each thread loads elements assigned to it // and we do the backward math element-wise or with simple reductions.
// S_in for this chunk: indexed same as workspace int state_idx = head_idx * total_tiles + tile_base + t; BF16 const* s_in_ptr = all_states_ptr + int64_t(state_idx) * D * D;
// ========== Load tile data into shared memory ========== // We need: kd[C,D], qd[C,D], kr[C,D], ki[C,D], v[C,D], do[C,D], // INV[C,C], Mqk[C,C], gT[D], beta[C], S_in[D,D] // Total smem needed beyond dS: quite a lot. Let's use registers where possible.
// Load gT[D] into shared scratch float* gT_smem = scratch_f32; // [D] for (int i = tid; i < D; i += NumThreads) { gT_smem[i] = gt_tile[i]; } __syncthreads();
// Load beta[CHUNK] into first CHUNK floats of scratch after gT float* beta_smem = scratch_f32 + D; // [CHUNK] if (tid < CHUNK) { float b_raw = bf16_to_f32(beta_tile_ptr[tid]); beta_smem[tid] = sigmoid_tanh_approx_f32(b_raw); } __syncthreads();
// Load INV[C,C] and Mqk[C,C] into shared float* inv_smem = scratch_f32 + D + CHUNK; // [C*C] float* mqk_smem = inv_smem + CHUNK * CHUNK; // [C*C] for (int i = tid; i < CHUNK * CHUNK; i += NumThreads) { inv_smem[i] = bf16_to_f32(inv_tile[i]); mqk_smem[i] = bf16_to_f32(mqk_tile[i]); } __syncthreads();
// ========== Step 1: Recompute U ========== // vcorr = (v - kd @ S_in) * beta // U = INV @ vcorr // We compute these in shared memory using simple GEMM loops.
// First: kd @ S_in -> tmp[C, D] (stored in smem) // S_in is [D, D], kd is [C, D], result is [C, D] // This is expensive: C*D*D = 16*128*128 = 262144 FMAs // With 256 threads, each thread does ~1024 FMAs
// We'll compute kd_S_in[c][d] = sum_k kd[c][k] * S_in[k][d] // Thread assignment: each thread handles a subset of (c,d) pairs float* kd_S_smem = mqk_smem + CHUNK * CHUNK; // [C*D]
// all_states stores S^T[V,K], so s_in_ptr[v*D+k] = S^T[v][k] = S[k][v] // We need kd_S[c][d] = sum_k kd[c][k] * S[k][d] // S[k][d] = s_in_ptr[d*D + k] for (int idx = tid; idx < CD; idx += NumThreads) { int c = idx / D; int d = idx % D; float acc = 0.0f; for (int k = 0; k < D; ++k) { acc += kd_tile[c * D + k] * bf16_to_f32(s_in_ptr[d * D + k]); } kd_S_smem[idx] = acc; } __syncthreads();
// U = INV @ vcorr: U[c][d] = sum_j INV[c][j] * vcorr[j][d] float* U_smem = vcorr_smem + CD; // need new buffer for U since vcorr is still needed // Actually we can't reuse vcorr yet. Let's use a different region. // Let's reorganize: put U after the scratch we've used. // scratch layout: gT[D] | beta[C] | INV[C*C] | Mqk[C*C] | vcorr[C*D] | U[C*D] // That's D + C + 2*C*C + 2*C*D = 128+16+512+4096 = 4752 floats = 19008 bytes // Plus dS[D*D] = 16384 floats = 65536 bytes. Total ~84KB, within SM90 smem limits. U_smem = vcorr_smem + CD; // [C*D]
for (int idx = tid; idx < CD; idx += NumThreads) { int c = idx / D; int d = idx % D; float acc = 0.0f; for (int j = 0; j < CHUNK; ++j) { acc += inv_smem[c * CHUNK + j] * vcorr_smem[j * D + d]; } U_smem[idx] = acc; } __syncthreads();
// ========== Step 2: Compute gradients ========== // do_tile[C, D] is the output gradient for this chunk
// dout = do[C, D] // out = qd @ S_in + Mqk @ U // // From out = qd @ S_in: // dqd_cross[c][k] += sum_d do[c][d] * S_in[k][d] (= do @ S_in^T) // dS[k][d] += sum_c qd[c][k] * do[c][d] (= qd^T @ do) // // From out += Mqk @ U: // dMqk[c][j] = sum_d do[c][d] * U[j][d] (= do @ U^T) // dU[j][d] += sum_c Mqk[c][j]^T * do[c][d] (= Mqk^T @ do)
// First compute dU from Mqk path: dU_mqk = Mqk^T @ do float* dU_smem = U_smem + CD; // [C*D] // scratch: gT[D] | beta[C] | INV[C*C] | Mqk[C*C] | vcorr[C*D] | U[C*D] | dU[C*D] // = D + C + 2*C*C + 3*C*D = 128+16+512+6144 = 6800 floats = 27200B // plus dS = 65536B. Total ~92KB. Tight but should be OK on SM90.
for (int idx = tid; idx < CD; idx += NumThreads) { int j = idx / D; // row of dU (chunk dim) int d = idx % D; float acc = 0.0f; for (int c = 0; c < CHUNK; ++c) { float do_val = (c < actual_len) ? bf16_to_f32(do_tile_ptr[c * D + d]) : 0.0f; acc += mqk_smem[c * CHUNK + j] * do_val; // Mqk^T: [j][c] = Mqk[c][j] } dU_smem[idx] = acc; } __syncthreads();
// Add dU contribution from state update: S_new = S * exp(gT) + kr^T @ U // dkr[c][k] = sum_d U[c][d] * dS[k][d] ... but actually: // S_new[k][d] += sum_c kr[c][k] * U[c][d] => kr^T @ U // dkr[c][k] = sum_d dS_new[k][d] * U[c][d] = dS @ U^T then transpose? No: // dkr^T[k][c] = sum_d dS[k][d] * U[c][d] => dkr^T = dS @ U^T => dkr = (dS @ U^T)^T = U @ dS^T // dU[c][d] += sum_k kr[c][k] * dS_new[k][d] = kr @ dS_new (from the kr^T @ U term) // But dS_new = dS (current dS, since we iterate backward and dS comes from next chunk)
// Add kr @ dS to dU for (int idx = tid; idx < CD; idx += NumThreads) { int c = idx / D; int d = idx % D; float acc = 0.0f; for (int k = 0; k < D; ++k) { acc += kr_tile[c * D + k] * dS_smem[k * D + d]; } dU_smem[idx] += acc; } __syncthreads();
// Now compute dkr from dS and U: dkr[c][k] = sum_d U[c][d] * dS[k][d] // = (U @ dS^T) ... but we want dkr, and the update was S += kr^T @ U // More precisely: d(kr^T @ U)/dkr = ... dkr[c][k] = sum_d dS[k][d]*U[c][d] float* dkr_smem = dU_smem + CD; // [C*D]
for (int idx = tid; idx < CD; idx += NumThreads) { int c = idx / D; int k = idx % D; float acc = 0.0f; for (int d = 0; d < D; ++d) { acc += dS_smem[k * D + d] * U_smem[c * D + d]; } dkr_smem[idx] = acc; } __syncthreads();
// ========== Step 3: Triangular solve adjoint ========== // U = INV @ vcorr, so dvcorr = INV^T @ dU // (since U = A^{-1} @ vcorr where A = I+L, d/dvcorr = A^{-T} @ dU = INV^T @ dU) float* dvcorr_smem = dkr_smem + CD; // [C*D] float* dL_smem = dvcorr_smem + CD; // [C*C] — dedicated region, not aliased // scratch: gT[D]|beta[C]|INV[C*C]|Mqk[C*C]|vcorr[C*D]|U[C*D]|dU[C*D]|dkr[C*D]|dvcorr[C*D]|dL[C*C] // = D+C+2*C*C+5*C*D+C*C floats
for (int idx = tid; idx < CD; idx += NumThreads) { int c = idx / D; int d = idx % D; float acc = 0.0f; for (int j = 0; j < CHUNK; ++j) { acc += inv_smem[j * CHUNK + c] * dU_smem[j * D + d]; // INV^T[c][j] = INV[j][c] } dvcorr_smem[idx] = acc; } __syncthreads();
// ========== Step 3b: Compute dL while dvcorr is still alive ========== // dL[i][j] = -sum_d dvcorr[i][d] * U[j][d], strictly lower triangular (i > j) for (int idx = tid; idx < CHUNK * CHUNK; idx += NumThreads) { int i = idx / CHUNK; int j = idx % CHUNK; if (i > j) { float acc = 0.0f; for (int d = 0; d < D; ++d) { acc += dvcorr_smem[i * D + d] * U_smem[j * D + d]; } dL_smem[idx] = -acc; } else { dL_smem[idx] = 0.0f; } } __syncthreads();
// Store dv to bwd workspace BF16* dv_out = bwd_ws_dv_ptr + int64_t(ws_idx) * CHUNK * D; for (int idx = tid; idx < CD; idx += NumThreads) { int c = idx / D; int d = idx % D; float dv_val = dvcorr_smem[idx] * beta_smem[c]; dv_out[idx] = BF16(dv_val); }
// dbeta_chunk[c] = sum_d dvcorr[c][d] * vcorr_orig[c][d] (where vcorr_orig = (v-kd@S)) // We already overwrote vcorr_smem. We need vcorr_orig = vcorr / beta for non-zero beta. // Actually, vcorr = (v - kd@S) * beta, so (v - kd@S) = vcorr / beta. // dbeta[c] = sum_d dvcorr[c][d] * (v-kd@S)[c][d] // = sum_d dvcorr[c][d] * vcorr[c][d] / beta[c] // But vcorr was overwritten... Let me recompute from U: // vcorr was overwritten but we have U = INV @ vcorr, so vcorr = (I+L) @ U // Actually wait, let me check what vcorr_smem contains at this point... // vcorr_smem was computed early, then we put U_smem after it. vcorr_smem should still be valid. // Let me trace: vcorr_smem = kd_S_smem = scratch_f32 + D + CHUNK + 2*C*C // After computing vcorr, we computed U into vcorr + CD, then dU into vcorr + 2*CD, etc. // vcorr_smem itself should still contain vcorr values. Good.
// Compute dbeta per chunk row float* dbeta_out = bwd_ws_dbeta_ptr + int64_t(ws_idx) * CHUNK; if (tid < CHUNK) { int c = tid; float beta_c = beta_smem[c]; float sum = 0.0f; if (beta_c > 1e-8f) { for (int d = 0; d < D; ++d) { sum += dvcorr_smem[c * D + d] * vcorr_smem[c * D + d] / beta_c; } } dbeta_out[c] = sum; } __syncthreads();
// d(kd@S)/dkd[c][k] = -beta[c] * sum_d dvcorr[c][d] * S_in[k][d] (from -kd@S*beta part) // = -dvcorr_beta[c][d] @ S_in^T // dkd[c][k] = -sum_d (dvcorr[c][d]*beta[c]) * S_in[k][d] // Also from the Mqk@U backward, we need dMqk contributions. Let's handle dMqk: // dMqk[c][j] = sum_d do[c][d] * U[j][d] // But Mqk = tril(qd @ ki^T), so dqd += tril(dMqk) @ ki, dki += tril(dMqk)^T @ qd // For now, compute dMqk and apply tril mask.
// Compute dMqk[C, C] = do @ U^T (only lower triangular matters since Mqk was tril) // Reuse some scratch. We can rewrite mqk_smem since we don't need Mqk anymore. float* dMqk_smem = mqk_smem; // reuse [C*C] for (int idx = tid; idx < CHUNK * CHUNK; idx += NumThreads) { int c = idx / CHUNK; int j = idx % CHUNK; if (c >= j) { // tril including diagonal float acc = 0.0f; for (int d = 0; d < D; ++d) { float do_val = (c < actual_len) ? bf16_to_f32(do_tile_ptr[c * D + d]) : 0.0f; acc += do_val * U_smem[j * D + d]; } dMqk_smem[idx] = acc; } else { dMqk_smem[idx] = 0.0f; } } __syncthreads();
// ========== Step 5: Accumulate dqd from cross-chunk and within-chunk ========== // dqd[c][k] = sum_d do[c][d] * S_in[k][d] (cross-chunk: qd@S) // + sum_j tril(dMqk)[c][j] * ki[j][k] (within-chunk: Mqk@U, where Mqk=tril(qd@ki^T)) float* dqd_local = dvcorr_smem; // reuse [C*D] for (int idx = tid; idx < CD; idx += NumThreads) { int c = idx / D; int k = idx % D; // Cross-chunk contribution: do[c] @ S_in[k]^T (dot over D) // S[k][d] = s_in_ptr[d*D + k] (since all_states stores S^T) float acc = 0.0f; for (int d = 0; d < D; ++d) { float do_val = (c < actual_len) ? bf16_to_f32(do_tile_ptr[c * D + d]) : 0.0f; acc += do_val * bf16_to_f32(s_in_ptr[d * D + k]); } // Within-chunk: sum_j dMqk[c][j] * ki[j][k] (j <= c since tril) for (int j = 0; j <= c && j < CHUNK; ++j) { acc += dMqk_smem[c * CHUNK + j] * ki_tile[j * D + k]; } dqd_local[idx] = acc; } __syncthreads();
// ========== Step 6: dki from Mqk backward ========== // dki[j][k] = sum_c dMqk^T[j][c] * qd[c][k] = sum_{c>=j} dMqk[c][j] * qd[c][k] float* dki_local = dqd_local; // reuse [C*D] for (int idx = tid; idx < CD; idx += NumThreads) { int j = idx / D; int k = idx % D; float acc = 0.0f; for (int c = j; c < CHUNK; ++c) { acc += dMqk_smem[c * CHUNK + j] * qd_tile[c * D + k]; } dki_local[idx] = acc; } __syncthreads();
// ========== Step 7: Use dL (computed in step 3b while dvcorr was alive) ========== // L = tril(kd @ ki^T, -1) * beta[:,None] // dF = dL (where F = tril(kd@ki^T, -1), and L = F * beta) // d(F*beta)/dF = dL * beta[i], d(F*beta)/dbeta[i] = sum_j dL[i][j]*F[i][j] // dkd_L[c][k] += sum_{j<c} (dL[c][j] * beta[c]) * ki[j][k] = sum_{j<c} dF[c][j] * ki[j][k] // dki_L[j][k] += sum_{c>j} (dL[c][j] * beta[c]) * kd[c][k] = sum_{c>j} dF^T[j][c] * kd[c][k]
// Compute dkd from L backward // Also accumulate dbeta from L: dbeta_L[c] = sum_{j<c} dL[c][j] * F[c][j] // F[c][j] = sum_k kd[c][k]*ki[j][k] for j < c
// dkd contributions from the L-backward (kd@S term): // dkd_from_vcorr[c][k] = -beta[c] * sum_d dvcorr[c][d] * S_in[k][d] // dkd_from_L[c][k] = sum_{j<c} dL[c][j]*beta[c] * ki[j][k] float* dkd_local = dki_local + CD; // need fresh buffer // Hmm, we're running low on scratch. Let me reconsider the layout. // Actually dki_local already reused dqd_local which reused dvcorr_smem. // Let me use a simpler approach: store dki first, then reuse for dkd.
// Store dki float* dki_out = bwd_ws_dki_ptr + int64_t(ws_idx) * CHUNK * D; for (int idx = tid; idx < CD; idx += NumThreads) { // Add dki from L backward: dki_L[j][k] = sum_{c>j} dL[c][j]*beta[c] * kd[c][k] int j = idx / D; int k = idx % D; float extra = 0.0f; for (int c = j + 1; c < CHUNK; ++c) { extra += dL_smem[c * CHUNK + j] * beta_smem[c] * kd_tile[c * D + k]; } float dki_total = dki_local[idx] + extra; dki_out[idx] = dki_total; } __syncthreads();
// Now compute dkd // dkd[c][k] = -beta[c] * sum_d dvcorr[c][d] * S_in[k][d] (from vcorr = (v-kd@S)*beta) // + sum_{j<c} dL[c][j]*beta[c] * ki[j][k] (from L = tril(kd@ki^T,-1)*beta) // dvcorr_smem was overwritten, but INV is still in smem (we used dMqk_smem for dL). // Recompute dvcorr on the fly from smem INV and dU_smem: float* dkd_out = bwd_ws_dkd_ptr + int64_t(ws_idx) * CHUNK * D; for (int idx = tid; idx < CD; idx += NumThreads) { int c = idx / D; int k = idx % D;
// dkd from vcorr backward: dvcorr[c][d] = sum_j INV^T[c][j] * dU[j][d] // INV is still in smem (preserved), dU_smem is still valid. float dvcorr_S_sum = 0.0f; for (int d = 0; d < D; ++d) { // Recompute dvcorr[c][d] = sum_j INV^T[c][j] * dU[j][d] float dvcorr_cd = 0.0f; for (int j = 0; j < CHUNK; ++j) { dvcorr_cd += inv_smem[j * CHUNK + c] * dU_smem[j * D + d]; // INV^T[c][j] = INV[j][c], from smem } dvcorr_S_sum += dvcorr_cd * bf16_to_f32(s_in_ptr[d * D + k]); // S[k][d] = s_in^T[d][k] }
// Store final dS (this is dS for the first chunk = d_initial_state) // dS_smem is [K, V] row-major: dS_smem[k*D + v] = dS[k][v] // d_initial_state layout is [V, K] (transposed state), so we store transposed: // dst[v*D + k] = dS[k][v] if (ds_out_ptr != nullptr) { int state_linear = seq_idx * H + head_idx; BF16* dst = ds_out_ptr + int64_t(state_linear) * D * D; for (int i = tid; i < dS_size; i += NumThreads) { int k = i / D; int v_idx = i % D; dst[v_idx * D + k] = BF16(dS_smem[i]); } } }