FlashKDA之backward(K2)

鱿鱼圈 Lv4

前置知识:KDA算法公式、triton、cuda、c++

K2_bwd 反向递推 kernel 逐行精讲

对应源码:csrc/smxx/bwd_kernel2.cuh_flash_kda_bwd_recurrence

目标读者:没接触过自动微分 / CUDA kernel 的人。读完应能自己把每一条梯度公式推一遍,并理解代码为什么这么写。

约定:d<量> 表示 loss 对该量的梯度。代码里的 exp(x) 实际是 (exp2),求导带的 因子被统一推迟到 K1_bwd 处理,本文遇到时会注明。公式用 LaTeX,伪代码用代码块。


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
#pragma once

#include "utils.cuh"

// ==================== Kernel 2 Backward: Reverse Recurrence ====================
//
// Grid: (N, H) — each block handles one sequence, one head
// Block: 256 threads (no warp specialization — simpler than fwd K2)
//
// Iterates chunks from back to front.
// Reads from: do [T,H,D], workspace (kd,qd,kr,ki,gT,INV,Mqk), v, beta, all_states
// Writes to: bwd_workspace (dkd,dqd,dki,dkr,dv,dbeta_chunk,dgT per tile)
// dS propagated in registers/smem across chunks
//
// Math per chunk (reverse order):
// Load: do, kd, qd, kr, ki, gT, INV, Mqk, v, beta, S_in from all_states
// Recompute: vcorr = (v - kd@S_in) * beta; U = INV @ vcorr
// dout = do (incoming grad)
// dqd_cross = dout @ S_in^T contribution to dqd from cross-chunk
// dS += qd^T @ dout cross-chunk state grad
// dMqk = dout @ U^T within-chunk attention grad
// dU_mqk = Mqk^T @ dout contribution to dU from Mqk
// dU_cross = kr @ dS contribution to dU from state update
// dkr = dU_total^T @ ... => actually kr^T @ U contributes to S_new
// ... (full derivation in comments below)
//
// dS_prev = exp(gT) * dS + ... (propagate backward)

template <int CHUNK, int D, int NumThreads, bool IsVarlen = true>
__global__ void __launch_bounds__(NumThreads) _flash_kda_bwd_recurrence(
// Input tensors (read-only)
cutlass::bfloat16_t const* __restrict__ do_ptr, // [T_total, H, D] (reordered as [H, T_total, D] for TMA compat)
cutlass::bfloat16_t const* __restrict__ v_ptr, // [H, T_total, D]
cutlass::bfloat16_t const* __restrict__ beta_ptr, // [H, T_total] (transposed)
float scale,
// Workspace from forward K1 (read-only, fp32 for backward precision)
float const* __restrict__ ws_kd_ptr, // [H*total_tiles, CHUNK, D] fp32
float const* __restrict__ ws_qd_ptr,
float const* __restrict__ ws_kr_ptr,
float const* __restrict__ ws_ki_ptr,
float const* __restrict__ ws_gt_ptr, // [H*total_tiles, D]
cutlass::bfloat16_t const* __restrict__ ws_inv_ptr, // [H*total_tiles, CHUNK, CHUNK]
cutlass::bfloat16_t const* __restrict__ ws_mqk_ptr,
// All states from forward K2
cutlass::bfloat16_t const* __restrict__ all_states_ptr, // [N*H*max_tiles, D, D]
// Backward workspace (write) — fp32 to preserve precision for dgc in K1_bwd
float* __restrict__ bwd_ws_dkd_ptr, // [H*total_tiles, CHUNK, D] fp32
float* __restrict__ bwd_ws_dqd_ptr,
float* __restrict__ bwd_ws_dki_ptr,
float* __restrict__ bwd_ws_dkr_ptr,
float* __restrict__ bwd_ws_dgt_ptr, // [H*total_tiles, D]
cutlass::bfloat16_t* __restrict__ bwd_ws_dv_ptr, // [H*total_tiles, CHUNK, D]
float* __restrict__ bwd_ws_dbeta_ptr, // [H*total_tiles, CHUNK]
// dS_init (from dfinal_state, or zero)
cutlass::bfloat16_t const* __restrict__ ds_init_ptr, // [N*H, D, D] or nullptr
// dS_out (d initial_state output)
cutlass::bfloat16_t* __restrict__ ds_out_ptr, // [N*H, D, D] or nullptr
// Dimensions
int T_total,
int H,
int N,
int64_t const* __restrict__ cu_seqlens,
int total_tiles
) {
using BF16 = cutlass::bfloat16_t;
constexpr int kWarpSize = 32;

int seq_idx = blockIdx.x;
int head_idx = blockIdx.y;
int tid = threadIdx.x;

int64_t bos, eos;
int tile_base;

if constexpr (IsVarlen) {
bos = cu_seqlens[seq_idx];
eos = cu_seqlens[seq_idx + 1];
tile_base = 0;
for (int i = 0; i < seq_idx; i++) {
tile_base += (int(cu_seqlens[i + 1] - cu_seqlens[i]) + CHUNK - 1) / CHUNK;
}
} else {
int T_seq = T_total / N;
bos = seq_idx * T_seq;
eos = bos + T_seq;
tile_base = seq_idx * ((T_seq + CHUNK - 1) / CHUNK);
}
int seq_len = int(eos - bos);
int t_tiles = (seq_len + CHUNK - 1) / CHUNK;

// --- Shared memory for dS accumulator [D, D] in fp32
// Also used for loading workspace tiles
extern __shared__ __align__(128) unsigned char shared_mem[];
float* dS_smem = reinterpret_cast<float*>(shared_mem); // [D*D]
// After dS: scratch for loading tiles
constexpr int 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);
constexpr int 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.

