FlashKDA之backward(K1)

鱿鱼圈 Lv4

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

K1_bwd “prepare” kernel 逐行精讲

对应源码:csrc/smxx/bwd_kernel1.cuh_flash_kda_bwd_prepare

目标读者:没接触过自动微分 / CUDA kernel 的人。读完应能自己推出 L2 归一化反向、数值稳定的 、以及门控的逆向 cumsum,并理解代码为何这样写。

约定:d<量> 表示 loss 对该量的梯度。exp(x) 实际是 (exp2)。公式用 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
#pragma once

#include "utils.cuh"

// ==================== Kernel 1 Backward: Prepare Backward ====================
//
// Grid: (total_tiles, H) — each CTA handles one chunk's one head
// Block: 256 threads (simple, no warp specialization)

template <int CHUNK, int D, int NumThreads, bool IsVarlen = true>
__global__ void __launch_bounds__(NumThreads) _flash_kda_bwd_prepare(
// Forward inputs (for backward computation)
cutlass::bfloat16_t const* __restrict__ q_ptr, // [H, T_total, D]
cutlass::bfloat16_t const* __restrict__ k_ptr, // [H, T_total, D]
cutlass::bfloat16_t const* __restrict__ g_bf16_ptr, // [H, T_total, D] (raw gate, bf16)
cutlass::bfloat16_t const* __restrict__ beta_ptr, // [H, T_total] (transposed, pre-sigmoid)
float const* __restrict__ A_log_ptr, // [H]
float const* __restrict__ dt_bias_ptr, // [H, D]
// Forward workspace (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] (g_total fp32)
float const* __restrict__ ws_gc_ptr, // [H*total_tiles, CHUNK, D] (gate cumsum fp32)
// Backward workspace from K2_bwd (fp32 for precision)
float const* __restrict__ bwd_ws_dkd_ptr,
float const* __restrict__ bwd_ws_dqd_ptr,
float const* __restrict__ bwd_ws_dki_ptr,
float const* __restrict__ bwd_ws_dkr_ptr,
float const* __restrict__ bwd_ws_dgt_ptr, // [H*total_tiles, D]
cutlass::bfloat16_t const* __restrict__ bwd_ws_dv_ptr, // [H*total_tiles, CHUNK, D]
float const* __restrict__ bwd_ws_dbeta_ptr, // [H*total_tiles, CHUNK]
// Output gradients
cutlass::bfloat16_t* __restrict__ dq_ptr, // [H, T_total, D]
cutlass::bfloat16_t* __restrict__ dk_ptr,
cutlass::bfloat16_t* __restrict__ dv_ptr,
cutlass::bfloat16_t* __restrict__ dg_ptr, // [H, T_total, D] (raw gate grad)
cutlass::bfloat16_t* __restrict__ dbeta_ptr, // [H, T_total] (logit space grad)
float* __restrict__ dA_log_ptr, // [H] (atomicAdd)
float* __restrict__ ddt_bias_ptr, // [H, D] (atomicAdd)
// Dimensions
float scale,
float gate_scale,
int T_total,
int H,
int N,
int64_t const* __restrict__ cu_seqlens,
int total_tiles
) {
using BF16 = cutlass::bfloat16_t;
constexpr int CD = CHUNK * D;
constexpr int ELEMS_PER_THREAD = 8;
constexpr int THREADS_PER_ROW = D / ELEMS_PER_THREAD; // 16

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

// Thread layout matching forward K1 (16 threads per row, 8 elems each)
int my_row = tid / THREADS_PER_ROW; // 0..15
int my_col = (tid % THREADS_PER_ROW) * ELEMS_PER_THREAD; // 0, 8, 16, ..., 120

// --- Resolve tile → sequence mapping
int seq_idx, local_t;
int64_t bos, eos;
int seq_len, t_tiles_this_seq;

if constexpr (IsVarlen) {
seq_idx = 0;
int tiles_before = 0;
for (int i = 0; i < N; i++) {
int slen = int(cu_seqlens[i + 1] - cu_seqlens[i]);
int n_tiles = (slen + CHUNK - 1) / CHUNK;
if (tiles_before + n_tiles > global_tile_idx) {
seq_idx = i;
break;
}
tiles_before += n_tiles;
}
local_t = global_tile_idx - tiles_before;
bos = cu_seqlens[seq_idx];
eos = cu_seqlens[seq_idx + 1];
} else {
int T_seq = T_total / N;
int tiles_per_seq = (T_seq + CHUNK - 1) / CHUNK;
seq_idx = global_tile_idx / tiles_per_seq;
int tiles_before = seq_idx * tiles_per_seq;
local_t = global_tile_idx - tiles_before;
bos = seq_idx * T_seq;
eos = bos + T_seq;
}
seq_len = int(eos - bos);
t_tiles_this_seq = (seq_len + CHUNK - 1) / CHUNK;
if (local_t >= t_tiles_this_seq) return;

int ws_idx = head_idx * total_tiles + global_tile_idx;
int t_start = int(bos) + local_t * CHUNK;
int actual_len = min(CHUNK, seq_len - local_t * CHUNK);

float a_log_exp = expf(A_log_ptr[head_idx]);

// Shared memory: dgc[C*D]
extern __shared__ __align__(128) unsigned char shared_mem[];
float* dgc_smem = reinterpret_cast<float*>(shared_mem);

BF16 const* g_tile_ptr = g_bf16_ptr + (int64_t(head_idx) * T_total + t_start) * D;
BF16 const* q_tile_ptr = q_ptr + (int64_t(head_idx) * T_total + t_start) * D;
BF16 const* k_tile_ptr = k_ptr + (int64_t(head_idx) * T_total + t_start) * D;

float const* gc_tile = ws_gc_ptr + int64_t(ws_idx) * CHUNK * D;
float const* gt_tile = ws_gt_ptr + int64_t(ws_idx) * D;
float const* dkd_tile = bwd_ws_dkd_ptr + int64_t(ws_idx) * CHUNK * D;
float const* dqd_tile = bwd_ws_dqd_ptr + int64_t(ws_idx) * CHUNK * D;
float const* dki_tile = bwd_ws_dki_ptr + int64_t(ws_idx) * CHUNK * D;
float const* dkr_tile = bwd_ws_dkr_ptr + int64_t(ws_idx) * CHUNK * D;
float const* dgt_tile = bwd_ws_dgt_ptr + int64_t(ws_idx) * D;
BF16 const* dv_tile = bwd_ws_dv_ptr + int64_t(ws_idx) * CHUNK * D;
float const* dbeta_tile = bwd_ws_dbeta_ptr + int64_t(ws_idx) * CHUNK;

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* ki_tile = ws_ki_ptr + int64_t(ws_idx) * CHUNK * D;
float const* kr_tile = ws_kr_ptr + int64_t(ws_idx) * CHUNK * D;

BF16* dq_out = dq_ptr + (int64_t(head_idx) * T_total + t_start) * D;
BF16* dk_out = dk_ptr + (int64_t(head_idx) * T_total + t_start) * D;

// ========== Phase 1+2: Fused dgc + dq/dk via L2 norm backward ==========
// Numerically stable dgc formula:
// dgc = kn * (dkd*exp(gc) - dki*exp(-gc) - dkr*exp(gT-gc)) + qn*scale*dqd*exp(gc)
// This avoids catastrophic cancellation in the original dkd*kd + dqd*qd - dki*ki - dkr*kr
// because dkd*exp(gc), dki*exp(-gc), dkr*exp(gT-gc) are all O(1) magnitude,
// while kd, ki can differ by ~1e20 making the original formula lose precision.
{
float q_vals[ELEMS_PER_THREAD], k_vals[ELEMS_PER_THREAD];
float q_sq = 0.0f, k_sq = 0.0f;

#pragma unroll
for (int i = 0; i < ELEMS_PER_THREAD; ++i) {
float qv = (my_row < actual_len) ? bf16_to_f32(q_tile_ptr[my_row * D + my_col + i]) : 0.0f;
float kv = (my_row < actual_len) ? bf16_to_f32(k_tile_ptr[my_row * D + my_col + i]) : 0.0f;
q_vals[i] = qv;
k_vals[i] = kv;
q_sq += qv * qv;
k_sq += kv * kv;
}

// 16-thread reduction matching forward K1 exactly
#pragma unroll
for (int delta = 8; delta >= 1; delta >>= 1) {
q_sq += __shfl_xor_sync(0xFFFFFFFF, q_sq, delta);
k_sq += __shfl_xor_sync(0xFFFFFFFF, k_sq, delta);
}

float q_inv_norm = rsqrtf(q_sq + 1e-6f);
float k_inv_norm = rsqrtf(k_sq + 1e-6f);

// Compute dqn, dkn, dgc using fp32 gc from workspace
float dqn_vals[ELEMS_PER_THREAD], dkn_vals[ELEMS_PER_THREAD];
float qn_vals[ELEMS_PER_THREAD], kn_vals[ELEMS_PER_THREAD];

#pragma unroll
for (int i = 0; i < ELEMS_PER_THREAD; ++i) {
int col = my_col + i;
float gc_v = gc_tile[my_row * D + col];
float exp_gc = ex2_approx_ftz_f32(gc_v);
float exp_neg_gc = ex2_approx_ftz_f32(-gc_v);
float exp_gt_gc = gt_tile[col] * exp_neg_gc;

float dqd_v = dqd_tile[my_row * D + col];
float dkd_v = dkd_tile[my_row * D + col];
float dki_v = dki_tile[my_row * D + col];
float dkr_v = dkr_tile[my_row * D + col];

// dqn, dkn (for L2 norm backward)
float dqn_v = dqd_v * exp_gc * scale;
float dkn_v = dkd_v * exp_gc + dki_v * exp_neg_gc + dkr_v * exp_gt_gc;

float qn_v = q_vals[i] * q_inv_norm;
float kn_v = k_vals[i] * k_inv_norm;

// Numerically stable dgc:
// dkn_signed = dkd*exp(gc) - dki*exp(-gc) - dkr*exp(gT-gc)
// Each term is O(1), subtraction is well-conditioned
float dkn_signed = dkd_v * exp_gc - dki_v * exp_neg_gc - dkr_v * exp_gt_gc;
dgc_smem[my_row * D + col] = kn_v * dkn_signed + qn_v * dqn_v;

dqn_vals[i] = dqn_v;
dkn_vals[i] = dkn_v;
qn_vals[i] = qn_v;
kn_vals[i] = kn_v;
}

float dot_dqn_qn = 0.0f, dot_dkn_kn = 0.0f;
#pragma unroll
for (int i = 0; i < ELEMS_PER_THREAD; ++i) {
dot_dqn_qn += dqn_vals[i] * qn_vals[i];
dot_dkn_kn += dkn_vals[i] * kn_vals[i];
}
#pragma unroll
for (int delta = 8; delta >= 1; delta >>= 1) {
dot_dqn_qn += __shfl_xor_sync(0xFFFFFFFF, dot_dqn_qn, delta);
dot_dkn_kn += __shfl_xor_sync(0xFFFFFFFF, dot_dkn_kn, delta);
}

if (my_row < actual_len) {
#pragma unroll
for (int i = 0; i < ELEMS_PER_THREAD; ++i) {
int col = my_col + i;
float dq_v = (dqn_vals[i] - qn_vals[i] * dot_dqn_qn) * q_inv_norm;
float dk_v = (dkn_vals[i] - kn_vals[i] * dot_dkn_kn) * k_inv_norm;
dq_out[my_row * D + col] = BF16(dq_v);
dk_out[my_row * D + col] = BF16(dk_v);
}
}
}
__syncthreads();

// ========== Phase 3: dv — copy from bwd workspace ==========
BF16* dv_out = dv_ptr + (int64_t(head_idx) * T_total + t_start) * D;
for (int idx = tid; idx < actual_len * D; idx += NumThreads) {
dv_out[idx] = dv_tile[idx];
}

// ========== Phase 4: Reverse cumsum of dgc + dgT → gate backward ==========
BF16* dg_out = dg_ptr + (int64_t(head_idx) * T_total + t_start) * D;

for (int col = tid; col < D; col += NumThreads) {

float dt = dt_bias_ptr[head_idx * D + col];
float rev_sum = 0.0f;
float dgT_val = dgt_tile[col];

float dA_log_acc = 0.0f;
float ddt_bias_acc = 0.0f;

for (int row = CHUNK - 1; row >= 0; --row) {
rev_sum += dgc_smem[row * D + col];
float dg_nat = rev_sum + dgT_val;

if (row < actual_len) {
float g_raw = bf16_to_f32(g_tile_ptr[row * D + col]);
float z = a_log_exp * (g_raw + dt);
float sig = sigmoid_tanh_approx_f32(z);
float dsig = sig * (1.0f - sig);
constexpr float kLn2 = 0.6931471805599453f;
float dz = dg_nat * gate_scale * dsig * kLn2;
float dg_raw = dz * a_log_exp;

dg_out[row * D + col] = BF16(dg_raw);

ddt_bias_acc += dz * a_log_exp;
dA_log_acc += dz * (g_raw + dt);
} else {
dg_out[row * D + col] = BF16(0.0f);
}
}

atomicAdd(&ddt_bias_ptr[head_idx * D + col], ddt_bias_acc);
atomicAdd(&dA_log_ptr[head_idx], dA_log_acc * a_log_exp);
}

// ========== Phase 5: dbeta — convert from chunk-space to logit-space ==========
BF16 const* beta_tile_ptr = beta_ptr + int64_t(head_idx) * T_total + t_start;
BF16* dbeta_out = dbeta_ptr + int64_t(head_idx) * T_total + t_start;
if (tid < actual_len) {
float dbeta_c = dbeta_tile[tid];
float beta_raw = bf16_to_f32(beta_tile_ptr[tid]);
float sig = sigmoid_tanh_approx_f32(beta_raw);
float dbeta_logit = dbeta_c * sig * (1.0f - sig);
dbeta_out[tid] = BF16(dbeta_logit);
}
}

