KDA源码剖析之FlashKDA(上)

鱿鱼圈 Lv4

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

仓库链接:MoonshotAI/FlashKDA: FlashKDA: high-performance Kimi Delta Attention kernels

公式回顾

kernel 1 代码

代码链接:FlashKDA/csrc/smxx/fwd_kernel1.cuh at master · MoonshotAI/FlashKDA

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
 // ===== Launch Kernel 1 (prepare) =====
#if BLOCK_LEVEL_K1 >= 0
{
constexpr int kK1Threads = 256;
using SharedStorageK1T = SharedStorageK1<K1L>;
int smem_size_k1 = sizeof(SharedStorageK1T);

auto kernel1 = _flash_kda_fwd_prepare<
decltype(tma_load_q), decltype(tma_load_k),
decltype(tma_load_beta),
decltype(tma_load_g), decltype(tma_load_dt_bias),
decltype(tma_store_ws_kd), decltype(tma_store_ws_qd), decltype(tma_store_ws_kr),
decltype(tma_store_ws_gt), decltype(tma_store_ws_inv), decltype(tma_store_ws_mqk),
CHUNK, D, kK1Threads, IsVarlen
>;

cudaFuncSetAttribute(kernel1, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size_k1);

dim3 grid_k1(total_tiles, H);
dim3 block_k1(kK1Threads);

kernel1<<<grid_k1, block_k1, smem_size_k1, stream>>>(
tma_load_q, tma_load_k, tma_load_beta,
tma_load_g, tma_load_dt_bias,
tma_store_ws_kd, tma_store_ws_qd, tma_store_ws_kr,
tma_store_ws_gt, tma_store_ws_inv, tma_store_ws_mqk,
scale, T_total, H, N, cu_seqlens_ptr, total_tiles,
A_log_ptr, gate_scale
);
}
#endif


#pragma once

#include "utils.cuh"

template <int D, int CHUNK = 16>
struct K1Layouts {
using QKLayout = decltype(make_layout(make_shape(Int<CHUNK>{}, Int<D>{}), LayoutRight{}));
using GLayout = decltype(make_layout(make_shape(Int<CHUNK>{}, Int<D>{}), LayoutRight{}));
using MMALayout = decltype(tile_to_shape(
GMMA::Layout_K_INTER_Atom<cute::bfloat16_t>{},
make_shape(Int<CHUNK>{}, Int<D>{}),
LayoutLeft{}
));
using BetaSmemLayout = Layout<Shape<Int<32>>, Stride<Int<1>>>;
using GTotalLayout = Layout<Shape<Int<D>>, Stride<Int<1>>>;
using LMLayout = decltype(tile_to_shape(
GMMA::Layout_K_INTER_Atom<cute::bfloat16_t>{},
make_shape(Int<CHUNK>{}, Int<CHUNK>{}),
LayoutLeft{}
));
using TransposedLMLayout = decltype(tile_to_shape(
GMMA::Layout_MN_INTER_Atom<cute::bfloat16_t>{},
make_shape(Int<CHUNK>{}, Int<CHUNK>{}),
LayoutRight{}
));

using TMABetaSmemLayout = BetaSmemLayout; // 1D TMA, no dummy dim
using TMAQKLayout = decltype(prepend(QKLayout{}));
using TMAVOLayout = decltype(composition(
MMALayout{}.layout_a(),
MMALayout{}.offset(),
prepend(MMALayout{}.layout_b())
));
using TMAGLayout = decltype(prepend(GLayout{}));
using TMALMLayout = decltype(composition(
LMLayout{}.layout_a(),
LMLayout{}.offset(),
prepend(LMLayout{}.layout_b())
));
using TMAGTotalSmemLayout = decltype(prepend(GTotalLayout{}));
};

template <class Layouts>
struct SharedStorageK1 {
using BF16 = cutlass::bfloat16_t;
using QKLayout = typename Layouts::QKLayout;
using GLayout = typename Layouts::GLayout;
using BetaSmemLayout = typename Layouts::BetaSmemLayout;
using GTotalLayout = typename Layouts::GTotalLayout;
using LMLayout = typename Layouts::LMLayout;
using MMALayout = typename Layouts::MMALayout;

// Phase A: q, k, g alive
// Phase B: k_decayed, q_decayed, k_inv, L, INV, Mqk alive
// These don't overlap → union saves ~14KB shared memory
union {
struct {
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<QKLayout>> q;
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<QKLayout>> k;
alignas(128) cute::ArrayEngine<float, cute::cosize_v<GLayout>> g;
};
struct {
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<MMALayout>> k_decayed;
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<MMALayout>> q_decayed;
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<MMALayout>> k_inv;
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<LMLayout>> L;
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<LMLayout>> INV;
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<LMLayout>> Mqk;
};
};

alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<BetaSmemLayout>> beta;

union {
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<QKLayout>> g_bf16; // TMA load target
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<MMALayout>> k_restored;
};
union {
alignas(128) cute::ArrayEngine<float, cute::cosize_v<GTotalLayout>> dt_bias; // TMA load target
alignas(128) cute::ArrayEngine<float, cute::cosize_v<GTotalLayout>> g_total;
};
alignas(16) cutlass::arch::ClusterTransactionBarrier tma_load_barrier;
};