// Pointers for this tile's workspace data
float const* kd_tile = ws_kd_ptr + int64_t(ws_idx) * CHUNK * D;
float const* qd_tile = ws_qd_ptr + int64_t(ws_idx) * CHUNK * D;
float const* kr_tile = ws_kr_ptr + int64_t(ws_idx) * CHUNK * D;
float const* ki_tile = ws_ki_ptr + int64_t(ws_idx) * CHUNK * D;
float const* gt_tile = ws_gt_ptr + int64_t(ws_idx) * D;
BF16 const* inv_tile = ws_inv_ptr + int64_t(ws_idx) * CHUNK * CHUNK;
BF16 const* mqk_tile = ws_mqk_ptr + int64_t(ws_idx) * CHUNK * CHUNK;

// v and do tiles: layout is [H, T_total, D], tile starts at [head_idx, bos + t*CHUNK, 0]
int t_start = int(bos) + t * CHUNK;
BF16 const* v_tile_ptr = v_ptr + (int64_t(head_idx) * T_total + t_start) * D;
BF16 const* do_tile_ptr = do_ptr + (int64_t(head_idx) * T_total + t_start) * D;

// Beta: [H, T_total], linear index = head_idx * T_total + t_start
BF16 const* beta_tile_ptr = beta_ptr + int64_t(head_idx) * T_total + t_start;

// 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();

// vcorr[c][d] = (v[c][d] - kd_S[c][d]) * beta[c]
float* vcorr_smem = kd_S_smem; // reuse same buffer
for (int idx = tid; idx < CD; idx += NumThreads) {
int c = idx / D;
float v_val = (c < actual_len) ? bf16_to_f32(v_tile_ptr[c * D + idx % D]) : 0.0f;
float beta_c = beta_smem[c];
vcorr_smem[idx] = (v_val - kd_S_smem[idx]) * beta_c;
}
__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();

// ========== Step 4: Backward through vcorr = (v - kd@S) * beta ==========
// dv[c][d] = dvcorr[c][d] * beta[c]
// d(kd@S)[c][d] = -dvcorr[c][d] * beta[c]
// dbeta_contrib[c] = sum_d dvcorr[c][d] * (v[c][d] - kd_S[c][d])
// = sum_d dvcorr[c][d] * vcorr[c][d] / beta[c] (if beta != 0)

// 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();