0. 它在整条反向链里的位置

1
Backward:  K2_bwd (反向递推, 后->前)  ->  K1_bwd (本文, 还原输入梯度)

K2_bwd 算出了一堆"派生量的梯度" 。 K1_bwd 是前向 K1 的逆过程:把这些梯度还原回原始输入的梯度—— (门)、,以及参数

0.1 前向 K1 做了什么(理解反向的前提)

设每行(一个 token)的原始 ,门控相关原始量 、参数 。前向:

所以反向要做三件事:把 收拢成 ;把 过 L2 反向得 ;把 (+)过 cumsum 和门函数得


1. 线程模型(line 10-62)

1
2
constexpr int ELEMS_PER_THREAD = 8;
constexpr int THREADS_PER_ROW = D / 8; // = 16
  • Grid = (total_tiles, H)blockIdx.x = global_tile_idxblockIdx.y = head_idx。每个 CTA 处理一个 chunk 的一个 head(line 7)。
  • 每行 16 个线程,每线程负责 8 个元素

这个布局和前向 K1 完全一致(line 60),保证归约顺序相同 → 数值可复现。


2. tile → 序列定位(line 64-99)

global_tile_idx 映射到 seq_idx / local_t,并算出:

越界 tile(local_t >= t_tiles_this_seq)直接 return(line 95)。 a_log_exp = exp(A_log[head])(line 101)。