// ==================== Kernel 1: Prepare ====================
template <
class TmaLoadQ,
class TmaLoadK,
class TmaLoadBeta,
class TmaLoadG,
class TmaLoadDtBias,
class TmaStoreWsKD, class TmaStoreWsQD, class TmaStoreWsKR,
class TmaStoreWsGT, class TmaStoreWsINV, class TmaStoreWsMqk,
int CHUNK,
int D,
int NumThreads,
bool IsVarlen = true
>
__global__ void __launch_bounds__(NumThreads, 8) _flash_kda_fwd_prepare(
CUTE_GRID_CONSTANT TmaLoadQ const tma_load_q,
CUTE_GRID_CONSTANT TmaLoadK const tma_load_k,
CUTE_GRID_CONSTANT TmaLoadBeta const tma_load_beta,
CUTE_GRID_CONSTANT TmaLoadG const tma_load_g,
CUTE_GRID_CONSTANT TmaLoadDtBias const tma_load_dt_bias,
CUTE_GRID_CONSTANT TmaStoreWsKD const tma_store_ws_kd,
CUTE_GRID_CONSTANT TmaStoreWsQD const tma_store_ws_qd,
CUTE_GRID_CONSTANT TmaStoreWsKR const tma_store_ws_kr,
CUTE_GRID_CONSTANT TmaStoreWsGT const tma_store_ws_gt,
CUTE_GRID_CONSTANT TmaStoreWsINV const tma_store_ws_inv,
CUTE_GRID_CONSTANT TmaStoreWsMqk const tma_store_ws_mqk,
float scale,
int T_total,
int H,
int N,
int64_t const* cu_seqlens,
int total_tiles,
float const* A_log_ptr,
float gate_scale
) {
// --- constants
using BF16 = cutlass::bfloat16_t;
using FP16 = cutlass::half_t;
using Layouts = K1Layouts<D, CHUNK>;
using MMALayout = typename Layouts::MMALayout;
using QKLayout = typename Layouts::QKLayout;
using GLayout = typename Layouts::GLayout;
using BetaSmemLayout = typename Layouts::BetaSmemLayout;
using GTotalLayout = typename Layouts::GTotalLayout;
using LMLayout = typename Layouts::LMLayout;
using TransposedLMLayout = typename Layouts::TransposedLMLayout;
using TMAQKLayout = typename Layouts::TMAQKLayout;
using TMABetaSmemLayout = typename Layouts::TMABetaSmemLayout;
using TMAVOLayout = typename Layouts::TMAVOLayout;
using TMALMLayout = typename Layouts::TMALMLayout;
using TMAGTotalSmemLayout = typename Layouts::TMAGTotalSmemLayout;
constexpr uint32_t kTmaTransactionBytes =
uint32_t(cute::cosize_v<QKLayout>) * uint32_t(3 * sizeof(BF16)) + // q + k + g_bf16
uint32_t(32) * uint32_t(sizeof(BF16)) + // beta (bf16, sigmoid fused)
uint32_t(D) * uint32_t(sizeof(float)); // dt_bias

// --- shared memory
extern __shared__ __align__(128) unsigned char shared_mem[];
using SharedStorageT = SharedStorageK1<Layouts>;
SharedStorageT& shared_storage = *reinterpret_cast<SharedStorageT*>(shared_mem);

// --- per-CTA tile info
int global_tile_idx = blockIdx.x;
int head_idx = blockIdx.y;
int seq_idx, tiles_before, local_t;
int64_t bos, eos;
int seq_len, t_tiles_this_seq;

if constexpr (IsVarlen) {
// Linear scan on cu_seqlens to find (seq_idx, local_t)
seq_idx = 0;
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;
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;
// Early exit for excess CTAs (total_tiles is an upper bound)
if (local_t >= t_tiles_this_seq) return;
// --- TMA load inputs (single-shot, no pipeline)
// Only thread 0 issues TMA loads (not elect_one_sync which is per-warp)
if (threadIdx.x == 0) {
using BarrierType = cutlass::arch::ClusterTransactionBarrier::ValueType;
shared_storage.tma_load_barrier.init(1);
shared_storage.tma_load_barrier.arrive_and_expect_tx(kTmaTransactionBytes);

Tensor g_q = tma_load_q.get_tma_tensor(make_shape(H, T_total, D));
Tensor g_k = tma_load_k.get_tma_tensor(make_shape(H, T_total, D));
Tensor g_beta = tma_load_beta.get_tma_tensor(make_shape(H * T_total));

auto cta_tma_load_q = tma_load_q.get_slice(Int<0>{});
auto cta_tma_load_k = tma_load_k.get_slice(Int<0>{});
auto cta_tma_load_beta = tma_load_beta.get_slice(Int<0>{});

auto qk_off = g_q.layout()(head_idx, int(bos) + local_t * CHUNK, 0);
auto tile_shape_3d = make_shape(Int<1>{}, Int<CHUNK>{}, Int<D>{});
auto tile_stride_3d = stride(g_q.layout());
Tensor g_q_tile = make_tensor(g_q.data() + qk_off, make_layout(tile_shape_3d, tile_stride_3d));
Tensor g_k_tile = make_tensor(g_k.data() + qk_off, make_layout(tile_shape_3d, tile_stride_3d));

int beta_linear = head_idx * T_total + (int(bos) + local_t * CHUNK);
int beta_aligned = beta_linear & ~7;
auto beta_off = g_beta.layout()(beta_aligned);
Tensor g_beta_tile = make_tensor(g_beta.data() + beta_off, BetaSmemLayout{});

Tensor s_q_tile = make_tensor(make_smem_ptr(shared_storage.q.begin()), TMAQKLayout{});
Tensor s_k_tile = make_tensor(make_smem_ptr(shared_storage.k.begin()), TMAQKLayout{});
Tensor s_beta_tile = make_tensor(make_smem_ptr(shared_storage.beta.begin()), TMABetaSmemLayout{});

cute::copy(tma_load_q.with(reinterpret_cast<BarrierType&>(shared_storage.tma_load_barrier)),
cta_tma_load_q.partition_S(g_q_tile), cta_tma_load_q.partition_D(s_q_tile));
cute::copy(tma_load_k.with(reinterpret_cast<BarrierType&>(shared_storage.tma_load_barrier)),
cta_tma_load_k.partition_S(g_k_tile), cta_tma_load_k.partition_D(s_k_tile));
cute::copy(tma_load_beta.with(reinterpret_cast<BarrierType&>(shared_storage.tma_load_barrier)),
cta_tma_load_beta.partition_S(g_beta_tile), cta_tma_load_beta.partition_D(s_beta_tile));

// TMA load g_bf16 (same gmem layout as q/k)
Tensor g_g = tma_load_g.get_tma_tensor(make_shape(H, T_total, D));
auto cta_tma_load_g = tma_load_g.get_slice(Int<0>{});
Tensor g_g_tile = make_tensor(g_g.data() + qk_off, make_layout(tile_shape_3d, tile_stride_3d));
Tensor s_g_bf16_tile = make_tensor(make_smem_ptr(shared_storage.g_bf16.begin()), TMAQKLayout{});
cute::copy(tma_load_g.with(reinterpret_cast<BarrierType&>(shared_storage.tma_load_barrier)),
cta_tma_load_g.partition_S(g_g_tile), cta_tma_load_g.partition_D(s_g_bf16_tile));

// TMA load dt_bias [H, D] → [D] slice for current head
Tensor g_dt = tma_load_dt_bias.get_tma_tensor(make_shape(H, D));
auto cta_tma_load_dt = tma_load_dt_bias.get_slice(Int<0>{});
auto dt_off = g_dt.layout()(head_idx, 0);
Tensor g_dt_tile = make_tensor(g_dt.data() + dt_off,
make_layout(make_shape(Int<1>{}, Int<D>{}), stride(g_dt.layout())));
Tensor s_dt_tile = make_tensor(make_smem_ptr(shared_storage.dt_bias.begin()), TMAGTotalSmemLayout{});
cute::copy(tma_load_dt_bias.with(reinterpret_cast<BarrierType&>(shared_storage.tma_load_barrier)),
cta_tma_load_dt.partition_S(g_dt_tile), cta_tma_load_dt.partition_D(s_dt_tile));
}

// --- Compute a_log_exp (overlaps with TMA)
float a_log_exp = expf(A_log_ptr[head_idx]);
// --- Wait for TMA (q, k, beta, g_bf16, dt_bias)
__syncthreads();
shared_storage.tma_load_barrier.wait(0);
cutlass::arch::fence_view_async_shared();
__syncthreads();

// --- QK L2 Normalization ---
int compute_tid = threadIdx.x;
{
constexpr int ELEMS_PER_THREAD = 8;
constexpr int THREADS_PER_ROW = D / ELEMS_PER_THREAD; // 16
int my_row = threadIdx.x / THREADS_PER_ROW;
int my_col = (threadIdx.x % THREADS_PER_ROW) * ELEMS_PER_THREAD;

BF16* q_smem = shared_storage.q.begin();
BF16* k_smem = shared_storage.k.begin();

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 = bf16_to_f32(q_smem[my_row * D + my_col + i]);
float kv = bf16_to_f32(k_smem[my_row * D + my_col + i]);
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 = rsqrtf(q_sq + 1e-6f);
float k_inv = rsqrtf(k_sq + 1e-6f);

#pragma unroll
for (int i = 0; i < ELEMS_PER_THREAD; ++i) {
q_smem[my_row * D + my_col + i] = BF16(q_vals[i] * q_inv);
k_smem[my_row * D + my_col + i] = BF16(k_vals[i] * k_inv);
}
}
__syncthreads();

// --- Fused gate activation + cumsum + k tail zero-fill ---
// Threads 0-127: gate(g_bf16 + dt_bias) → cumulative sum, eliminates raw-g smem round-trip
// Threads 128-255: zero k for tail rows
{
int actual_len = min(CHUNK, seq_len - local_t * CHUNK);
if (compute_tid < 128) {
int col = compute_tid;
BF16 const* g_bf16_smem = shared_storage.g_bf16.begin();
float dt = shared_storage.dt_bias.begin()[col];
float* g_smem = shared_storage.g.begin();
float sum = 0.0f;
#pragma unroll
for (int row = 0; row < CHUNK; ++row) {
float g_val;
if (row < actual_len) {
g_val = bf16_to_f32(g_bf16_smem[row * D + col]) + dt;
g_val = a_log_exp * g_val;
g_val = gate_scale * sigmoid_tanh_approx_f32(g_val);
} else {
g_val = 0.0f;
}
sum += g_val;
g_smem[row * D + col] = sum;
}
shared_storage.g_total.begin()[col] = sum;
} else {
int col = compute_tid - 128;
BF16* k_smem = shared_storage.k.begin();
for (int row = actual_len; row < CHUNK; ++row) {
k_smem[row * D + col] = BF16(0);
}
}
}
__syncthreads();

Tensor q_tile = make_tensor(make_smem_ptr(shared_storage.q.begin()), QKLayout{});
Tensor k_tile = make_tensor(make_smem_ptr(shared_storage.k.begin()), QKLayout{});
Tensor g_tile = make_tensor(make_smem_ptr(shared_storage.g.begin()), GLayout{});
Tensor beta_tile = make_tensor(make_smem_ptr(shared_storage.beta.begin()), BetaSmemLayout{});
int beta_smem_offset = (head_idx * T_total + int(bos) + local_t * CHUNK) & 7;

Tensor k_restored = make_tensor(make_smem_ptr(shared_storage.k_restored.begin()), MMALayout{});
Tensor k_decayed = make_tensor(make_smem_ptr(shared_storage.k_decayed.begin()), MMALayout{});
Tensor q_decayed = make_tensor(make_smem_ptr(shared_storage.q_decayed.begin()), MMALayout{});
Tensor k_inv = make_tensor(make_smem_ptr(shared_storage.k_inv.begin()), MMALayout{});
Tensor g_total = make_tensor(make_smem_ptr(shared_storage.g_total.begin()), GTotalLayout{});

// exp_g_total: compute exp(g_total) in smem before decay_apply
if (compute_tid < 128) {
float x = g_total(compute_tid);
g_total(compute_tid) = ex2_approx_ftz_f32(x);
}
__syncthreads();

// decay_apply
if (compute_tid < 256) {
static_assert(D % 64 == 0);
static_assert(CHUNK % 8 == 0);

int lane = compute_tid % 32;
int warp_id = compute_tid / 32;
int g = lane / 4;
int t = lane % 4;

auto vec8_2d = make_shape(_1{}, _8{});
auto vec8_1d = make_shape(_8{});
auto thr2_2d = make_shape(_1{}, _2{});
auto thr2_1d = make_shape(_2{});

constexpr int N_M = CHUNK / 8;
constexpr int N_N = D / 64;
constexpr int N_TILES = N_M * N_N;

float reg_g[N_TILES][2];
BF16 reg_q[N_TILES][2];
BF16 reg_k[N_TILES][2];
float reg_gt[N_TILES][2];

#pragma unroll
for (int m_blk = 0; m_blk < CHUNK; m_blk += 8) {
#pragma unroll
for (int n_blk = 0; n_blk < D; n_blk += 64) {
int tile_idx = (m_blk / 8) * N_N + (n_blk / 64);
int row = m_blk + ((warp_id + g) % 8);
int col_base = n_blk + g * 8;
int col_tile = col_base / 8;

Tensor tile_g = local_tile(g_tile, vec8_2d, make_coord(row, col_tile));
Tensor tile_q = local_tile(q_tile, vec8_2d, make_coord(row, col_tile));
Tensor tile_k = local_tile(k_tile, vec8_2d, make_coord(row, col_tile));
Tensor tile_gt = local_tile(g_total, vec8_1d, make_coord(col_tile));

Tensor s_g = local_tile(tile_g, thr2_2d, make_coord(0, t));
Tensor s_q = local_tile(tile_q, thr2_2d, make_coord(0, t));
Tensor s_k = local_tile(tile_k, thr2_2d, make_coord(0, t));
Tensor s_gt = local_tile(tile_gt, thr2_1d, make_coord(t));

Tensor r_g = make_tensor_like<float>(s_g);
Tensor r_q = make_tensor_like<BF16>(s_q);
Tensor r_k = make_tensor_like<BF16>(s_k);
Tensor r_gt = make_tensor_like<float>(s_gt);

cute::copy(AutoVectorizingCopy{}, s_g, r_g);
cute::copy(AutoVectorizingCopy{}, s_q, r_q);
cute::copy(AutoVectorizingCopy{}, s_k, r_k);
cute::copy(AutoVectorizingCopy{}, s_gt, r_gt);

#pragma unroll
for (int v = 0; v < 2; ++v) {
reg_g[tile_idx][v] = r_g(0, v);
reg_q[tile_idx][v] = r_q(0, v);
reg_k[tile_idx][v] = r_k(0, v);
reg_gt[tile_idx][v] = r_gt(v);
}
}
}

// Sync before writing to union'd smem (q/k/g → k_decayed/q_decayed/k_inv)
// Safe: all 256 threads enter this if block (compute_tid < 256 always true)
__syncthreads();

#pragma unroll
for (int m_blk = 0; m_blk < CHUNK; m_blk += 8) {
#pragma unroll
for (int n_blk = 0; n_blk < D; n_blk += 64) {
int tile_idx = (m_blk / 8) * N_N + (n_blk / 64);
int row = m_blk + ((warp_id + g) % 8);
int col_base = n_blk + g * 8;
int col_tile = col_base / 8;

Tensor tile_qd = local_tile(q_decayed, vec8_2d, make_coord(row, col_tile));
Tensor tile_kd = local_tile(k_decayed, vec8_2d, make_coord(row, col_tile));
Tensor tile_kr = local_tile(k_restored, vec8_2d, make_coord(row, col_tile));
Tensor tile_ki = local_tile(k_inv, vec8_2d, make_coord(row, col_tile));

Tensor s_qd = local_tile(tile_qd, thr2_2d, make_coord(0, t));
Tensor s_kd = local_tile(tile_kd, thr2_2d, make_coord(0, t));
Tensor s_kr = local_tile(tile_kr, thr2_2d, make_coord(0, t));
Tensor s_ki = local_tile(tile_ki, thr2_2d, make_coord(0, t));

Tensor r_qd = make_tensor_like<BF16>(s_qd);
Tensor r_kd = make_tensor_like<BF16>(s_kd);
#pragma unroll
for (int v = 0; v < 2; ++v) {
float g = reg_g[tile_idx][v];
BF16 q = reg_q[tile_idx][v];
BF16 k = reg_k[tile_idx][v];
BF16 exp_cumsum = BF16(ex2_approx_ftz_f32(g));
r_qd(0, v) = q * exp_cumsum * BF16(scale);
r_kd(0, v) = k * exp_cumsum;
}
cute::copy(AutoVectorizingCopy{}, r_qd, s_qd);
cute::copy(AutoVectorizingCopy{}, r_kd, s_kd);

Tensor r_ki = make_tensor_like<BF16>(s_ki);
Tensor r_kr = make_tensor_like<BF16>(s_kr);
#pragma unroll
for (int v = 0; v < 2; ++v) {
float g = reg_g[tile_idx][v];
BF16 k = reg_k[tile_idx][v];
BF16 inv_cumsum = BF16(ex2_approx_ftz_f32(-g));
r_ki(0, v) = k * inv_cumsum;
r_kr(0, v) = k * inv_cumsum * BF16(reg_gt[tile_idx][v]);
}
cute::copy(AutoVectorizingCopy{}, r_ki, s_ki);
cute::copy(AutoVectorizingCopy{}, r_kr, s_kr);
}
}

}
__syncthreads();

Tensor L = make_tensor(make_smem_ptr(shared_storage.L.begin()), LMLayout{});
Tensor Mqk = make_tensor(make_smem_ptr(shared_storage.Mqk.begin()), LMLayout{});
Tensor L_fp16 = make_tensor(make_smem_ptr(reinterpret_cast<FP16*>(shared_storage.L.begin())), LMLayout{});

// L_Mqk
if (compute_tid < 32) {
mma_m16n16_bf16bf16fp16_1warp(k_decayed, k_inv, L_fp16, compute_tid);
} else if (compute_tid >= 32 && compute_tid < 64) {
mma_m16n16_bf16bf16bf16_1warp(q_decayed, k_inv, Mqk, compute_tid - 32);
}
__syncthreads();

Tensor INV = make_tensor(make_smem_ptr(shared_storage.INV.begin()), LMLayout{});
Tensor INV_fp16 = make_tensor(make_smem_ptr(reinterpret_cast<FP16*>(shared_storage.INV.begin())), LMLayout{});

// tril_IL + INV = I - L (merged, same thread same element)
if (compute_tid < 256) {
const int col_block_size = 8;
int block_idx = compute_tid / (CHUNK * col_block_size);
int i = (compute_tid / col_block_size) % CHUNK;
int j = compute_tid % col_block_size + block_idx * col_block_size;
if (i <= j) {
L_fp16(i, j) = FP16::bitcast(0);
} else {
L_fp16(i, j) = L_fp16(i, j) * FP16(sigmoid_tanh_approx_f32(float(beta_tile(beta_smem_offset + i))));
}
if (i < j) {
Mqk(i, j) = BF16::bitcast(0);
}
// INV = I - L (same thread reads L(i,j) it just wrote)
FP16 x = L_fp16(i, j);
INV_fp16(i, j) = (i == j ? FP16(1.0f) - x : -x);
}
__syncthreads();

// inv (Neumann series, fused in registers)
neumann_inv_fused_1warp(L_fp16, INV_fp16, INV, compute_tid);
// Fence + sync combined: completion + TMA visibility
cutlass::arch::fence_view_async_shared();
__syncthreads();
if (threadIdx.x == 0) {
int ws_idx = head_idx * total_tiles + global_tile_idx;
// Store k_decayed [CHUNK, D] bf16
{
auto g_ws = tma_store_ws_kd.get_tma_tensor(make_shape(H * total_tiles, CHUNK, D));
auto ws_off = g_ws.layout()(ws_idx, 0, 0);
Tensor g_ws_tile = make_tensor(g_ws.data() + ws_off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<D>{}), stride(g_ws.layout())));
Tensor s_kd = make_tensor(make_smem_ptr(shared_storage.k_decayed.begin()), TMAVOLayout{});
auto cta_tma = tma_store_ws_kd.get_slice(Int<0>{});
cute::copy(tma_store_ws_kd, cta_tma.partition_S(s_kd), cta_tma.partition_D(g_ws_tile));
tma_store_arrive();
}
// Store q_decayed
{
auto g_ws = tma_store_ws_qd.get_tma_tensor(make_shape(H * total_tiles, CHUNK, D));
auto ws_off = g_ws.layout()(ws_idx, 0, 0);
Tensor g_ws_tile = make_tensor(g_ws.data() + ws_off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<D>{}), stride(g_ws.layout())));
Tensor s_qd = make_tensor(make_smem_ptr(shared_storage.q_decayed.begin()), TMAVOLayout{});
auto cta_tma = tma_store_ws_qd.get_slice(Int<0>{});
cute::copy(tma_store_ws_qd, cta_tma.partition_S(s_qd), cta_tma.partition_D(g_ws_tile));
tma_store_arrive();
}
// Store k_restored
{
auto g_ws = tma_store_ws_kr.get_tma_tensor(make_shape(H * total_tiles, CHUNK, D));
auto ws_off = g_ws.layout()(ws_idx, 0, 0);
Tensor g_ws_tile = make_tensor(g_ws.data() + ws_off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<D>{}), stride(g_ws.layout())));
Tensor s_kr = make_tensor(make_smem_ptr(shared_storage.k_restored.begin()), TMAVOLayout{});
auto cta_tma = tma_store_ws_kr.get_slice(Int<0>{});
cute::copy(tma_store_ws_kr, cta_tma.partition_S(s_kr), cta_tma.partition_D(g_ws_tile));
tma_store_arrive();
}
// Store g_total [D] float
{
auto g_ws = tma_store_ws_gt.get_tma_tensor(make_shape(H * total_tiles, D));
auto ws_off = g_ws.layout()(ws_idx, 0);
Tensor g_ws_tile = make_tensor(g_ws.data() + ws_off,
make_layout(make_shape(Int<1>{}, Int<D>{}), stride(g_ws.layout())));
Tensor s_gt = make_tensor(make_smem_ptr(shared_storage.g_total.begin()), TMAGTotalSmemLayout{});
auto cta_tma = tma_store_ws_gt.get_slice(Int<0>{});
cute::copy(tma_store_ws_gt, cta_tma.partition_S(s_gt), cta_tma.partition_D(g_ws_tile));
tma_store_arrive();
}
// Store INV [CHUNK, CHUNK] bf16
{
auto g_ws = tma_store_ws_inv.get_tma_tensor(make_shape(H * total_tiles, CHUNK, CHUNK));
auto ws_off = g_ws.layout()(ws_idx, 0, 0);
Tensor g_ws_tile = make_tensor(g_ws.data() + ws_off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<CHUNK>{}), stride(g_ws.layout())));
Tensor s_inv = make_tensor(make_smem_ptr(shared_storage.INV.begin()), TMALMLayout{});
auto cta_tma = tma_store_ws_inv.get_slice(Int<0>{});
cute::copy(tma_store_ws_inv, cta_tma.partition_S(s_inv), cta_tma.partition_D(g_ws_tile));
tma_store_arrive();
}
// Store Mqk [CHUNK, CHUNK] bf16
{
auto g_ws = tma_store_ws_mqk.get_tma_tensor(make_shape(H * total_tiles, CHUNK, CHUNK));
auto ws_off = g_ws.layout()(ws_idx, 0, 0);
Tensor g_ws_tile = make_tensor(g_ws.data() + ws_off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<CHUNK>{}), stride(g_ws.layout())));
Tensor s_mqk = make_tensor(make_smem_ptr(shared_storage.Mqk.begin()), TMALMLayout{});
auto cta_tma = tma_store_ws_mqk.get_slice(Int<0>{});
cute::copy(tma_store_ws_mqk, cta_tma.partition_S(s_mqk), cta_tma.partition_D(g_ws_tile));
tma_store_arrive();
}
}
tma_store_wait<0>();
__syncthreads();
}