// Store dqd
float* dqd_out = bwd_ws_dqd_ptr + int64_t(ws_idx) * CHUNK * D;
for (int idx = tid; idx < CD; idx += NumThreads) {
dqd_out[idx] = dqd_local[idx];
}

// ========== 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]
}

float dkd_val = -beta_smem[c] * dvcorr_S_sum;

// Add L backward contribution: sum_{j<c} dL[c][j]*beta[c] * ki[j][k]
for (int j = 0; j < c; ++j) {
dkd_val += dL_smem[c * CHUNK + j] * beta_smem[c] * ki_tile[j * D + k];
}

dkd_out[idx] = dkd_val;
}
__syncthreads();

// Add dbeta from L backward: sum_j dL[c][j] * F[c][j]
// F[c][j] = sum_k kd[c][k]*ki[j][k] for j < c
if (tid < CHUNK) {
int c = tid;
float dbeta_L = 0.0f;
for (int j = 0; j < c; ++j) {
float F_cj = 0.0f;
for (int k = 0; k < D; ++k) {
F_cj += kd_tile[c * D + k] * ki_tile[j * D + k];
}
dbeta_L += dL_smem[c * CHUNK + j] * F_cj;
}
// Add to existing dbeta
dbeta_out[c] += dbeta_L;
}
__syncthreads();

// Store dkr
float* dkr_out = bwd_ws_dkr_ptr + int64_t(ws_idx) * CHUNK * D;
for (int idx = tid; idx < CD; idx += NumThreads) {
dkr_out[idx] = dkr_smem[idx];
}

// ========== Step 8: Update dS for previous chunk ==========
// dS_prev = exp(gT) * dS + qd^T @ do
// From state update: S_new = S * exp(gT) + kr^T @ U
// dS += exp(gT) * dS_next (already in dS_smem)
// Also: dS += qd^T @ do (from out = qd @ S)
// And: dS -= beta * dvcorr @ kd (from vcorr = (v-kd@S)*beta, d(kd@S)/dS = kd^T)
// Wait, let me re-derive:
// vcorr = (v - kd@S)*beta => d_loss/dS from this path = -kd^T @ (dvcorr * beta_vec)
// But dvcorr already accounts for the beta multiplication? No:
// vcorr[c][d] = (v[c][d] - sum_k kd[c][k]*S[k][d]) * beta[c]
// d_loss/dS[k][d] = sum_c -kd[c][k] * beta[c] * d_loss/d_vcorr[c][d]
// = sum_c -kd[c][k] * dvcorr_beta[c][d]
// where dvcorr_beta[c][d] = beta[c] * d_vcorr/d_input? No, dvcorr is already the gradient w.r.t. vcorr.
// dS[k][d] += sum_c (-kd[c][k] * beta[c]) * dvcorr[c][d]

// For qd@S: dS[k][d] += sum_c qd[c][k] * do[c][d]

// For S * exp(gT): dS_prev[k][d] = exp(gT[k]) * dS_next[k][d]

// Recompute dvcorr for dS update from smem INV and dU_smem.
// dS[k][d] = exp(gT[k]) * dS_current[k][d]
// + sum_c qd[c][k] * do[c][d]
// + sum_c (-kd[c][k] * beta[c]) * dvcorr[c][d]

// ========== Step 8a: Compute dgT BEFORE updating dS ==========
// dgT has TWO contributions:
// 1) State update: S_new = S_in * exp2(gT) + kr^T @ U
// => dgT_state[k] = exp2(gT[k]) * sum_d dS_next[k][d] * S_in[k][d]
// 2) kr dependency: kr = kn * exp2(gT - gc), so d(kr)/d(gT) = kr * ln2
// => dgT_kr[k] = sum_c dkr[c][k] * kr[c][k] (without ln2, matching dgc convention)
// Note: dgc already has -dkr*kr for d(kr)/d(gc), but d(kr)/d(gT) is +dkr*kr
float* dgT_out = bwd_ws_dgt_ptr + int64_t(ws_idx) * D;
for (int k = tid; k < D; k += NumThreads) {
float gT_k = gT_smem[k];
// Contribution 1: state update
float dgT_k = 0.0f;
for (int d = 0; d < D; ++d) {
// S[k][d] = s_in_ptr[d*D + k] (since all_states stores S^T)
dgT_k += dS_smem[k * D + d] * bf16_to_f32(s_in_ptr[d * D + k]);
}
dgT_k *= gT_k;

// Contribution 2: kr dependency — sum_c dkr[c][k] * kr[c][k]
for (int c = 0; c < CHUNK; ++c) {
dgT_k += dkr_smem[c * D + k] * kr_tile[c * D + k];
}

dgT_out[k] = dgT_k;
}
__syncthreads();