shared memory 只有一块:dgc_smem[C*D](line 105)。


3. Phase 1+2:融合算 dgc 与 dq/dk(line 129-218)

这是全 kernel 最核心、最讲数值稳定的部分。

3.1 重算 L2 范数(line 136-157)

每线程读自己负责的 8 个 元素,累加平方和;尾块越界行读 0;再用 16 线程的 __shfl_xor 归约()把整行平方和汇总:

归约方式必须与前向逐位一致。

3.2 重算门控指数、算 dqn/dkn 与 dgc(line 163-193)

对每个元素(列 ):

(代码里 gt_tile[col] 存的就是 ,所以 exp_gt_gc = gt_tile[col]*exp_neg_gc。)

(a) dqn / dkn(供 L2 反向)

前向 ,所以对

同时进了 三个分支,梯度相加:

(b) dgc —— 数值稳定写法(关键 trick)

下面先给从零详解(b.0–b.4),已熟悉的读者可直接看 b.3 的结论框。

(b.0) 先搞清 是什么、要求什么

是"门控累加和",一个标量(对每行每维)。前向它出现在四个派生量的指数里:

上游已经给了这四个量的梯度 。一个变量被用到多处(扇出 / fan-out),它的总梯度 = 各处回传之(多元链式法则)。所以:

(b.1) 求每个偏导(指数函数求导)

对指数函数 。这里底数 ,所以 。本 kernel 约定把公共因子 推迟到 Phase 4 统一乘(因为 对所有项线性,提一个常数出来最后乘即可),故这里先省略 。导数符号跟着指数的 号走:

指数是 ,链式带出一个负号。)代回得天真公式(数学正确):

(b.2) 为什么天真公式在 fp32 里会崩(灾难性抵消)

。当 较大时这两者量级天差地别。举个具体例子,设

两者相差约 。fp32 只有约 7 位十进制有效数字( 相对精度)。一旦把 (可能上千万量级)和 (极小)放进同一个加减式:

  • 大 + 小:小项落在大项的有效位之外,被直接舍掉 → 信息丢失;
  • 大 − 大:若两个大项接近,相减后高位全抵消、只剩低位舍入噪声 → 灾难性抵消(catastrophic cancellation)

结果 基本是垃圾。注释 line 134 说的 “kd, ki can differ by ~1e20” 就是这个意思。

(b.3) 稳定写法:把 提到括号外

注意 ,对 同理。把公共的 (和 )提出来:

其中 。对应源码 line 186-187。

(b.4) 为什么这样就稳了

关键观察:括号里的乘积 本身就是 (正常大小)。原因是上游梯度 是按 的比例传下来的(前向 ,反向链式正好带一个 ),二者相乘把 抵消掉。于是:

  • 括号内是"几个正常大小的数相减" → 良态,没有大数吞小数;
  • 再乘一个 (单位向量的分量)→ 仍然正常。