计算流程图

image

image

image

image

整体剖析

grid: (total_tiles, H) threads: 256


第一部分: 数据布局定义 (第5-41行)

1
2
template <int D, int CHUNK = 16>
struct K1Layouts {

这是一个模板结构体,定义K1 kernel里所有shared memory的排列方式。 D=128 (head维度), CHUNK=16 (每个chunk的token数)。

QKLayout (第7行)

1
using QKLayout = make_layout(make_shape(16, 128), LayoutRight{});

形状: [16, 128] LayoutRight: 行主序(row-major)

  • 地址 = i * 128 + j
  • 总大小 = 16 × 128 = 2048 个元素
  • 用途: 存放 TMA 加载进来的原始 q 和 k

MMALayout (第9-13行)

1
2
3
4
5
using MMALayout = tile_to_shape(
GMMA::Layout_K_INTER_Atom<bf16>{},
make_shape(16, 128),
LayoutLeft{}
);

什么是 “swizzle 布局”?

  • GPU shared memory 分成 32 个 bank,每 bank 4字节宽
  • 如果一个 warp 的 32 个线程同时访问同一 bank → bank conflict → 串行化
  • Swizzle 布局通过"打乱"地址映射,让线程尽量访问不同 bank

GMMA::Layout_K_INTER_Atom**<bf16>**: CUTLASS 预定义的 swizzle 模式,专门为 bf16 MMA 指令优化,确保加载操作数时无 bank conflict。

tile_to_shape: 把这个原子 swizzle 模式"铺开"到 [16, 128] 的大小 LayoutLeft: 列优先的铺开方式

简单理解: MMALayout 和 QKLayout 存的是同样 16×128 的数据,但 MMALayout 的地址映射经过特殊打乱,让 MMA 指令读取时不冲突。

其他布局 (第14-20行)

布局 形状 用途
BetaSmemLayout [32], stride=1 beta 加载 (实际只用16个,TMA要求对齐)
GTotalLayout [128], stride=1 g_total 存128个float
LMLayout [16,16] swizzle L 和 Mqk 矩阵
TransposedLMLayout [16,16] 转置方向 转置版本的 LMLayout

TMA 布局 (第27-40行)

TMA (Tensor Memory Accelerator) 是 Hopper GPU 的硬件单元:

传统做法 TMA做法
线程发load指令 一个线程下一条命令
通过L1/L2 cache 硬件DMA直接搬数据
数据到寄存器 → 写到smem 直接搬整块到smem

好处:

  • 不占计算线程的时间,异步执行
  • 自动处理2D/3D地址计算(给行列坐标,自动算出全局内存地址)
  • 自动处理 swizzle (搬运时直接按 MMALayout 排列)

TMA 布局需要比对应的 smem 布局多一个维度(prepend加了一个size=1的维度),这是 CUTLASS TMA 接口的要求。


第二部分: Shared Memory 结构 (第43-83行)

union 核心 (第56-70行)

1
2
3
4
union {
struct { q, k, g }; // Phase A: 加载阶段
struct { k_decayed, q_decayed, k_inv, L, INV, Mqk }; // Phase B: 计算阶段
};

C++ 的 union 意味着两个 struct 共享同一块内存!

Phase A: 存原始 q, k, g

变量 大小 类型
q 16×128×2 = 4096 字节 bf16
k 16×128×2 = 4096 字节 bf16
g 16×128×4 = 8192 字节 float (cumsum需要fp32)
合计 16384 字节

Phase B: 存计算结果 (共用同一块内存)

变量 大小 类型
k_decayed 16×128×2 = 4096 字节 bf16 (MMALayout)
q_decayed 16×128×2 = 4096 字节 bf16 (MMALayout)
k_inv 16×128×2 = 4096 字节 bf16 (MMALayout)
L 16×16×2 = 512 字节 fp16
INV 16×16×2 = 512 字节 bf16
Mqk 16×16×2 = 512 字节 bf16
合计 13824 字节

两个阶段不重叠(Phase A用完才开始Phase B),所以可以共用。省了约 14KB shared memory,让 SM 能放更多 CTA (occupancy 更高)。

其他内存区域

1
2
3
4
alignas(128) bf16 beta[32];              // 不在union里,两个阶段都要用
bf16 k_restored[16][128]; // 与g_bf16共用空间
float g_total[128]; // 与dt_bias共用空间
ClusterTransactionBarrier tma_load_barrier; // Hopper硬件同步原语

ClusterTransactionBarrier:

  • TMA 是异步的: thread 0 发命令后立即返回,数据还没到

  • barrier 用来等 TMA 完成:

  • 发 TMA 前: barrier.arrive_and_expect_tx(字节数) — 告诉barrier"将有这么多数据到达"

  • 等待: barrier.wait(0) — 阻塞直到所有预期的数据到达


第三部分: kernel 函数签名 (第85-119行)

1
__global__ void __launch_bounds__(NumThreads, 8) _flash_kda_fwd_prepare(...)
关键字 含义
__global__ CUDA kernel 函数
__launch_bounds__(256, 8) 每个CTA最多256线程,目标每SM同时运行8个CTA → 编译器限制寄存器使用,让SM能塞下更多CTA

参数列表 (第100-118行)

1
2
CUTE_GRID_CONSTANT TmaLoadQ const tma_load_q
// ... 6个load descriptor和6个store descriptor

CUTE_GRID_CONSTANT 告诉编译器这些参数放在 GPU 的 constant memory 里。

TMA descriptor 包含: 全局内存地址、形状、stride、swizzle模式。thread 用 descriptor 来发起 TMA 操作,硬件自动完成数据搬运。

类型 数量 用途
Load descriptor 6 q, k, beta, g, dt_bias
Store descriptor 6 kd, qd, kr, gt, inv, mqk 到workspace

标量参数:

  • scale: 注意力缩放
  • T_total: 总token数
  • H: head数
  • N: 序列数
  • total_tiles: 总chunk数
  • gate_scale: gate下界×log2e

第四部分: 每个CTA确定自己处理哪个chunk (第147-181行)

1
2
3
int global_tile_idx = blockIdx.x;
int head_idx = blockIdx.y;
grid = (total_tiles, H)
  • blockIdx.x: 全局chunk编号 (0, 1, …, total_tiles-1)
  • blockIdx.y: head编号 (0, 1, …, H-1)

找到当前chunk属于哪个序列

Varlen模式 (多个变长序列拼在一起):

1
2
3
4
5
6
7
8
9
10
11
12
线性扫描 cu_seqlens 数组:
for i in 0..N:
slen = cu_seqlens[i+1] - cu_seqlens[i] // 第i个序列长度
n_tiles = ceil(slen / 16) // 第i个序列有几个chunk
if tiles_before + n_tiles > global_tile_idx:
seq_idx = i // 找到了!
break
tiles_before += n_tiles

local_t = global_tile_idx - tiles_before // chunk在序列内的编号
bos = cu_seqlens[seq_idx] // 序列起始token位置
eos = cu_seqlens[seq_idx + 1] // 序列结束token位置

非Varlen模式 (所有序列等长):

1
2
3
4
5
6
T_seq = T_total / N
tiles_per_seq = ceil(T_seq / 16)
seq_idx = global_tile_idx / tiles_per_seq
local_t = global_tile_idx % tiles_per_seq
bos = seq_idx * T_seq
eos = bos + T_seq

边界检查

1
if (local_t >= t_tiles_this_seq) return;

因为 total_tiles 是上界(varlen时可能多分配了一些),多余的CTA直接退出。


第五部分: TMA加载输入数据 (第182-244行)

thread 0 发起所有TMA

1
if (threadIdx.x == 0) {

只有 thread 0 发 TMA 命令。TMA 是硬件操作,一个线程发命令就够了,其他 255 个线程不需要参与。

初始化 barrier

1
2
shared_storage.tma_load_barrier.init(1);
shared_storage.tma_load_barrier.arrive_and_expect_tx(kTmaTransactionBytes);

告诉 barrier “接下来会有 kTmaTransactionBytes 字节的数据到来”。

1
2
3
kTmaTransactionBytes = 3×16×128×2 + 32×2 + 128×4
= 3×4096 + 64 + 512
= 12864 字节

这是 q + k + g_bf16 + beta + dt_bias 的总大小。

准备TMA地址

1
2
3
4
5
6
7
8
9
Tensor g_q = tma_load_q.get_tma_tensor(make_shape(H, T_total, D));
// get_tma_tensor: 用TMA descriptor构造"虚拟全局tensor"
// 不真正分配内存,只是用来计算坐标到地址的映射

qk_off = g_q.layout()(head_idx, bos + local_t * 16, 0)
// 算出当前(head, chunk起始token, dim=0)在全局内存中的偏移

g_q_tile = make_tensor(g_q.data() + qk_off, ...)
// 构造指向具体数据的 tile (大小 [1, 16, 128])

发起TMA

1
2
3
4
s_q_tile = make_tensor(smem_ptr(shared_storage.q), TMAQKLayout{})
// 构造 shared memory 侧的 tile

cute::copy(tma_load_q.with(barrier), partition_S(g_q_tile), partition_D(s_q_tile))

这一行就是 TMA 操作:

  • 源(S): 全局内存的 g_q_tile
  • 目的(D): shared memory 的 s_q_tile
  • with(barrier): 绑定到 barrier,完成时自动通知

TMA 硬件会:

  1. 根据 descriptor 算出全局内存地址
  2. 发起 DMA 请求
  3. 数据到达 shared memory 后,自动给 barrier 的计数器减(已传输字节数)
  4. 当 barrier 的计数器减到 0,barrier.wait() 解除阻塞

q, k, g_bf16, beta, dt_bias 共 5 次 TMA,都绑定同一个 barrier。

CPU计算与TMA重叠

1
2
float a_log_exp = expf(A_log_ptr[head_idx]);
// A_log 是 per-head 的 gate 参数,在等 TMA 时计算(overlap)

同步模式 (Hopper 标准)

1
2
3
4
__syncthreads();                                    // 确保barrier初始化可见
shared_storage.tma_load_barrier.wait(0); // 等待TMA完成
fence_view_async_shared(); // memory fence
__syncthreads(); // 确保所有线程看到数据

第六部分: L2归一化 (第246-285行)

线程分工

1
2
constexpr int ELEMS_PER_THREAD = 8;
constexpr int THREADS_PER_ROW = D / 8 = 128 / 8 = 16;

每行 128 维,16 个线程分工,每线程处理 8 个元素。 256 个线程 / 16 = 16 行,正好覆盖 CHUNK=16 行。

1
2
int my_row = threadIdx.x / 16;           // 我负责第几行
int my_col = (threadIdx.x % 16) * 8; // 我负责的起始列

示例: threadIdx.x = 35

  • my_row = 35 / 16 = 2 (第2行)
  • my_col = (35 % 16) × 8 = 3 × 8 = 24 (第24-31列)

读数据 + 计算平方和

1
2
3
4
5
6
7
for i in 0..7:
qv = bf16_to_f32(q_smem[my_row * 128 + my_col + i])
kv = bf16_to_f32(k_smem[my_row * 128 + my_col + i])
q_vals[i] = qv
k_vals[i] = kv
q_sq += qv * qv
k_sq += kv * kv

每个线程累加自己 8 个元素的平方和。

warp shuffle 规约

1
2
3
for delta = 8, 4, 2, 1:
q_sq += __shfl_xor_sync(0xFFFFFFFF, q_sq, delta)
k_sq += __shfl_xor_sync(0xFFFFFFFF, k_sq, delta)

__shfl_xor_sync: warp 内的寄存器数据交换。 lane i 和 lane i^delta 互换 q_sq 值,然后各自相加。

delta 效果
8 lane 0-7 和 lane 8-15 交换 → 每个lane有16个元素的平方和
4 相邻4个交换
2 相邻2个交换
1 最终同一行的16个线程都有完整的行平方和

为什么从 delta=8 开始而不是 16? 因为 THREADS_PER_ROW=16,同一行占了半个 warp (lane 0-15)。 delta=8 让 lane 0 和 lane 8 交换(同一行内),够了。

归一化

1
2
3
4
5
q_inv = rsqrtf(q_sq + 1e-6f)    // 1/sqrt(||q||^2 + eps)
k_inv = rsqrtf(k_sq + 1e-6f)

q_smem[...] = BF16(q_vals[i] * q_inv)
k_smem[...] = BF16(k_vals[i] * k_inv)

结果直接写回 shared memory (覆盖原来的 q/k)。 这就是 L2 normalization: x = x / ||x||


第七部分: gate激活 + cumsum (第287-319行)

获取实际长度

1
int actual_len = min(CHUNK, seq_len - local_t * CHUNK);

当前chunk的实际token数(最后一个chunk可能不满16)。

threads 0-127 做 gate + cumsum (每线程一列)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
int col = compute_tid;              // 我负责第几列 (0-127)
float dt = dt_bias[col]; // 这列的 gate bias
float sum = 0.0f;

for row in 0..15:
if row < actual_len:
g_val = bf16_to_f32(g_bf16[row * 128 + col]) + dt
g_val = a_log_exp * g_val
g_val = gate_scale * sigmoid(g_val)
else:
g_val = 0.0f // padding行的gate=0

sum += g_val // 前缀和
g_smem[row * 128 + col] = sum // 存cumsum结果

g_total[col] = sum // chunk总gate

gate激活公式:

1
g_val = gate_scale × sigmoid(A_log_exp × (g_bf16 + dt_bias))

其中:

  • gate_scale = lower_bound × log2(e) ≈ -5 × 1.4427 = -7.2135
  • sigmoidtanh.approx.f32 实现: sigmoid(x) = 0.5 × tanh(0.5×x) + 0.5
  • 最终 g_val 的范围是 [gate_scale, 0] = [-7.2135, 0]
  • cumsum 16步后最大绝对值 ≈ 16 × 7.2135 ≈ 115.4

对比fla: fla用单独的 chunk_local_cumsum_vector_kernel 来做 cumsum,并且需要先在外部做 gate 激活。FlashKDA 全部融合在一起。

threads 128-255 清零 k 的 padding 行

1
2
for row = actual_len .. 15:
k_smem[row * 128 + col] = 0

如果 chunk 不满 16 个 token,把 k 的多余行清零。这样后面矩阵乘时自然不会把 padding 算进去。


第八部分: g_total 转为 exp2 形式 (第334-339行)

1
2
3
if (compute_tid < 128):
x = g_total(compute_tid)
g_total(compute_tid) = ex2_approx_ftz_f32(x)

g_total 从 cumsum 值变成 exp2(cumsum) 值。后面 decay_apply 需要用 exp2(g_total) 来计算 k_restored。

1
ex2_approx_ftz_f32`: PTX 指令 `ex2.approx.ftz.f32
  • ex2 = 2^x (base-2 指数)
  • approx = 近似(快速, 约23bit精度)
  • ftz = flush to zero (极小值直接归零)
  • expf(x) 快,因为 expf 需要先乘 log2(e) 再调 ex2

第九部分: decay_apply (第341-457行)

最复杂的一段。256个线程把 q, k, g 转换成 q_decayed, k_decayed, k_inv, k_restored。

warp/lane 信息

1
2
3
4
int lane = compute_tid % 32;    // warp内的lane编号 (0-31)
int warp_id = compute_tid / 32; // warp编号 (0-7)
int g = lane / 4; // 每4个lane一组,组号 (0-7)
int t = lane % 4; // 组内编号 (0-3)

命名 gt 是因为数据访问模式:

  • g 决定读哪一行: row = m_blk + (warp_id + g) % 8
  • t 决定读行内的哪 2 个元素: 第 t×2t×2+1

循环参数

1
2
3
N_M = CHUNK / 8 = 2     // 16行分成2个8行块
N_N = D / 64 = 2 // 128列分成2个64列块
N_TILES = 2 × 2 = 4 // 总共4个tile

寄存器数组

1
2
3
4
float reg_g[4][2]   -- 4个tile,每个tile存2个gate值
BF16 reg_q[4][2] -- 4个tile的q值
BF16 reg_k[4][2] -- 4个tile的k值
float reg_gt[4][2] -- 4个tile的g_total值

第一轮循环: 读数据到寄存器

1
2
3
4
5
6
7
8
9
10
11
for m_blk in {0, 8}:         // 两个行块
for n_blk in {0, 64}: // 两个列块
tile_idx = (m_blk/8)*2 + (n_blk/64) // 0,1,2,3
row = m_blk + (warp_id + g) % 8 // 每个线程读不同的行

// 每个线程读q/k/g的2个元素
for v in 0,1:
reg_g[tile_idx][v] = g_tile(row, col+t*2+v)
reg_q[tile_idx][v] = q_tile(row, col+t*2+v)
reg_k[tile_idx][v] = k_tile(row, col+t*2+v)
reg_gt[tile_idx][v] = g_total(col+t*2+v)

关键: (warp_id + g) % 8 让不同 warp 的不同 lane-group 读不同的行。这样 256 个线程读完了全部 16×128 = 2048 个 (q,k,g) 值。

每个线程在4个tile中各持有2个元素,共 4×2 = 8 个元素。

同步

1
__syncthreads();

同步! 接下来要写到 union 的另一半 (k_decayed/q_decayed/k_inv/k_restored)。必须确保所有线程都读完了 q/k/g,因为写入会覆盖它们。

第二轮循环: 计算并写结果

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
for m_blk, n_blk (同上):
for v in 0,1:
g = reg_g[tile_idx][v] // cumsum gate值
q = reg_q[tile_idx][v] // 归一化后的q
k = reg_k[tile_idx][v] // 归一化后的k

exp_cumsum = ex2_approx_ftz_f32(g) // 2^g = exp2(cumsum_gate)
r_qd = q * exp_cumsum * scale // q_decayed = q * exp2(g) * scale
r_kd = k * exp_cumsum // k_decayed = k * exp2(g)

inv_cumsum = ex2_approx_ftz_f32(-g) // 2^(-g)
r_ki = k * inv_cumsum // k_inv = k * exp2(-g)
r_kr = k * inv_cumsum * reg_gt[tile_idx][v] // k_restored = k * exp2(-g) * exp2(g_total)
// = k * exp2(g_total - g)

// 写到 MMALayout 的 shared memory
copy(r_qd, s_qd)
copy(r_kd, s_kd)
copy(r_ki, s_ki)
copy(r_kr, s_kr)

注意这里写的是 MMALayout (swizzle布局),不是 QKLayout (行主序)。因为后面 MMA 操作需要 swizzle 布局。

对应fla的关系

FlashKDA fla对应 公式
k_decayed k × beta × exp2(gk) = k * exp2(g)
q_decayed qg × scale = q * exp2(g) * scale
k_inv 用于构造L矩阵 = k * exp2(-g)
k_restored kg = k * exp2(g_total - gk)

第十部分: 构造 L 和 Mqk 矩阵 (第459-468行)

1
2
3
Tensor L   = make_tensor(smem_ptr(shared_storage.L), LMLayout{});
Tensor Mqk = make_tensor(smem_ptr(shared_storage.Mqk), LMLayout{});
Tensor L_fp16 = make_tensor(smem_ptr(reinterpret_cast<FP16*>(L.data())), LMLayout{});

L 用 fp16 存(和bf16共用同样的 LMLayout 因为都是2字节)。reinterpret_cast 只是改变类型解释,不改变内存。

MMA计算

1
2
3
4
5
if (compute_tid < 32) {
mma_m16n16_bf16bf16fp16_1warp(k_decayed, k_inv, L_fp16, compute_tid);
} else if (compute_tid >= 32 && compute_tid < 64) {
mma_m16n16_bf16bf16bf16_1warp(q_decayed, k_inv, Mqk, compute_tid - 32);
}
Warp 线程 计算 结果类型
warp 0 0-31 L = k_decayed @ k_inv^T fp16
warp 1 32-63 Mqk = q_decayed @ k_inv^T bf16

其他 warp 闲着。

mma_m16n16_xxx 内部 (utils.cuh 第146-165行)

  • SM80_16x8x16_F32BF16BF16F32_TN 的 MMA atom
  • 输入 bf16,累加 fp32,然后转目标精度存
  • 一个 16×16 矩阵乘 = 两个 16×8×16 atom (沿N维拼)
  • cooperative_gemm 让一个warp协作完成

数学: L_ij = sum_d k_decayed[i,d] × k_inv[j,d]

对比fla: fla在token_parallel kernel里逐token循环计算 Akk 的元素;FlashKDA一次MMA直接算出完整的16×16矩阵。

为什么L用fp16而Mqk用bf16? L后面要送进 Neumann 求逆,fp16有10bit尾数(vs bf16的7bit),求逆过程中误差累积更小。Mqk直接用于最终计算,bf16够用。


第十一部分: 下三角化 + INV = I - L (第474-492行)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
if (compute_tid < 256):    # 所有线程
# 每个线程负责L和INV的一个元素
block_idx, i, j = ... (从compute_tid算出行列)

# L: 下三角化 + 乘beta
if i <= j:
L_fp16(i,j) = 0 # 上三角和对角线清零
else:
L_fp16(i,j) = L_fp16(i,j) * sigmoid(beta[i]) # 下三角乘beta

# Mqk: 因果mask
if i < j:
Mqk(i,j) = 0 # 上三角清零 (对角线保留)

# INV = I - L
x = L_fp16(i,j)
INV_fp16(i,j) = (i == j) ? (1 - x) : (-x)

256个线程处理 16×16 = 256 个元素,正好一对一。

对比fla:

  • fla在token_parallel kernel里,Akk对角线上是1,Akk(i,j) = sum_d k[i]×k[j]×exp(g_i-g_j) × beta[j] (j<i)
  • 然后在inter_solve_fused里做 I - Akk
  • FlashKDA这里 L = k_decayed @ k_inv^T 已经包含了 exp(g_i-g_j) 的效果
  • 再乘 sigmoid(beta) 就完成了

注意: beta的sigmoid是在这里才做的(K1开头TMA加载的是原始logits)。fla里beta是在外部提前sigmoid过的。


第十二部分: Neumann 级数求逆 (第494-498行)

1
neumann_inv_fused_1warp(L_fp16, INV_fp16, INV, compute_tid);

只有线程 0-31 (一个warp) 实际计算,其他线程跳过。

算法 (utils.cuh 第189-312行)

1
2
3
4
5
6
L^2 = L x L           # MMA
INV += INV x L^2 # MMA
L^4 = L^2 x L^2 # MMA
INV += INV x L^4 # MMA
L^8 = L^4 x L^4 # MMA
INV += INV x L^8 # MMA

输入: L (严格下三角, fp16), INV = I - L (fp16) 输出: INV_bf16 = (I-L)^{-1} (bf16)

数学原理

要求 (I + L)^{-1}

Neumann级数: (I + L)^{-1} = I - L + L^2 - L^3 + ... - L^{15} (L是16×16严格下三角,L^{16} = 0,级数有限项,精确!)

倍增法加速:

1
2
3
4
INV_0 = I - L
INV_1 = INV_0 + INV_0 x L^2 = (I-L)(I + L^2)
INV_2 = INV_1 + INV_1 x L^4 = (I-L)(I + L^2)(I + L^4)
INV_3 = INV_2 + INV_2 x L^8 = (I-L)(I + L^2)(I + L^4)(I + L^8)

展开: (I-L)(I+L^2)(I+L^4)(I+L^8) = I - L + L^2 - L^3 + ... - L^{15} → 这是 (I+L)^{-1} 的精确展开 (因为 L^{16}=0)

关键优化 — 全在寄存器里算

每个 16×16 MMA 的操作数用 uint32_t[4] 表示 (每个 lane 持有 4 个 fp16 值)

transpose_u32x4: 用 SM75_U32x1_MOVM_T 在寄存器内转置

  • A操作数格式 → B操作数格式
  • 不需要写回 shared memory!

mma_16x16: 两个 SM80_16x8x16_F16F16F16F16_TN 拼成 16×16

add_fp16x2_u32x4: fp16x2 向量加法 (__hadd2)

最后: 把 fp16 结果转为 bf16 写回 INV smem。


第十三部分: TMA存到workspace (第499-569行)

1
2
if (threadIdx.x == 0):
int ws_idx = head_idx * total_tiles + global_tile_idx;

workspace 按 [H × total_tiles, ...] 排列,ws_idx 是当前tile的全局索引。

6次TMA store

每次的模式相同:

  1. 构造全局 tensor 描述
  2. 算出目标偏移
  3. 构造 smem 侧的 tile
  4. cute::copy(tma_store_xxx, S, D) — 发起TMA store
  5. tma_store_arrive() — 通知硬件"有一个store要做"

存储的6个中间结果:

输出 形状 大小 含义
k_decayed [16,128] bf16 4KB k × exp2(cumsum_gate)
q_decayed [16,128] bf16 4KB q × exp2(cumsum_gate) × scale
k_restored [16,128] bf16 4KB k × exp2(g_total - cumsum_gate)
g_total [128] fp32 512B exp2(chunk总gate)
INV [16,16] bf16 512B (I - L)^{-1}
Mqk [16,16] bf16 512B q_decayed @ k_inv^T

每个tile总workspace = 3×4096 + 512 + 512 + 512 = 13824 字节 ≈ 13.5KB

等待完成

1
2
tma_store_wait<0>();      // 等待所有TMA store完成 (<0>表示等到剩余0个未完成)
__syncthreads(); // 最终同步

第十四部分: K1 总结

输入

输入 来源
q, k, g 全局内存 (TMA加载)
beta, dt_bias, A_log 全局内存 (TMA加载)

计算流程

1
2
3
4
5
6
7
8
1. L2归一化 q, k               (256线程并行)
2. gate激活 + cumsum (128线程, 每线程一列)
3. exp2(g_total) (128线程)
4. 计算 k_decayed, q_decayed, k_inv, k_restored (256线程)
5. L = k_decayed @ k_inv^T (1个warp, fp16)
Mqk = q_decayed @ k_inv^T (1个warp, bf16)
6. 下三角化 + 乘beta + INV = I - L (256线程)
7. Neumann求逆: INV = (I-L)^{-1} (1个warp, fp16→bf16)

输出 (TMA存到workspace)

1
k_decayed, q_decayed, k_restored, g_total, INV, Mqk

这些workspace数据会被 K2 kernel 读取,用于chunk间的隐状态递推。

对应fla中的步骤

FlashKDA K1 fla对应
全部融合 l2norm + gate激活 + cumsum + Step1(token_parallel) + Step2(inter_solve, 矩阵求逆) + Step3(recompute_w_u的部分)

附录:decay_apply 线程映射详解

文件位置: fwd_kernel1.cuh 的 decay_apply 部分

256个线程从 shared memory 读取 q/k/g,计算 q_decayed/k_decayed/k_inv/k_restored,写回到 shared memory 的 union 复用空间。


一、256个线程的编号规则

1
2
3
4
int lane = compute_tid % 32;    // warp内编号 0-31
int warp_id = compute_tid / 32; // warp编号 0-7 (256线程=8个warp)
int g = lane / 4; // 组号 0-7 (每4线程一组)
int t = lane % 4; // 组内编号 0-3

每个warp内部:

lane 0 1 2 3 4 5 6 7 8 9 10 11 28 29 30 31
g 0 0 0 0 1 1 1 1 2 2 2 2 7 7 7 7
t 0 1 2 3 0 1 2 3 0 1 2 3 0 1 2 3

8个warp × 每warp8组 = 64组,每组4线程,总 64×4 = 256 线程


二、16×128 矩阵切成 4 个 tile

1
2
3
4
5
CHUNK=16, D=128

N_M = CHUNK/8 = 2 行方向切2块
N_N = D/64 = 2 列方向切2块
N_TILES = 4

tile 布局:

1
2
3
4
5
6
7
8
9
10
11
12
            列 0         列 63   列 64        列 127
+-------------+ +-------------+
行 0 | | | |
行 1 | tile 0 | | tile 1 |
... | 8行x64列 | | 8行x64列 |
行 7 | | | |
+-------------+ +-------------+
行 8 | | | |
行 9 | tile 2 | | tile 3 |
... | 8行x64列 | | 8行x64列 |
行 15 | | | |
+-------------+ +-------------+

tile_idx 计算:

1
2
3
4
5
6
tile_idx = (m_blk/8) * 2 + (n_blk/64)

m_blk=0, n_blk=0 -> tile_idx=0
m_blk=0, n_blk=64 -> tile_idx=1
m_blk=8, n_blk=0 -> tile_idx=2
m_blk=8, n_blk=64 -> tile_idx=3

循环顺序: tile0 → tile1 → tile2 → tile3

每个线程每次循环取2个元素,4次循环共取 4×2 = 8 个元素


三、每线程负责的行和列

1
2
int row = m_blk + ((warp_id + g) % 8);
int col_base = n_blk + g * 8;
  • g 同时决定行和列段:

  • 列段: g=0 → col[0…7], g=1 → col[8…15], …, g=7 → col[56…63]

  • 行: (warp_id + g) % 8 做循环移位

单个warp (warp0) 的分配:

段0 col0-7 段1 col8-15 段2 col16-23 段3 col24-31 段4 col32-39 段5 col40-47 段6 col48-55 段7 col56-63
row0 w0,g=0
row1 w0,g=1
row2 w0,g=2
row3 w0,g=3
row4 w0,g=4
row5 w0,g=5
row6 w0,g=6
row7 w0,g=7

warp0 走对角线!

warp1 (row = (1+g)%8) 偏移一格:

段0 col0-7 段1 col8-15 段2 col16-23 段3 col24-31 段4 col32-39 段5 col40-47 段6 col48-55 段7 col56-63
row1 w1,g=0
row2 w1,g=1
row3 w1,g=2
row4 w1,g=3
row5 w1,g=4
row6 w1,g=5
row7 w1,g=6
row0 w1,g=7

四、8个warp叠加: 循环移位矩阵

全部8个warp叠加后,每个格子恰好被一个warp负责:

段0 段1 段2 段3 段4 段5 段6 段7
row0 w0 w7 w6 w5 w4 w3 w2 w1
row1 w1 w0 w7 w6 w5 w4 w3 w2
row2 w2 w1 w0 w7 w6 w5 w4 w3
row3 w3 w2 w1 w0 w7 w6 w5 w4
row4 w4 w3 w2 w1 w0 w7 w6 w5
row5 w5 w4 w3 w2 w1 w0 w7 w6
row6 w6 w5 w4 w3 w2 w1 w0 w7
row7 w7 w6 w5 w4 w3 w2 w1 w0

循环移位矩阵! 保证同一时刻每行每列段只有一个warp在访问。


五、一个格子 (1行×8列) 内4线程的分工

每个格子 8 列,由同组的 4 个线程 (t=0,1,2,3) 瓜分,每线程 2 列:

1
2
3
4
5
6
格子: row=R, col_base=C (8个元素)

col: C+0 C+1 C+2 C+3 C+4 C+5 C+6 C+7
\_____/ \_____/ \_____/ \_____/
t=0 t=1 t=2 t=3
2元素 2元素 2元素 2元素

六、local_tile 图解

第一步: local_tile(g_tile, (1,8), make_coord(row, col_tile))

1
2
3
4
5
6
7
8
9
10
g_tile [16行][128列]:

col_tile: 0 1 2 3 ... 15
8列 8列 8列 8列 8列
row 0: [...] [...] [...] ...
row 1: [...] [...] [...] ...
...
row 3: [tile_g] ... <-- make_coord(3, 2)
...
row 15: [...] [...] [...] ...

含义:

  • 把 [16,128] 按 (1行, 8列) 粒度切分
  • 行方向: 每1行一个tile,共16个,编号0-15
  • 列方向: 每8列一个tile,共16个,编号0-15
  • make_coord(row, col_tile) = 取第row行、第col_tile个列段
  • 结果 tile_g: shape=(1,8),指向 g_tile[row][col_tile×8 .. col_tile×8+7]

第二步: local_tile(tile_g, (1,2), make_coord(0, t))

1
2
3
4
tile_g: [g[3][16], g[3][17], g[3][18], g[3][19], g[3][20], g[3][21], g[3][22], g[3][23]]
\_____________/ \_____________/ \_____________/ \_____________/
t=0 的 s_g t=1 的 s_g t=2 的 s_g t=3 的 s_g
make_coord(0,0) make_coord(0,1) make_coord(0,2) make_coord(0,3)

含义:

  • 把 (1,8) 按 (1,2) 切分
  • 列方向: 8/2=4个tile,编号0-3
  • make_coord(0, t) = 取第0行tile、第t个列tile
  • 结果 s_g: shape=(1,2),指向 shared memory 中的2个连续元素

对于 1D 的 g_total [128]:

1
2
3
4
5
6
7
vec8_1d = (8)
local_tile(g_total, (8), make_coord(col_tile))
// 从128元素中按8个一组取第col_tile组,得到8元素

thr2_1d = (2)
local_tile(tile_gt, (2), make_coord(t))
// 从8元素中按2个一组取第t组,得到2元素

七、为什么用循环移位: 避免 bank conflict

Shared memory 有 32 个 bank,每 bank 4字节宽。bf16 元素 2字节,所以每 bank 放 2 个 bf16。

如果 warp 内 32 线程访问同一行连续 128 列:

线程 访问列 bank
0 col0,1 bank 0
1 col2,3 bank 1
31 col62,63 bank 31

完美无冲突!

但如果多线程访问同一列段的不同行:

可能多线程撞同一 bank → 串行化 → 慢!

循环移位保证: 同一 warp 内的 8 组 (g=0…7) 访问不同行、不同列段 → 落入不同 bank → 无冲突


八、完整流程总图

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
Phase A: shared memory 存着 q[16][128], k[16][128], g[16][128] (gate cumsum后)

+-- 循环 4 个 tile (8x64 each) --+
| |
| 每线程读 2 元素到寄存器 |
| reg_g[tile][0..1] |
| reg_q[tile][0..1] |
| reg_k[tile][0..1] |
| reg_gt[tile][0..1] |
+----------------------------------+
|
__syncthreads()
(等所有人读完, 释放Phase A空间)
|
+-- 循环 4 个 tile --+
| |
| 从寄存器计算: |
| q_decayed = q * exp2(g) * scale |
| k_decayed = k * exp2(g) |
| k_inv = k * exp2(-g) |
| k_restored= k * exp2(-g) * exp2(gt) |
| |
| 写回 Phase B shared memory |
+---------------------+

Phase B: shared memory 现在存着 q_decayed, k_decayed, k_inv, k_restored
(与Phase A共用同一片物理空间, union)

九、计算公式

1
2
3
4
5
q_decayed[row][col]  = q[row][col] × exp2(g[row][col]) × scale
k_decayed[row][col] = k[row][col] × exp2(g[row][col])
k_inv[row][col] = k[row][col] × exp2(-g[row][col])
k_restored[row][col] = k[row][col] × exp2(-g[row][col]) × exp2(g_total[col])
= k[row][col] × exp2(g_total[col] - g[row][col])

代码实现:

1
2
3
4
5
6
7
8
float g = reg_g[tile_idx][v];
BF16 exp_cumsum = BF16(ex2_approx_ftz_f32(g)); // exp2(g)
r_qd = q * exp_cumsum * scale; // q_decayed
r_kd = k * exp_cumsum; // k_decayed

BF16 inv_cumsum = BF16(ex2_approx_ftz_f32(-g)); // exp2(-g)
r_ki = k * inv_cumsum; // k_inv
r_kr = k * inv_cumsum * BF16(reg_gt[tile_idx][v]); // k_restored

注意: reg_gt 存的是 exp2(g_total[col]),在 gate+cumsum 阶段已经算好。 所以 k_restored = k × exp2(-g) × exp2(g_total) = k × exp2(g_total - g)

  • 标题: KDA源码剖析之FlashKDA(上)
  • 作者: 鱿鱼圈
  • 创建于 : 2026-06-05 23:50:00
  • 更新于 : 2026-06-14 15:45:15
  • 链接: https://yuyanqi.com/2026/06/05/KDA源码剖析之FlashKDA(上)/
  • 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。
评论
目录
KDA源码剖析之FlashKDA(上)