// ========== Step 8b: Update dS for previous chunk ==========
for (int idx = tid; idx < dS_size; idx += NumThreads) {
int k = idx / D;
int d = idx % D;

float dS_val = gT_smem[k] * dS_smem[idx]; // gT_smem already stores exp2(gT)

// qd^T @ do contribution
float qd_do = 0.0f;
for (int c = 0; c < actual_len; ++c) {
float qd_ck = qd_tile[c * D + k];
float do_cd = bf16_to_f32(do_tile_ptr[c * D + d]);
qd_do += qd_ck * do_cd;
}
dS_val += qd_do;

// -kd^T @ (beta * dvcorr) contribution
float kd_dv = 0.0f;
for (int c = 0; c < CHUNK; ++c) {
float kd_ck = kd_tile[c * D + k];
// Recompute dvcorr[c][d] = sum_j INV[j][c] * 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]; // from smem, not bf16 gmem
}
kd_dv += kd_ck * beta_smem[c] * dvcorr_cd;
}
dS_val -= kd_dv;

dS_smem[idx] = dS_val;
}
__syncthreads();
}

// 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]);
}
}
}

0. 先建立全局图景

0.1 FlashKDA 把序列切块做线性注意力

序列按 CHUNK 长度切块。块内用注意力,块间用一个状态矩阵 做递推。前向每块做:

反向就是对这一串求伴随(adjoint,即反向传播)。 表示按行/按元素广播乘。

0.2 形状与记号

符号 形状 含义
CHUNK () 标量 块长,例 16
标量 head 维度,例 128
进入本块时的状态。存储是转置的,见 0.3
前向 K1 生成的派生量(fp32)
整块门控总和,workspace 里存
每行一个门控标量(已过 sigmoid)
,下三角
输入 是输出 的上游梯度
中间量(前向没存,反向重算)
跨块传播的状态梯度累加器(常驻 smem)

0.3 状态的转置存储(极其重要,否则下标全错)

all_states / ds_init / ds_out 在显存里存的是 ,布局 。即:

源码中凡出现 s_in_ptr[d*D + k],语义都是数学上的

0.4 为什么必须从后往前

块入口的状态梯度 依赖第 块算出的 。所以反向必须 从最后一块往第一块 迭代(for t = t_tiles-1 ... 0)。


1. Kernel 启动配置与线程模型(line 29-70)

1
2
template <int CHUNK, int D, int NumThreads, bool IsVarlen = true>
__global__ void __launch_bounds__(NumThreads) _flash_kda_bwd_recurrence(...)
  • Grid = (N, H)blockIdx.x = seq_idxblockIdx.y = head_idx。每个 block 独占一条序列的一个 head。
  • Block = 256 线程,无 warp specialization(注释 line 8)。
  • 纯标量的 shared-memory GEMM 循环,不用 tensor core;逻辑直白但有冗余重算(用算力换显存)。

2. 序列定位(line 72-89)

seq_idx 映射到 token 维范围 和全局 tile 偏移 tile_base

  • Varlen:用 cu_seqlensbos/eostile_base 累加前面序列的块数。
  • 定长:直接乘除。

3. 共享内存布局与 dS 初始化(line 91-117)

1
2
float* dS_smem     = (float*)shared_mem;   // [D*D], 整个 kernel 常驻
float* scratch_f32 = dS_smem + D*D; // 后面所有临时 tile 堆这里