全程不出现 这种怪物。对比一句话:

两者数学恒等,只是运算顺序不同,浮点结果天差地别。这是典型的"重新结合以避免中间量爆炸"的数值技巧。

注意区分两个量(来源相同、符号不同):

不带 的符号(前向 是线性因子,故 全加);而 指数 号(故 带符号)。两者只差中间两项的正负号,别混用。

3.3 L2 归一化的反向 → dq/dk(line 195-216)

前向 要求 已知

下面 3.3.0 是写给线代基础薄弱读者的从零详解;已熟悉的读者可直接看 3.3.1 的结论。

3.3.0 从零推导

为什么不简单:若 这种逐元素映射,求导就完事。但归一化里 装着所有分量 ,所以改动任一 会通过 影响每一个 。这就是最终会冒出耦合项 的根源。

第一块: 的导数。。只有第 项含 ,故 ,再用 的链式:

第二块: 的导数(核心)。 ,用乘积法则,(克罗内克符号, 为 1 否则 0),(代入 ★):

两部分含义: 是"自己对自己"的直接项(仅 ); 是通过 产生的"所有分量互相耦合"项( 任意都有)。

第三块:链式法则汇总。 影响所有 ,把贡献全加起来,并代入 (☆):

第一项里 使求和塌缩成单项

第四块:用 化简。 关键代换

代回 (◇):

3.3.1 结论

对应代码:

点积 用 16 线程 shfl 归约出整行(line 195-205),仅 my_row < actual_len 才写出(line 207)。这个归约就是在算第二块说的"耦合标量"——所以归一化反向必须先做一次行内归约,无法纯逐元素

直觉 是单位向量,沿 方向拉伸 不改变 (只改长度),所以梯度里"沿 的径向分量"无效。 上的投影长度, 是该径向分量,减掉它只留切向分量;再乘 越长,同样的 影响越小)。一句话:把梯度投影到 方向,再按 长度 缩放

3.3.2 小数字验证(2 维)

,则 。设

自检 恒垂直于 (沿 推不改变 ):。一般地

可作为写单测时的快速断言。


4. Phase 3:dv 直接搬运(line 220-224)

在 K2_bwd 已算好(bf16),这里只是从 workspace 的 tile 布局拷到最终 布局,仅拷 actual_len*D 个元素。


5. Phase 4:门控反向 = dgc 的逆向 cumsum + dgT(line 226-261)

下面 5.0 是写给基础薄弱读者的从零详解(核心就一句:前缀和的反向是后缀和);已熟悉的读者可直接看 5.1。

5.0 从零理解"cumsum 的反向"

(a) 什么是 cumsum(前缀和)

前向把一串数 累加成"到当前为止的和":

展开看就是(以 为例):

(b) 关键观察:每个输入"扇出"到哪些输出

竖着看上面这张表: 出现在 行及其下面所有行里。比如 出现在 (不在 )。用偏导写:

© 反向:把扇出的梯度收回来 = 后缀和

多元链式法则: 被用到多处,它的总梯度 = 这些输出回传的梯度之,每个偏导都是 1,所以就是直接相加

即"从 到末尾"的求和——这叫后缀和(suffix sum)。还是

一句话规律:前向是前缀和(从前往后累加),反向就是后缀和(从后往前累加)。 方向正好反过来,所以叫"逆向 cumsum"。

(d) 怎么 算出来:一个滚动变量

不用对每个 都重新求和(那是 )。从最后一行往前走,维护一个累加器 rev_sum,每步把当前行的 加进去:

1
2
3
4
rev_sum = 0
for row = C-1 down to 0: // 从后往前
rev_sum += dgc[row] // 此刻 rev_sum == sum_{r>=row} dgc[r]
dg_total[row] = rev_sum

