前置知识: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" template <int CHUNK, int D, int NumThreads, bool IsVarlen = true >__global__ void __launch_bounds__(NumThreads) _flash_kda_bwd_prepare( cutlass::bfloat16_t const * __restrict__ q_ptr, cutlass::bfloat16_t const * __restrict__ k_ptr, cutlass::bfloat16_t const * __restrict__ g_bf16_ptr, cutlass::bfloat16_t const * __restrict__ beta_ptr, float const * __restrict__ A_log_ptr, float const * __restrict__ dt_bias_ptr, float const * __restrict__ ws_kd_ptr, float const * __restrict__ ws_qd_ptr, float const * __restrict__ ws_kr_ptr, float const * __restrict__ ws_ki_ptr, float const * __restrict__ ws_gt_ptr, float const * __restrict__ ws_gc_ptr, 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, cutlass::bfloat16_t const * __restrict__ bwd_ws_dv_ptr, float const * __restrict__ bwd_ws_dbeta_ptr, cutlass::bfloat16_t * __restrict__ dq_ptr, cutlass::bfloat16_t * __restrict__ dk_ptr, cutlass::bfloat16_t * __restrict__ dv_ptr, cutlass::bfloat16_t * __restrict__ dg_ptr, cutlass::bfloat16_t * __restrict__ dbeta_ptr, float * __restrict__ dA_log_ptr, float * __restrict__ ddt_bias_ptr, 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; int global_tile_idx = blockIdx.x; int head_idx = blockIdx.y; int tid = threadIdx.x; int my_row = tid / THREADS_PER_ROW; int my_col = (tid % THREADS_PER_ROW) * ELEMS_PER_THREAD; 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]); 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; { 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; } #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 ); 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]; 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; 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(); 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]; } 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); } 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 ;
Grid = (total_tiles, H) :blockIdx.x = global_tile_idx,blockIdx.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*a、dA_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 的输入 全部来自那里。