初值来自 dfinal_state(对最终状态的梯度),没有则为 0。转置加载:


4. 主循环:加载本块数据(line 120-177)

1
2
3
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); // 尾块可能不满
  • actual_len:尾块真实行数;越界行读 时取 0。
  • (已 exp2)、(读原始再过 sigmoid_tanh_approx)、 搬进 smem。
  • scratch 区堆叠顺序(贯穿全 kernel):
1
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]

多数 buffer 复用别名,唯独 dL 独立(见第 8 节)。


5. Step 1:重算 U(line 179-235)

前向的 未存,反向重算。

5.1 的 次乘加,最重。vcorr_smem 复用 kd_S_smem


6. Step 2:输出梯度反传(line 237-303)

前向 ,分两路。

6.1 dU 来自 路(line 251-267)

,对 求偏导得

6.2 dU 加状态路贡献(line 269-287)

也影响 loss:

其中 后一块传回的状态梯度。

6.3 dkr(line 289-303)

对同一式中的 求偏导得


7. Step 3:三角求解的伴随(line 305-321)—— 重点之一

前向 ,等价解

本节是全文最难的一处。下面 7.0 是写给线性代数基础薄弱读者的从零详解,已经熟悉矩阵微分的读者可直接跳到 7.1 看结论。

7.0 从零搭起:只用一条核心规则

(a) 三个积木

积木 1 — 矩阵乘法 的第 行第 列 = “A 的第 行"点乘"X 的第 列”。

一句话:matmul 就是一堆点积。

积木 2 — 转置 :行列互换,(沿对角线翻一下)。

积木 3 — 逆矩阵 :矩阵版的"倒数",定义 是单位矩阵,相当于数字 1)。你不需要会手算它,只要知道:

即"解线性方程组"和"求逆再乘"是同一件事。本例 就等价于解

(b) 反向传播唯一要背的核心规则

其中 是"loss 对 的梯度"(与 同形状)。推导(只用积木 1 的点积定义 + 链式法则):

转置就是这么冒出来的。 口诀:求输入 的梯度,就把另一个因子 转置乘到 左边;求权重 的梯度,就把 转置乘到 右边。

© 标量类比(建立"为什么有逆/转置"的直觉)

把矩阵换成普通数字,看除法

记住这两条标量结果,矩阵版长得一模一样(只是 、注意左右与转置)。

(d) 第一步:求 —— 直接套核心规则

前向 就是 ,对应 。求"对输入 的梯度":

(e) 第二步:求 —— 需要"逆矩阵的梯度"这块拼图

藏在 里,要多走一步。

先用核心规则求"对权重 的梯度":

再用逆矩阵反向规则(标量 的矩阵版,推导见 (g)):

,且

(f) 化简到代码里的简洁式

代入并重新分组(结合律):

(用了 ,以及 。)于是得到

又和标量版结构完全对应。

(g)(选读)逆矩阵反向规则的来历

求微分:;再转成伴随(内积配对,见 7.1)得

小结:第一步 ,第二步 。全程只用了"matmul 反向"这一条核心规则 + 一次"逆矩阵反向"。和标量 一一对应。

7.1 通用公式(务必记住)

,已知 (即 ):

推导(内积配对)。对 取全微分,用

loss 一阶变化 ,其中