走到 row 时,rev_sum 恰好就是 ,正是我们要的后缀和。对应源码 line 237-238。

(e) 数字小例(

(下标 0…3)。从后往前滚:

步 row 加入 rev_sum
3 3 3 3
2 1 4 4
1 5 9 9
0 2 11 11

可手动验证 ✓。

(f) 再叠加 那条路

其实还流进了整块总和 (见 5.1)。所有 ,所以每行都再加一份相同的

小结:cumsum 反向 = 后缀和(一个滚动累加器搞定); 再贡献一个对每行都相同的常数。下面 5.1 把这两件事正式写出来。

5.1 g_total 扇出到两个去处

前向 同时进入两处(计算图里是并列两条边):

多元链式法则:一个量被用到多处(扇出),其总梯度 = 各下游路径回传之

  • gc 这条:因 仅当 ,所以是"后缀和"。代码从 row=C-1 往下走,维护 rev_sum += dgc[row],到 rev_sum 正是 (这就是"逆向 cumsum")。
  • gT 这条:因 所有 ,所以每行都加同一个常数 dgT_val(line 232 循环外读一次,line 239 每行加)。

(此处仍是"未乘 "约定值。)

5.2 过门函数反传到 dg_raw 与参数(line 241-260)

(=dg_nat)沿 往回传,并补上延后的

下面 5.2.0 是从零详解(单变量链式法则 + sigmoid 求导 + 一个量被多处使用怎么办);已熟悉的读者可直接看 (a)(b)©。

5.2.0 从零理解这条链

(i) 前向是一串"嵌套函数"。 从最里到最外:

原始输入 和两个参数 都在这条链上。反向就是从右端的已知梯度 一站一站往左乘"每站的局部导数"。

(ii) 单变量链式法则(唯一要会的)。,则 。也就是"上游梯度 × 本站导数"。一站一站接力即可。

(iii) sigmoid 的导数为什么是 。求导:

(用了 。)这就是代码里的 dsig = sig*(1-sig)

(iv) 从哪冒出来。 回忆 3.2(b):算 时我们故意省略了 求导该带的公因子 ,约定"最后统一补"。 的来源( 的前缀和),所以传到这里的 dg_nat 也是"少乘了一个 "的版本。现在就把它补回来:真正的梯度 = dg_nat 。因为整条链对 dg_nat 是线性的,补在哪一步都行,代码补在算 这一步。

(v) 一个量被用在多处怎么办( 的情形)。 同时乘进了 ——它只出现一次在 里,简单。但反过来看 对三个东西都有依赖:。求 对它们各自的偏导就是把另外的当常数:

每个偏导再乘上游 ,就是各自的梯度。下面 (a)(b)© 正是把 (ii) 的接力走完。

(a) 过

(代码 line 247 把 写在最后,等价。)

(b) 过

© 过

代码 line 252-253 先累加 ddt_bias_acc += dz*adA_log_acc += dz*(g_raw+dt),循环后:

为什么用 atomicAdd(line 259-260): 是每 head 一个标量、 是每 head 每维,同一 head 的很多 chunk-CTA 会并发写同一地址,必须原子加。而 是每 token 独立位置,直接写。

尾块掩码row >= actual_len 的行 dg_out=0(line 255),且不累加参数梯度。


6. Phase 5:dbeta 从 chunk 空间转回 logit 空间(line 263-272)

K2_bwd 给的 是对"已过 sigmoid 的 "的梯度。前向 ,反传过 sigmoid 回 logit:

仅前 actual_len 个线程各写一个 token 行。


7. 全部梯度公式速查表


8. 设计要点总结

主题 做法
数值稳定 dgc 把量级悬殊()的 改写成 ,避免 fp32 抵消
16 线程/行 + shfl 严格对齐前向 K1 的归约顺序,保证 bit 级可复现
dkn vs dkn_signed 前者全加(对 ,给 L2 反向);后者带符号(对
门控逆向 cumsum 是前缀和 → 反向是后缀和(rev_sum),再叠加常数
ln2 延后 求导的 统一在 Phase 4 乘入
atomicAdd 跨 chunk-CTA 累加;其余每 token 独立直接写
尾块掩码 actual_len 控制越界行读 0 / 写 0 / 不累加参数

配套阅读:上游 K2_bwd 的逐行精讲见 k2_bwd_recurrence_deep_dive.md。本 kernel 的输入 全部来自那里。

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