7.2 套进本例(


8. Step 3b:dL(line 324-338)—— 重点之二

由 7.1 的 ,且

(这条公式的完整从零推导见 7.0(e)(f);直觉:它是"输出敏感度 "与"解 "的外积、带负号,正对应标量 。)

严格下三角(上三角与对角前向被 mask 成 0,那些位置根本不是自由参数),故梯度也 mask:

工程关键dL 依赖 dvcorr,而 dvcorr_smem 在 Step 5 会被覆盖。所以必须趁 dvcorrU 还在时立刻把 dL 算进独立不别名dL_smem(line 309 注释 “not aliased”)。这就是它被命名为 “Step 3b”、紧贴 Step 3a 之后的原因。


9. Step 4:vcorr 反传到 dv、dbeta(line 341-382)

前向

代码仅在 时做除法(line 375)。


10. Step 5-6:dMqk → dqd、dki(line 384-451)

前向

dqd 两路(line 411-437):

dki 来自 Mqk 路(line 439-451):


11. Step 7:L 路反传 → dki、dkd、dbeta(line 453-535)

前向 。记 ,则

dki 加 L 贡献(line 472-485):

dkd 两路(line 487-518):

此时 dvcorr_smem 已被覆盖,但 inv_smemdU_smem 还在,于是就地重算 (line 500-505)。

dbeta 加 L 贡献(line 520-534):


12. Step 8:dgT 与 dS 递推更新(line 543-625)

12.1 dgT(必须在更新 dS 之前算,line 567-592)

有两个去处:状态衰减 ,以及

代码里 gT_smem[k] 已是 。用的是当前(后块传来) dS,故必须先于 12.2。

下面 12.1.0 是从零详解( 是个长度 的向量, 扇出到很多地方,逐路收集);已熟悉的读者可跳过。

12.1.0 从零推导 dgT

(a) 先看 在前向被用在哪。 是长度 的向量(每个 head 维一个数)。 出现在两个地方,所以它的总梯度 = 两路回传之和(扇出 → 求和,多元链式法则):

  • 路 1 — 状态衰减(注意只有,因为衰减是按 这一维广播的)。
  • 路 2 — 进入 (每个 都用到同一个 )。

(b) 路 1 的偏导。 固定 影响第 行所有列 。用指数求导 (这里同样把公因子 按"dgT 约定"省略,留到 K1_bwd 统一处理):

乘上游 并对所有列 求和(因为 影响了整行):

代码里 gT_smem[k] 已经是 ,所以源码先 sum_d dS*S_in*= gT_k(line 581-583)。

© 路 2 的偏导。 出现在每个 里,同样 (因 )。乘上游 并对所有行 求和:

这里的 是本块第 6.3 节刚算出的(dkr_smem), 从 workspace 读。

(d) 两路相加。

(e) 为什么必须"先于 12.2"。 路 1 用的 进入本块时(后一块传回来)的 ;而 12.2 会把 dS_smem 原地覆盖成"前一块的 "。若先跑 12.2,路 1 就会读到错的 。所以源码严格按 Step 8a (dgT) → Step 8b (更新 dS) 的顺序(line 567 注释 “compute dgT BEFORE updating dS”)。

符号小提醒 都来自门控, 的指数是 。所以 增大让 变大( 号,本节路 2);而 增大让 变小( 号,那部分梯度走 ,在 K1_bwd 处理,见 3.2(b))。同一个 贡献符号相反,别搞混。

12.2 dS 反向递推(line 594-625)

更新为前一块的状态梯度,三项相加:

三项偏导依据:

dvcorr 就地重算。写回 dS_smem,进入更前一块。


13. 收尾:写出 d_initial_state(line 628-641)

循环结束时 dS_smem 即第一块入口的状态梯度。转置回 写出:


14. 全部梯度公式速查表


15. 设计要点总结

主题 做法
从后往前 依赖后块,必须逆序迭代
重算换显存 不存,靠 在 smem 重算
状态转置存储 all_states/ds_init/ds_out 都是 ;取 s_in_ptr[d*D+k]
buffer 复用与别名 scratch 严格堆叠;dMqk 复用 Mqk,dqd/dki 复用 dvcorr;唯 dL 独立
gT 已是 exp2 workspace 存
尾块掩码 actual_len 控制越界行读 0
数值 workspace fp32;dbeta 除 时加

配套阅读:K1_bwd 的逐行精讲见 k1_bwd_prepare_deep_dive.md。本 kernel 产出的 正是 K1_bwd 的输入。

  • 标题: FlashKDA之backward(K2)
  • 作者: 鱿鱼圈
  • 创建于 : 2026-06-20 23:50:00
  • 更新于 : 2026-06-22 21:17:09
  • 链接: https://yuyanqi.com/2026/06/20/FlashKDA之backward(K2)/
  • 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。
评论