KDA源码剖析之FlashKDA(下)

鱿鱼圈 Lv4

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

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

公式回顾

kernel 2 代码

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
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
#pragma once

// TMA_DISABLE_ALL: when defined, disable load/store warps entirely
// and let MMA warps work without pipeline synchronization
// #define TMA_DISABLE_ALL

#include "utils.cuh"

template <int D, int CHUNK = 16>
struct K2Layouts {
using MMALayout = decltype(tile_to_shape(
GMMA::Layout_K_INTER_Atom<cute::bfloat16_t>{},
make_shape(Int<CHUNK>{}, Int<D>{}),
LayoutLeft{}
));
using TransposedMMALayout = decltype(tile_to_shape(
GMMA::Layout_MN_INTER_Atom<cute::bfloat16_t>{},
make_shape(Int<D>{}, Int<CHUNK>{}),
LayoutRight{}
));
using VOLayout = MMALayout;
using TransposedVOLayout = TransposedMMALayout;
using BetaSmemLayout = Layout<Shape<Int<32>>, Stride<Int<1>>>;
using StateSmemLayout = decltype(tile_to_shape(
GMMA::Layout_K_INTER_Atom<cute::bfloat16_t>{},
make_shape(Int<D>{}, Int<D>{}),
LayoutLeft{}
));
using TransposedStateSmemLayout = decltype(tile_to_shape(
GMMA::Layout_MN_INTER_Atom<cute::bfloat16_t>{},
make_shape(Int<D>{}, Int<D>{}),
LayoutRight{}
));
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 TMABetaSmemLayout = BetaSmemLayout; // 1D TMA, no dummy dim
using TMAVOLayout = decltype(composition(
VOLayout{}.layout_a(),
VOLayout{}.offset(),
prepend(VOLayout{}.layout_b())
));
using TMAStateSmemLayout = decltype(composition(
StateSmemLayout{}.layout_a(),
StateSmemLayout{}.offset(),
prepend(StateSmemLayout{}.layout_b())
));
using TMALMLayout = decltype(composition(
LMLayout{}.layout_a(),
LMLayout{}.offset(),
prepend(LMLayout{}.layout_b())
));
using TMAGTotalSmemLayout = decltype(prepend(GTotalLayout{}));

// FP32 state layout (K_SW32 atom, same 8x8 atom structure as K_INTER bf16)
using FP32StateSmemLayout = decltype(tile_to_shape(
GMMA::Layout_K_SW32_Atom<float>{},
make_shape(Int<D>{}, Int<D>{}),
LayoutLeft{}
));
using TMAFP32StateSmemLayout = decltype(composition(
FP32StateSmemLayout{}.layout_a(),
FP32StateSmemLayout{}.offset(),
prepend(FP32StateSmemLayout{}.layout_b())
));
};

template <class Layouts, int InputStages, int OutputStages>
struct SharedStorageK2 {
using BF16 = cutlass::bfloat16_t;
using VOLayout = typename Layouts::VOLayout;
using BetaSmemLayout = typename Layouts::BetaSmemLayout;
using StateSmemLayout = typename Layouts::StateSmemLayout;
using GTotalLayout = typename Layouts::GTotalLayout;
using LMLayout = typename Layouts::LMLayout;
using MMALayout = typename Layouts::MMALayout;

alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<StateSmemLayout>> state_acc;

struct InputStorage {
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<VOLayout>> v;
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<BetaSmemLayout>> beta;
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_restored;
alignas(128) cute::ArrayEngine<float, cute::cosize_v<GTotalLayout>> g_total;
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<LMLayout>> INV;
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<LMLayout>> Mqk;
};

struct OutputStorage {
alignas(128) cute::ArrayEngine<BF16, cute::cosize_v<VOLayout>> out;
};

// Anonymous union: pipeline buffers share space with fp32 state conversion buffer.
// FP32 state load/store happens before/after the pipeline loop, so no overlap.
union {
struct {
InputStorage input[InputStages];
OutputStorage output[OutputStages];
};
alignas(128) char state_fp32_buf[cute::cosize_v<StateSmemLayout> * sizeof(float)];
};

typename cutlass::PipelineTmaAsync<InputStages>::SharedStorage load_pipeline;
typename cutlass::PipelineAsync<OutputStages>::SharedStorage store_pipeline;
alignas(16) cutlass::arch::ClusterTransactionBarrier state_acc_tma_barrier;
};

// ==================== Kernel 2: Recurrence ====================
template <
class TmaLoadV,
class TmaLoadBeta,
class TmaLoadWsKD, class TmaLoadWsQD, class TmaLoadWsKR,
class TmaLoadWsGT, class TmaLoadWsINV, class TmaLoadWsMqk,
class TmaLoadState,
class TmaStoreState,
class TmaStoreOut,
int CHUNK,
int D,
int InputStages,
int OutputStages,
int NumThreads,
bool HasStateIn = true,
bool HasStateOut = true,
bool StateFP32 = false,
bool IsVarlen = true
>
__global__ void __launch_bounds__(NumThreads) _flash_kda_fwd_recurrence(
CUTE_GRID_CONSTANT TmaLoadV const tma_load_v,
CUTE_GRID_CONSTANT TmaLoadBeta const tma_load_beta,
CUTE_GRID_CONSTANT TmaLoadWsKD const tma_load_ws_kd,
CUTE_GRID_CONSTANT TmaLoadWsQD const tma_load_ws_qd,
CUTE_GRID_CONSTANT TmaLoadWsKR const tma_load_ws_kr,
CUTE_GRID_CONSTANT TmaLoadWsGT const tma_load_ws_gt,
CUTE_GRID_CONSTANT TmaLoadWsINV const tma_load_ws_inv,
CUTE_GRID_CONSTANT TmaLoadWsMqk const tma_load_ws_mqk,
CUTE_GRID_CONSTANT TmaLoadState const tma_load_initial_state,
CUTE_GRID_CONSTANT TmaStoreState const tma_store_final_state,
CUTE_GRID_CONSTANT TmaStoreOut const tma_store_out,
cutlass::bfloat16_t* out_raw_ptr,
int T_total,
int H,
int N,
int64_t const* cu_seqlens,
int total_tiles
) {
using BF16 = cutlass::bfloat16_t;
using FP16 = cutlass::half_t;
using Layouts = K2Layouts<D, CHUNK>;
using MMALayout = typename Layouts::MMALayout;
using TransposedMMALayout = typename Layouts::TransposedMMALayout;
using VOLayout = typename Layouts::VOLayout;
using TransposedVOLayout = typename Layouts::TransposedVOLayout;
using BetaSmemLayout = typename Layouts::BetaSmemLayout;
using StateSmemLayout = typename Layouts::StateSmemLayout;
using TransposedStateSmemLayout = typename Layouts::TransposedStateSmemLayout;
using GTotalLayout = typename Layouts::GTotalLayout;
using LMLayout = typename Layouts::LMLayout;
using TMAVOLayout = typename Layouts::TMAVOLayout;
using TMABetaSmemLayout = typename Layouts::TMABetaSmemLayout;
using TMAStateSmemLayout = typename Layouts::TMAStateSmemLayout;
using TMALMLayout = typename Layouts::TMALMLayout;
using TMAGTotalSmemLayout = typename Layouts::TMAGTotalSmemLayout;
constexpr int kWarpSize = 32;
constexpr int kComputeThreads = 128;

// Transaction bytes: v + beta + k_decayed + q_decayed + k_restored + g_total + INV + Mqk
constexpr uint32_t kTmaTransactionBytes =
#ifndef TMA_DISABLE_ALL
uint32_t(cute::cosize_v<VOLayout>) * uint32_t(sizeof(BF16)) +
uint32_t(32) * uint32_t(sizeof(BF16)) + // beta (bf16, sigmoid fused)
uint32_t(cute::cosize_v<MMALayout>) * uint32_t(sizeof(BF16)) * 3 + // k_decayed, q_decayed, k_restored
uint32_t(cute::cosize_v<GTotalLayout>) * uint32_t(sizeof(float)) + // g_total
uint32_t(cute::cosize_v<LMLayout>) * uint32_t(sizeof(BF16)) * 2 + // INV, Mqk
#endif
0u;

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

// --- warp specialization
int warp_id = threadIdx.x / kWarpSize;
WarpRole warp_role = WarpRole::NonParticipant;
if (warp_id < kComputeThreads / kWarpSize) {
warp_role = WarpRole::MMA;
} else if (warp_id < kComputeThreads / kWarpSize + 1) {
warp_role = WarpRole::LOAD_QKG;
} else if (warp_id < kComputeThreads / kWarpSize + 2) {
warp_role = WarpRole::STORE;
}

#ifndef TMA_DISABLE_ALL
using LoadPipelineState = cutlass::PipelineState<InputStages>;
using LoadPipeline = cutlass::PipelineTmaAsync<InputStages>;
LoadPipeline load_pipeline = make_load_pipeline<InputStages>(
shared_storage.load_pipeline,
kTmaTransactionBytes,
warp_role, 1, kComputeThreads
);
using StorePipelineState = cutlass::PipelineState<OutputStages>;
using StorePipeline = cutlass::PipelineAsync<OutputStages>;
StorePipeline store_pipeline = make_store_pipeline<OutputStages>(
shared_storage.store_pipeline,
warp_role, kComputeThreads, 1
);
#endif

// --- per-block sequence info
int seq_idx = blockIdx.x;
int head_idx = blockIdx.y;
int64_t bos, eos;
int tile_base;

if constexpr (IsVarlen) {
bos = cu_seqlens[seq_idx];
eos = cu_seqlens[seq_idx + 1];
// Compute tile_base via linear scan (no host-precomputed table)
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;
bool lane_predicate = cute::elect_one_sync();

// --- Load initial state
#ifndef TMA_DISABLE_ALL
if constexpr (HasStateIn && !StateFP32) {
// BF16 state: TMA load directly into state_acc
if (warp_role == WarpRole::LOAD_QKG && lane_predicate) {
using BarrierType = cutlass::arch::ClusterTransactionBarrier::ValueType;
constexpr uint32_t kStateTransactionBytes = cute::cosize_v<StateSmemLayout> * sizeof(BF16);

shared_storage.state_acc_tma_barrier.init(1);
shared_storage.state_acc_tma_barrier.arrive_and_expect_tx(kStateTransactionBytes);

Tensor g_init = tma_load_initial_state.get_tma_tensor(make_shape(N * H, D, D));
auto init_off = g_init.layout()(seq_idx * H + head_idx, 0, 0);
Tensor g_init_tile = make_tensor(g_init.data() + init_off,
make_layout(make_shape(Int<1>{}, Int<D>{}, Int<D>{}), stride(g_init.layout())));
Tensor s_state = make_tensor(make_smem_ptr(shared_storage.state_acc.begin()), TMAStateSmemLayout{});

auto cta_tma_load_state = tma_load_initial_state.get_slice(Int<0>{});
cute::copy(
tma_load_initial_state.with(reinterpret_cast<BarrierType&>(shared_storage.state_acc_tma_barrier)),
cta_tma_load_state.partition_S(g_init_tile),
cta_tma_load_state.partition_D(s_state)
);
}
__syncthreads();
shared_storage.state_acc_tma_barrier.wait(0);
cutlass::arch::fence_view_async_shared();
} else if constexpr (HasStateIn && StateFP32) {
// FP32 state: TMA load fp32 into pipeline buffer, then convert to bf16 in state_acc
using FP32StateSmemLayout = typename Layouts::FP32StateSmemLayout;
using TMAFP32StateSmemLayout = typename Layouts::TMAFP32StateSmemLayout;

if (warp_role == WarpRole::LOAD_QKG && lane_predicate) {
using BarrierType = cutlass::arch::ClusterTransactionBarrier::ValueType;
constexpr uint32_t kFP32StateTransactionBytes = cute::cosize_v<StateSmemLayout> * sizeof(float);

shared_storage.state_acc_tma_barrier.init(1);
shared_storage.state_acc_tma_barrier.arrive_and_expect_tx(kFP32StateTransactionBytes);

Tensor g_init = tma_load_initial_state.get_tma_tensor(make_shape(N * H, D, D));
auto init_off = g_init.layout()(seq_idx * H + head_idx, 0, 0);
Tensor g_init_tile = make_tensor(g_init.data() + init_off,
make_layout(make_shape(Int<1>{}, Int<D>{}, Int<D>{}), stride(g_init.layout())));
Tensor s_fp32 = make_tensor(
make_smem_ptr(reinterpret_cast<float*>(shared_storage.state_fp32_buf)),
TMAFP32StateSmemLayout{});

auto cta_tma_load_state = tma_load_initial_state.get_slice(Int<0>{});
cute::copy(
tma_load_initial_state.with(reinterpret_cast<BarrierType&>(shared_storage.state_acc_tma_barrier)),
cta_tma_load_state.partition_S(g_init_tile),
cta_tma_load_state.partition_D(s_fp32)
);
}
__syncthreads();
shared_storage.state_acc_tma_barrier.wait(0);
cutlass::arch::fence_view_async_shared();

// All threads: convert fp32 -> bf16 with layout transformation
smem_cvt_fp32_to_bf16<FP32StateSmemLayout, StateSmemLayout, D, NumThreads>(
reinterpret_cast<float*>(shared_storage.state_fp32_buf),
shared_storage.state_acc.begin(),
threadIdx.x);
__syncthreads();
} else {
// No state in: zero-initialize state_acc
{
BF16* buf = shared_storage.state_acc.begin();
constexpr int kTotal = cute::cosize_v<StateSmemLayout>;
for (int i = threadIdx.x; i < kTotal; i += NumThreads) {
buf[i] = BF16(0);
}
}
__syncthreads();
}
#endif

#ifndef TMA_DISABLE_ALL
__syncthreads();

// --- LOAD warp: issue TMA loads for v, beta, and workspace intermediates
if (warp_role == WarpRole::LOAD_QKG && lane_predicate) {
Tensor g_v = tma_load_v.get_tma_tensor(make_shape(H, T_total, D));
Tensor g_beta = tma_load_beta.get_tma_tensor(make_shape(H * T_total));

// Workspace gmem tensors
auto g_ws_kd = tma_load_ws_kd.get_tma_tensor(make_shape(H * total_tiles, CHUNK, D));
auto g_ws_qd = tma_load_ws_qd.get_tma_tensor(make_shape(H * total_tiles, CHUNK, D));
auto g_ws_kr = tma_load_ws_kr.get_tma_tensor(make_shape(H * total_tiles, CHUNK, D));
auto g_ws_gt = tma_load_ws_gt.get_tma_tensor(make_shape(H * total_tiles, D));
auto g_ws_inv = tma_load_ws_inv.get_tma_tensor(make_shape(H * total_tiles, CHUNK, CHUNK));
auto g_ws_mqk = tma_load_ws_mqk.get_tma_tensor(make_shape(H * total_tiles, CHUNK, CHUNK));

LoadPipelineState load_write = cutlass::make_producer_start_state<LoadPipeline>();
auto cta_tma_load_v = tma_load_v.get_slice(Int<0>{});
auto cta_tma_load_beta = tma_load_beta.get_slice(Int<0>{});
auto cta_ws_kd = tma_load_ws_kd.get_slice(Int<0>{});
auto cta_ws_qd = tma_load_ws_qd.get_slice(Int<0>{});
auto cta_ws_kr = tma_load_ws_kr.get_slice(Int<0>{});
auto cta_ws_gt = tma_load_ws_gt.get_slice(Int<0>{});
auto cta_ws_inv = tma_load_ws_inv.get_slice(Int<0>{});
auto cta_ws_mqk = tma_load_ws_mqk.get_slice(Int<0>{});

for (int t = 0; t < t_tiles; ++t) {
load_pipeline.producer_acquire(load_write);
using LoadBarrierType = typename LoadPipeline::ProducerBarrierType;
LoadBarrierType* tma_barrier = load_pipeline.producer_get_barrier(load_write);
int stage = load_write.index();
int ws_idx = head_idx * total_tiles + tile_base + t;

// TMA load v
auto v_off = g_v.layout()(head_idx, int(bos) + t * CHUNK, 0);
Tensor g_v_tile = make_tensor(g_v.data() + v_off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<D>{}), stride(g_v.layout())));
Tensor s_v_tile = make_tensor(make_smem_ptr(shared_storage.input[stage].v.begin()), TMAVOLayout{});
cute::copy(tma_load_v.with(*tma_barrier),
cta_tma_load_v.partition_S(g_v_tile), cta_tma_load_v.partition_D(s_v_tile));

// TMA load beta (1D)
int beta_linear = head_idx * T_total + (int(bos) + 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_beta_tile = make_tensor(make_smem_ptr(shared_storage.input[stage].beta.begin()), TMABetaSmemLayout{});
cute::copy(tma_load_beta.with(*tma_barrier),
cta_tma_load_beta.partition_S(g_beta_tile), cta_tma_load_beta.partition_D(s_beta_tile));

// TMA load workspace: k_decayed
{
auto off = g_ws_kd.layout()(ws_idx, 0, 0);
Tensor g_tile = make_tensor(g_ws_kd.data() + off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<D>{}), stride(g_ws_kd.layout())));
Tensor s_tile = make_tensor(make_smem_ptr(shared_storage.input[stage].k_decayed.begin()), TMAVOLayout{});
cute::copy(tma_load_ws_kd.with(*tma_barrier), cta_ws_kd.partition_S(g_tile), cta_ws_kd.partition_D(s_tile));
}
// q_decayed
{
auto off = g_ws_qd.layout()(ws_idx, 0, 0);
Tensor g_tile = make_tensor(g_ws_qd.data() + off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<D>{}), stride(g_ws_qd.layout())));
Tensor s_tile = make_tensor(make_smem_ptr(shared_storage.input[stage].q_decayed.begin()), TMAVOLayout{});
cute::copy(tma_load_ws_qd.with(*tma_barrier), cta_ws_qd.partition_S(g_tile), cta_ws_qd.partition_D(s_tile));
}
// k_restored
{
auto off = g_ws_kr.layout()(ws_idx, 0, 0);
Tensor g_tile = make_tensor(g_ws_kr.data() + off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<D>{}), stride(g_ws_kr.layout())));
Tensor s_tile = make_tensor(make_smem_ptr(shared_storage.input[stage].k_restored.begin()), TMAVOLayout{});
cute::copy(tma_load_ws_kr.with(*tma_barrier), cta_ws_kr.partition_S(g_tile), cta_ws_kr.partition_D(s_tile));
}
// g_total
{
auto off = g_ws_gt.layout()(ws_idx, 0);
Tensor g_tile = make_tensor(g_ws_gt.data() + off,
make_layout(make_shape(Int<1>{}, Int<D>{}), stride(g_ws_gt.layout())));
Tensor s_tile = make_tensor(make_smem_ptr(shared_storage.input[stage].g_total.begin()), TMAGTotalSmemLayout{});
cute::copy(tma_load_ws_gt.with(*tma_barrier), cta_ws_gt.partition_S(g_tile), cta_ws_gt.partition_D(s_tile));
}
// INV
{
auto off = g_ws_inv.layout()(ws_idx, 0, 0);
Tensor g_tile = make_tensor(g_ws_inv.data() + off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<CHUNK>{}), stride(g_ws_inv.layout())));
Tensor s_tile = make_tensor(make_smem_ptr(shared_storage.input[stage].INV.begin()), TMALMLayout{});
cute::copy(tma_load_ws_inv.with(*tma_barrier), cta_ws_inv.partition_S(g_tile), cta_ws_inv.partition_D(s_tile));
}
// Mqk
{
auto off = g_ws_mqk.layout()(ws_idx, 0, 0);
Tensor g_tile = make_tensor(g_ws_mqk.data() + off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<CHUNK>{}), stride(g_ws_mqk.layout())));
Tensor s_tile = make_tensor(make_smem_ptr(shared_storage.input[stage].Mqk.begin()), TMALMLayout{});
cute::copy(tma_load_ws_mqk.with(*tma_barrier), cta_ws_mqk.partition_S(g_tile), cta_ws_mqk.partition_D(s_tile));
}

++load_write;
}
load_pipeline.producer_tail(load_write);
}
#endif

// --- MMA warps
if (warp_role == WarpRole::MMA) {
cutlass::arch::NamedBarrier compute_barrier(kComputeThreads, 0);
#ifndef TMA_DISABLE_ALL
LoadPipelineState load_read;
StorePipelineState out_write = cutlass::make_producer_start_state<StorePipeline>();
#endif
int compute_tid = threadIdx.x;

for (int t = 0; t < t_tiles; ++t) {
#ifndef TMA_DISABLE_ALL
store_pipeline.producer_acquire(out_write);
load_pipeline.consumer_wait(load_read);
int load_stage = load_read.index();
int out_stage = out_write.index();
#else
constexpr int load_stage = 0;
constexpr int out_stage = 0;
#endif

Tensor v_tile = make_tensor(make_smem_ptr(shared_storage.input[load_stage].v.begin()), VOLayout{});
Tensor beta_tile = make_tensor(make_smem_ptr(shared_storage.input[load_stage].beta.begin()), BetaSmemLayout{});
int beta_smem_offset = (head_idx * T_total + int(bos) + t * CHUNK) & 7;
Tensor out_tile = make_tensor(make_smem_ptr(shared_storage.output[out_stage].out.begin()), VOLayout{});

Tensor k_decayed = make_tensor(make_smem_ptr(shared_storage.input[load_stage].k_decayed.begin()), MMALayout{});
Tensor q_decayed = make_tensor(make_smem_ptr(shared_storage.input[load_stage].q_decayed.begin()), MMALayout{});
Tensor k_restored = make_tensor(make_smem_ptr(shared_storage.input[load_stage].k_restored.begin()), MMALayout{});
Tensor g_total = make_tensor(make_smem_ptr(shared_storage.input[load_stage].g_total.begin()), GTotalLayout{});
Tensor INV = make_tensor(make_smem_ptr(shared_storage.input[load_stage].INV.begin()), LMLayout{});
Tensor Mqk = make_tensor(make_smem_ptr(shared_storage.input[load_stage].Mqk.begin()), LMLayout{});

Tensor s_acc = make_tensor(make_smem_ptr(shared_storage.state_acc.begin()), StateSmemLayout{});
Tensor s_acc_T = make_tensor(make_smem_ptr(shared_storage.state_acc.begin()), TransposedStateSmemLayout{});

// Fused MMA: v_sub, v_beta, U=INV@v, out=q@s, out+=Mqk@U, s_acc_update
// Each warp handles TWO 16x16 column blocks (N=128 / 4 warps = 32 = 2 x 16)
// U stays in registers via SM75_U32x1_MOVM_T (no smem round-trip)
{
Tensor k_restored_t = make_tensor(make_smem_ptr(shared_storage.input[load_stage].k_restored.begin()), TransposedMMALayout{});

constexpr int PREFETCH = 1;

auto mma = make_tiled_mma(
MMA_Atom<SM80_16x8x16_F32BF16BF16F32_TN>{},
Layout<Shape<_1,_1>>{},
Tile<_16,_16,_16>{}
);

const int warp_id = compute_tid / 32;
const int lane_id = compute_tid % 32;
const int group_id = (lane_id / 4) % 8;

ThrMMA thr_mma = mma.get_slice(lane_id);

// A copy: K_INTER → LDSM_N (for k_decayed, q_decayed, INV, Mqk)
auto smem_tiled_copy_A = make_tiled_copy_A(Copy_Atom<SM75_U32x4_LDSM_N, BF16>{}, mma);
auto smem_thr_copy_A = smem_tiled_copy_A.get_thread_slice(lane_id);

// A copy: MN_INTER → LDSM_T (for k_restored_t in Phase 7)
auto smem_tiled_copy_A_T = make_tiled_copy_A(Copy_Atom<SM75_U16x8_LDSM_T, BF16>{}, mma);
auto smem_thr_copy_A_T = smem_tiled_copy_A_T.get_thread_slice(lane_id);

// B copy: K_INTER → LDSM_N
auto smem_tiled_copy_B = make_tiled_copy_B(Copy_Atom<SM75_U32x4_LDSM_N, BF16>{}, mma);
auto smem_thr_copy_B = smem_tiled_copy_B.get_thread_slice(lane_id);

// C load/store
auto smem_tiled_load_C = make_tiled_copy_C(Copy_Atom<SM75_U32x4_LDSM_N, BF16>{}, mma);
auto smem_thr_load_C = smem_tiled_load_C.get_slice(lane_id);
auto smem_tiled_store_C = make_tiled_copy_C(Copy_Atom<SM90_U32x4_STSM_N, BF16>{}, mma);
auto smem_thr_store_C = smem_tiled_store_C.get_slice(lane_id);

// C load/store transposed (for Phase 6 state access via s_acc_T)
auto smem_tiled_load_C_T = make_tiled_copy_C(Copy_Atom<SM75_U16x8_LDSM_T, BF16>{}, mma);
auto smem_thr_load_C_T = smem_tiled_load_C_T.get_slice(lane_id);
auto smem_tiled_store_C_T = make_tiled_copy_C(Copy_Atom<SM90_U16x8_STSM_T, BF16>{}, mma);
auto smem_thr_store_C_T = smem_tiled_store_C_T.get_slice(lane_id);

Tensor A_ref = local_tile(k_decayed, make_shape(Int<16>{}, Int<16>{}), make_coord(0, 0));
Tensor B_ref = local_tile(s_acc, make_shape(Int<16>{}, Int<16>{}), make_coord(0, 0));
Tensor C_ref = local_tile(v_tile, make_shape(Int<16>{}, Int<16>{}), make_coord(0, 0));

Tensor tCrAi_k = make_fragment_like<BF16>(thr_mma.partition_fragment_A(A_ref));
auto tCrAi_k_view = smem_thr_copy_A.retile_D(tCrAi_k);
auto tCrA_k = thr_mma.partition_fragment_A(A_ref);

Tensor tCrAi_q = make_fragment_like<BF16>(thr_mma.partition_fragment_A(A_ref));
auto tCrAi_q_view = smem_thr_copy_A.retile_D(tCrAi_q);
auto tCrA_q = thr_mma.partition_fragment_A(A_ref);

Tensor tCrBi = make_fragment_like<BF16>(thr_mma.partition_fragment_B(B_ref));
auto tCrBi_view = smem_thr_copy_B.retile_D(tCrBi);
auto tCrB = thr_mma.partition_fragment_B(B_ref);

auto tCrC_ref = thr_mma.partition_C(C_ref);

using AccFragT = decltype(thr_mma.make_fragment_C(tCrC_ref));
using SFragT = decltype(make_fragment_like<BF16>(thr_mma.make_fragment_C(tCrC_ref)));
using AFragT = decltype(thr_mma.partition_fragment_A(A_ref));
using BFragT_u = decltype(thr_mma.partition_fragment_B(B_ref));

AccFragT u_acc[2], out_acc[2];
#pragma unroll
for (int i = 0; i < 2; ++i) { u_acc[i] = thr_mma.make_fragment_C(tCrC_ref); clear(u_acc[i]); }
#pragma unroll
for (int i = 0; i < 2; ++i) { out_acc[i] = thr_mma.make_fragment_C(tCrC_ref); clear(out_acc[i]); }

// ======== Phase 1: Dual GEMM k@s and q@s (k-loop, 2 blocks per warp) ========
constexpr int K_BLOCKS = decltype(cute::size<1>(k_decayed))::value / 16;

copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(
local_tile(k_decayed, make_shape(Int<16>{}, Int<16>{}), make_coord(0, 0))), tCrAi_k_view);
copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(
local_tile(q_decayed, make_shape(Int<16>{}, Int<16>{}), make_coord(0, 0))), tCrAi_q_view);
copy(smem_tiled_copy_B, smem_thr_copy_B.partition_S(
local_tile(s_acc, make_shape(Int<16>{}, Int<16>{}), make_coord(warp_id * 2, 0))), tCrBi_view);

#pragma unroll
for (int k = 0; k < K_BLOCKS; ++k) {
cute::transform(tCrAi_k, tCrA_k, cute::identity{});
cute::transform(tCrAi_q, tCrA_q, cute::identity{});
cute::transform(tCrBi, tCrB, cute::identity{});

copy(smem_tiled_copy_B, smem_thr_copy_B.partition_S(
local_tile(s_acc, make_shape(Int<16>{}, Int<16>{}), make_coord(warp_id * 2 + 1, k))), tCrBi_view);

gemm(thr_mma, tCrA_k(_,_,Int<0>{}), tCrB(_,_,Int<0>{}), u_acc[0]);
gemm(thr_mma, tCrA_q(_,_,Int<0>{}), tCrB(_,_,Int<0>{}), out_acc[0]);

cute::transform(tCrBi, tCrB, cute::identity{});

if (k + 1 < K_BLOCKS) {
copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(
local_tile(k_decayed, make_shape(Int<16>{}, Int<16>{}), make_coord(0, k + 1))), tCrAi_k_view);
copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(
local_tile(q_decayed, make_shape(Int<16>{}, Int<16>{}), make_coord(0, k + 1))), tCrAi_q_view);
copy(smem_tiled_copy_B, smem_thr_copy_B.partition_S(
local_tile(s_acc, make_shape(Int<16>{}, Int<16>{}), make_coord(warp_id * 2, k + 1))), tCrBi_view);
}

gemm(thr_mma, tCrA_k(_,_,Int<0>{}), tCrB(_,_,Int<0>{}), u_acc[1]);
gemm(thr_mma, tCrA_q(_,_,Int<0>{}), tCrB(_,_,Int<0>{}), out_acc[1]);
}

// ======== Phase 2: Cast out (keep in regs), load v/INV/beta ========
SFragT out_bf16[2];
#pragma unroll
for (int i = 0; i < 2; ++i)
cute::transform(out_acc[i], out_bf16[i], [] __device__ (float x) { return BF16(x); });

SFragT v_bf16[2];
#pragma unroll
for (int i = 0; i < 2; ++i) {
Tensor v_block = local_tile(v_tile, make_shape(Int<16>{}, Int<16>{}), make_coord(0, warp_id * 2 + i));
copy(smem_tiled_load_C, smem_thr_load_C.partition_S(v_block), smem_thr_load_C.retile_D(v_bf16[i]));
}

copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(INV), tCrAi_k_view);
cute::transform(tCrAi_k, tCrA_k, cute::identity{});

BF16 beta0 = BF16(sigmoid_tanh_approx_f32(float(beta_tile(beta_smem_offset + group_id))));
BF16 beta1 = BF16(sigmoid_tanh_approx_f32(float(beta_tile(beta_smem_offset + group_id + 8))));

// ======== Phase 3: u = (v - u) * beta; u = INV @ u (per block) ========
SFragT u_bf16[2];
uint32_t u_b_regs[4];

#pragma unroll
for (int i = 0; i < 2; ++i) {
cute::transform(u_acc[i], u_bf16[i], [] __device__ (float x) { return BF16(x); });

#pragma unroll
for (int a = 0; a < 2; ++a) {
#pragma unroll
for (int d = 0; d < 2; ++d) {
auto c0 = make_coord(make_coord(a, 0), 0, d);
auto c1 = make_coord(make_coord(a, 1), 0, d);
u_bf16[i](c0) = (v_bf16[i](c0) - u_bf16[i](c0)) * beta0;
u_bf16[i](c1) = (v_bf16[i](c1) - u_bf16[i](c1)) * beta1;
}
}

uint32_t* u_c = reinterpret_cast<uint32_t*>(&u_bf16[i](0));
SM75_U32x1_MOVM_T::copy(u_c[0], u_b_regs[0]);
SM75_U32x1_MOVM_T::copy(u_c[1], u_b_regs[1]);
SM75_U32x1_MOVM_T::copy(u_c[2], u_b_regs[2]);
SM75_U32x1_MOVM_T::copy(u_c[3], u_b_regs[3]);

auto tCrB_u_tmp = thr_mma.partition_fragment_B(B_ref);
uint32_t* b_dst = reinterpret_cast<uint32_t*>(&tCrB_u_tmp(0));
b_dst[0] = u_b_regs[0]; b_dst[1] = u_b_regs[1];
b_dst[2] = u_b_regs[2]; b_dst[3] = u_b_regs[3];

clear(u_acc[i]);
gemm(thr_mma, tCrA_k(_,_,Int<0>{}), tCrB_u_tmp(_,_,Int<0>{}), u_acc[i]);

cute::transform(u_acc[i], u_bf16[i], [] __device__ (float x) { return BF16(x); });
}

// ======== Phase 4: Load Mqk, MOVM_T → tCrB_u_arr, Mqk@U + add out ========
copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(Mqk), tCrAi_k_view);
cute::transform(tCrAi_k, tCrA_k, cute::identity{});

BFragT_u tCrB_u_arr[2];

#pragma unroll
for (int i = 0; i < 2; ++i) {
uint32_t* u_c = reinterpret_cast<uint32_t*>(&u_bf16[i](0));
SM75_U32x1_MOVM_T::copy(u_c[0], u_b_regs[0]);
SM75_U32x1_MOVM_T::copy(u_c[1], u_b_regs[1]);
SM75_U32x1_MOVM_T::copy(u_c[2], u_b_regs[2]);
SM75_U32x1_MOVM_T::copy(u_c[3], u_b_regs[3]);

tCrB_u_arr[i] = thr_mma.partition_fragment_B(B_ref);
uint32_t* b_dst = reinterpret_cast<uint32_t*>(&tCrB_u_arr[i](0));
b_dst[0] = u_b_regs[0]; b_dst[1] = u_b_regs[1];
b_dst[2] = u_b_regs[2]; b_dst[3] = u_b_regs[3];

clear(out_acc[i]);
gemm(thr_mma, tCrA_k(_,_,Int<0>{}), tCrB_u_arr[i](_,_,Int<0>{}), out_acc[i]);

SFragT gemm_bf16;
cute::transform(out_acc[i], gemm_bf16, [] __device__ (float x) { return BF16(x); });
cute::transform(out_bf16[i], gemm_bf16, out_bf16[i], [] __device__ (BF16 c, BF16 a) { return c + a; });
}

// ======== Phase 5: Store final out ========
#pragma unroll
for (int i = 0; i < 2; ++i) {
Tensor out_block = local_tile(out_tile, make_shape(Int<16>{}, Int<16>{}), make_coord(0, warp_id * 2 + i));
copy(smem_tiled_store_C, smem_thr_store_C.retile_S(out_bf16[i]), smem_thr_store_C.partition_D(out_block));
}

// ======== Phase 6: s_acc update ========
// s_acc[D, D] = s_acc * g_total + k_restored_t[D, 16] @ U[16, D]
// Each warp handles columns [warp_id*32, (warp_id+1)*32] = 2 x 16x16 blocks
// U is already in tCrB_u_arr[0..1] as B operands (from Phase 4 MOVM_T)
constexpr int S_M_BLOCKS = decltype(cute::size<0>(k_restored_t))::value / 16;

Tensor tCrAi_kr = make_fragment_like<BF16>(thr_mma.partition_fragment_A(A_ref));
auto tCrAi_kr_view = smem_thr_copy_A_T.retile_D(tCrAi_kr);

AFragT ring_A_kr[PREFETCH];
SFragT ring_S_acc[2][PREFETCH];
float ring_g0[PREFETCH], ring_g1[PREFETCH];

#pragma unroll
for (int i = 0; i < PREFETCH; ++i) {
Tensor kr_block = local_tile(k_restored_t, make_shape(Int<16>{}, Int<16>{}), make_coord(i, 0));
copy(smem_tiled_copy_A_T, smem_thr_copy_A_T.partition_S(kr_block), tCrAi_kr_view);
cute::transform(tCrAi_kr, ring_A_kr[i], cute::identity{});

#pragma unroll
for (int bi = 0; bi < 2; ++bi) {
Tensor s_block = local_tile(s_acc_T, make_shape(Int<16>{}, Int<16>{}), make_coord(i, warp_id * 2 + bi));
copy(smem_tiled_load_C_T, smem_thr_load_C_T.partition_S(s_block), smem_thr_load_C_T.retile_D(ring_S_acc[bi][i]));
}

ring_g0[i] = g_total(i * 16 + group_id);
ring_g1[i] = g_total(i * 16 + group_id + 8);
}

#pragma unroll
for (int m = 0; m < S_M_BLOCKS; ++m) {
const int slot = m % PREFETCH;

float g0 = ring_g0[slot];
float g1 = ring_g1[slot];

#pragma unroll
for (int bi = 0; bi < 2; ++bi) {
clear(u_acc[bi]);
gemm(thr_mma, ring_A_kr[slot](_,_,Int<0>{}), tCrB_u_arr[bi](_,_,Int<0>{}), u_acc[bi]);
}

if (m + PREFETCH < S_M_BLOCKS) {
Tensor kr_next = local_tile(k_restored_t, make_shape(Int<16>{}, Int<16>{}), make_coord(m + PREFETCH, 0));
copy(smem_tiled_copy_A_T, smem_thr_copy_A_T.partition_S(kr_next), tCrAi_kr_view);
cute::transform(tCrAi_kr, ring_A_kr[slot], cute::identity{});

ring_g0[slot] = g_total((m + PREFETCH) * 16 + group_id);
ring_g1[slot] = g_total((m + PREFETCH) * 16 + group_id + 8);
}

#pragma unroll
for (int bi = 0; bi < 2; ++bi) {
#pragma unroll
for (int a = 0; a < 2; ++a) {
#pragma unroll
for (int d = 0; d < 2; ++d) {
auto c0 = make_coord(make_coord(a, 0), 0, d);
auto c1 = make_coord(make_coord(a, 1), 0, d);
ring_S_acc[bi][slot](c0) = BF16(bf16_to_f32(ring_S_acc[bi][slot](c0)) * g0 + u_acc[bi](c0));
ring_S_acc[bi][slot](c1) = BF16(bf16_to_f32(ring_S_acc[bi][slot](c1)) * g1 + u_acc[bi](c1));
}
}

Tensor s_block = local_tile(s_acc_T, make_shape(Int<16>{}, Int<16>{}), make_coord(m, warp_id * 2 + bi));
copy(smem_tiled_store_C_T, smem_thr_store_C_T.retile_S(ring_S_acc[bi][slot]), smem_thr_store_C_T.partition_D(s_block));

if (m + PREFETCH < S_M_BLOCKS) {
Tensor s_next = local_tile(s_acc_T, make_shape(Int<16>{}, Int<16>{}), make_coord(m + PREFETCH, warp_id * 2 + bi));
copy(smem_tiled_load_C_T, smem_thr_load_C_T.partition_S(s_next), smem_thr_load_C_T.retile_D(ring_S_acc[bi][slot]));
}
}
}
}
compute_barrier.arrive_and_wait();

#ifndef TMA_DISABLE_ALL
cutlass::arch::fence_view_async_shared();
store_pipeline.producer_commit(out_write);
load_pipeline.consumer_release(load_read);
++load_read;
++out_write;
#endif
}
}

#ifndef TMA_DISABLE_ALL
if (warp_role == WarpRole::STORE && lane_predicate) {
Tensor g_out = tma_store_out.get_tma_tensor(make_shape(H, T_total, D));
auto cta_tma_store = tma_store_out.get_slice(Int<0>{});
StorePipelineState out_read;
for (int t = 0; t < t_tiles; ++t) {
store_pipeline.consumer_wait(out_read);
int stage = out_read.index();
int actual_len = min(CHUNK, seq_len - t * CHUNK);

BF16* out_stage_ptr = shared_storage.output[stage].out.begin();

if (actual_len < CHUNK) {
// Manual store for tail tile to avoid overwriting next sequence
// Only one thread (lane_predicate) runs here, so loop over all D
Tensor s_out = make_tensor(make_smem_ptr(out_stage_ptr), VOLayout{});
for (int row = 0; row < actual_len; ++row) {
int64_t global_base = (bos + t * CHUNK + row) * H * D + head_idx * D;
for (int col = 0; col < D; ++col) {
out_raw_ptr[global_base + col] = s_out(row, col);
}
}
} else {
// TMA store for full tiles
auto out_off = g_out.layout()(head_idx, int(bos) + t * CHUNK, 0);
Tensor g_out_tile = make_tensor(g_out.data() + out_off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<D>{}), stride(g_out.layout())));
Tensor s_out_tile = make_tensor(make_smem_ptr(out_stage_ptr), TMAVOLayout{});
cute::copy(
tma_store_out,
cta_tma_store.partition_S(s_out_tile),
cta_tma_store.partition_D(g_out_tile)
);
tma_store_arrive();
}

tma_store_wait<0>();
store_pipeline.consumer_release(out_read);
++out_read;
}

if constexpr (HasStateOut && !StateFP32) {
// BF16 state: TMA store directly from state_acc
Tensor g_final = tma_store_final_state.get_tma_tensor(make_shape(N * H, D, D));
auto state_off = g_final.layout()(seq_idx * H + head_idx, 0, 0);
Tensor g_final_tile = make_tensor(g_final.data() + state_off,
make_layout(make_shape(Int<1>{}, Int<D>{}, Int<D>{}), stride(g_final.layout())));
Tensor s_state = make_tensor(make_smem_ptr(shared_storage.state_acc.begin()), TMAStateSmemLayout{});

auto cta_tma_store_state = tma_store_final_state.get_slice(Int<0>{});
cute::copy(
tma_store_final_state,
cta_tma_store_state.partition_S(s_state),
cta_tma_store_state.partition_D(g_final_tile)
);
tma_store_arrive();
}
}

if constexpr (HasStateOut && StateFP32) {
// FP32 state: all threads sync, convert bf16->fp32, then STORE warp does TMA
using FP32StateSmemLayout = typename Layouts::FP32StateSmemLayout;
using TMAFP32StateSmemLayout = typename Layouts::TMAFP32StateSmemLayout;

__syncthreads(); // all warps sync — pipeline smem now free

smem_cvt_bf16_to_fp32<StateSmemLayout, FP32StateSmemLayout, D, NumThreads>(
shared_storage.state_acc.begin(),
reinterpret_cast<float*>(shared_storage.state_fp32_buf),
threadIdx.x);
__syncthreads(); // conversion complete

if (warp_role == WarpRole::STORE && lane_predicate) {
Tensor g_final = tma_store_final_state.get_tma_tensor(make_shape(N * H, D, D));
auto state_off = g_final.layout()(seq_idx * H + head_idx, 0, 0);
Tensor g_final_tile = make_tensor(g_final.data() + state_off,
make_layout(make_shape(Int<1>{}, Int<D>{}, Int<D>{}), stride(g_final.layout())));
Tensor s_fp32 = make_tensor(
make_smem_ptr(reinterpret_cast<float*>(shared_storage.state_fp32_buf)),
TMAFP32StateSmemLayout{});

auto cta_tma_store_state = tma_store_final_state.get_slice(Int<0>{});
cute::copy(
tma_store_final_state,
cta_tma_store_state.partition_S(s_fp32),
cta_tma_store_state.partition_D(g_final_tile)
);
tma_store_arrive();
}
}

__syncthreads();
#endif
}

计算流程图

image

image

本文目标:从底层 MMA 指令出发,解释 Phase 1-6 中所有寄存器操作的"为什么"。


1 MMA 寄存器基础知识

1.1 MMA 指令基础

GPU 的矩阵乘法不是一个线程算一个元素,而是 一个 warp (32线程) 协作算一整个矩阵块

SM80+ 的核心指令:mma.sync.aligned.m16n8k16.f32.bf16.bf16.f32

含义:

1
2
3
C[16×8] += A[16×16] × B[16×16]^T
^^^^^^^ ^^^^^^^^ ^^^^^^^^
fp32累加 bf16输入 bf16输入

注意

  • 一条 mma 指令只能算 16×8 的输出
  • 要算 16×16 的输出,需要 两条 mma 指令(列方向拼接:8+8=16)
  • CuTe 的一次 gemm() 调用 = 两条 mma 指令 = 得到 16×16 的 C

1.2 三种格式:A格式、B格式、C格式

MMA 指令要求 A、B、C 三个操作数在 32 个线程的寄存器中以特定方式分布。三种分布方式完全不同!这是理解所有寄存器操作的关键。

1.2.1 C 格式(累加器格式 / 输出格式)

C 是 MMA 的输出。16×16 = 256 个元素,分给 32 线程,每线程 8 个。

每线程持有的 8 个 fp32 元素在矩阵中的位置:

1
2
3
4
5
6
groupID    = lane_id / 4      (0..7)
tid_in_grp = lane_id % 4 (0..3)

行: groupID 和 groupID+8 (只碰这两行)
列: tid_in_grp*2 和 tid_in_grp*2+1 (连续两列)
+ 列偏移 0 (第一条mma) 或 8 (第二条mma)

画出来(以 groupID=1, tid_in_grp=1 即 lane_id=5 为例):

1
2
3
4
5
6
7
8
9
10
11
16×16 矩阵:
第一条mma管的列 第二条mma管的列
col 0-7 col 8-15
row 0: . . . . . . . . . . . . . . . .
row 1: . . ★ ★ . . . . . . ★ ★ . . . . ← groupID=1
row 2: . . . . . . . . . . . . . . . .
...
row 8: . . . . . . . . . . . . . . . .
row 9: . . ★ ★ . . . . . . ★ ★ . . . . ← groupID+8=9

★ = 这个线程(lane5)持有的元素,共8个

特点

  • 每线程跨两行(row 和 row+8)
  • 每线程在一行中持有连续两列
  • 同行同列的 4 个线程(tid_in_grp=0,1,2,3)覆盖了该行的全部 8 列

这就是为什么:

  • beta0 对应 row_half=0 的 4 个元素(groupID 那行)
  • beta1 对应 row_half=1 的 4 个元素(groupID+8 那行)
  • 每行的 beta 值相同,一个线程只碰两行,所以只需要 2 个标量

1.2.2 A 格式(左操作数格式)

A 是 16×16 的左矩阵。每线程也持有 8 个 bf16(打包成 4 个 uint32)。

  • lane_id % 16(0…15,每线程只在一行)
  • :分成4组,由 lane_id / 16 和具体分配决定

关键区别:A 格式中每线程持有同一行的元素(按 K 维度排列),C 格式中每线程持有两行的元素(按输出行排列)。

1.2.3 B 格式(右操作数格式)

B 是 16×16 的右矩阵(实际是转置后使用)。每线程也持有 8 个 bf16,但排列方式又不同于 A 和 C。

每线程持有同一列的元素(按 K 维度排列)。

1.2.4 三种格式的核心区别

格式 每线程持有的几何含义 典型操作
A 格式 同一行的若干元素 LDSM_N 加载,作为 gemm 左操作数
B 格式 同一列的若干元素 LDSM_N 加载,作为 gemm 右操作数
C 格式 两行的若干元素 gemm 输出,逐元素运算,STSM 存储

三种格式是不可互换的!如果你在 C 格式的寄存器中有数据,想把它作为 B 操作数喂给下一个 gemm,你必须做格式转换(C→B)。这就是 MOVM_T 存在的意义。


1.3 从 smem 到寄存器:LDSM 指令

LDSM = Load Shared Memory(ldmatrix PTX 指令)

一条 LDSM 指令:32 个线程协作,从 shared memory 读取一个矩阵块,自动把数据按 A/B/C 格式分配到各线程的寄存器中。

1.3.1 LDSM_N(Normal)

按正常方向加载。

  • copy_A + LDSM_N:把 smem 中 16×16 块加载为 A 格式
  • copy_B + LDSM_N:把 smem 中 16×16 块加载为 B 格式
  • copy_C + LDSM_N:把 smem 中 16×16 块加载为 C 格式

1.3.2 LDSM_T(Transposed)

加载时做转置。

  • copy_A + LDSM_T:把 smem 中 16×16 块转置后加载为 A 格式
  • copy_C + LDSM_T:把 smem 中 16×16 块转置后加载为 C 格式

物理数据不变,但加载后在寄存器中的行列关系变了。

应用:Phase 6 中 s_acc 的物理布局是"列优先"(为 Phase1 的 B 格式设计),但 Phase 6 需要按行读取,所以用 LDSM_T 转置加载。

1.3.3 对应 Phase 中的用法

Phase LDSM 操作
Phase 1 LDSM_N 加载 kd, qd → A 格式(tCrA_k, tCrA_q) LDSM_N 加载 s_acc → B 格式(tCrB)
Phase 2 LDSM_N 加载 v → C 格式(v_bf16) LDSM_N 加载 INV → A 格式(tCrA_k)
Phase 4 LDSM_N 加载 Mqk → A 格式(tCrA_k)
Phase 6 LDSM_T 加载 kr^T → A 格式(ring_A_kr) LDSM_T 加载 s_acc_T → C 格式(ring_S_acc)

1.4、从寄存器到 smem:STSM 指令

STSM = Store Shared Memory(stmatrix PTX 指令)

和 LDSM 相反,32 线程协作把寄存器中的数据写回 smem。

Phase STSM 操作
Phase 5 STSM_N:out_bf16(C 格式)→ smem out_tile
Phase 6 STSM_T:ring_S_acc(C 格式)→ smem s_acc_T(转置写回)

1.5、格式转换:为什么需要、怎么做

1.5.1 问题场景

Phase 3 的情况:

  1. u_acc[i] = k_decayed @ s_acc ← gemm 输出,C 格式(fp32)
  2. u_bf16[i] = bf16(u_acc[i]) ← 类型转换,还是 C 格式
  3. u = (v - u) * beta ← 逐元素,还是 C 格式
  4. 现在要算 INV @ uu 需要作为 B 操作数!!但 u 在 C 格式里!!

C 格式中,lane 5 持有 (row1,col2), (row1,col3), (row9,col2)… B 格式中,lane 5 应该持有同一列的不同行的元素。

完全不同的排列!

1.5.2 传统方案:smem round-trip

1
C 格式寄存器 ──STSM──> shared memory ──LDSM──> B 格式寄存器
  • 两次 smem 访问,每次 ~30 cycle 延迟
  • 还要 __syncthreads() 确保写完再读

1.5.3 MOVM_T:纯寄存器转置

SM75_U32x1_MOVM_T:使用 warp shuffle 在线程间交换寄存器

每个线程把自己的 C 格式寄存器值 shuffle 给需要它的线程,接收方拿到的值自然形成 B 格式。

1
C 格式的 uint32[0..3] ──4次MOVM_T──> B 格式的 uint32[0..3]
  • 零 smem 访问
  • 延迟只有 shuffle 的 ~5 cycle

代码示例:

1
2
3
4
5
uint32_t* u_c = reinterpret_cast<uint32_t*>(&u_bf16[i](0));
SM75_U32x1_MOVM_T::copy(u_c[0], u_b_regs[0]); // shuffle
SM75_U32x1_MOVM_T::copy(u_c[1], u_b_regs[1]); // shuffle
SM75_U32x1_MOVM_T::copy(u_c[2], u_b_regs[2]); // shuffle
SM75_U32x1_MOVM_T::copy(u_c[3], u_b_regs[3]); // shuffle

每条 MOVM_T 处理 1 个 uint32 = 2 个 bf16,4 条 = 8 个 bf16 = 一个线程在 16×16 中持有的全部分量。

1.5.4 Phase 中哪里做了格式转换

Phase 转换操作
Phase 3(两次,每列块一次) u_bf16[i] C格式 ──MOVM_T──> tCrB_u_tmp B格式(为了 INV @ u)
Phase 4(两次,每列块一次) u_bf16[i] C格式 ──MOVM_T──> tCrB_u_arr[i] B格式(为了 Mqk @ U 和 Phase6 的 kr^T @ U)

1.6、load格式 vs 计算格式:tCrAi 和 tCrA 的区别

代码中每个 fragment 都有两份:

  • tCrAi_k(load 版)和 tCrA_k(计算版)
  • tCrAi_q(load 版)和 tCrA_q(计算版)
  • tCrBi(load 版)和 tCrB(计算版)

原因

  • LDSM 指令要求目标寄存器按特定排列(load layout)
  • MMA 指令要求操作数寄存器按另一种排列(compute layout)

这两种 layout 在 SM90 上实际是相同的(底层寄存器完全一致),但 CuTe 的类型系统不知道这一点,它把两种 layout 作为不同类型处理。

所以:

1
2
3
copy(smem → tCrAi_k)                 // LDSM 写入 load 格式
transform(tCrAi_k, tCrA_k, identity) // "转换"为计算格式(实际零开销)
gemm(tCrA_k, tCrB, acc) // MMA 使用计算格式

transform + identity 是编译期的类型仪式,不生成实际指令。两个变量名指向的是逻辑上同一组寄存器(或编译器内联后等价)。

为什么不合并成一个?

CuTe 的 copygemm 要求不同的类型签名。如果直接用 tCrA_k 做 LDSM 目标,类型不匹配会编译报错。tCrAi_k + retile_D 视图让 copy 指令能写入,然后 transform 让 gemm 指令能读取。这是模板元编程的代价。


1.7、累加器(AccFragT)的复用

u_acc[2]out_acc[2] 是 fp32 累加器,每线程持有 2×4 = 8 个 fp32。

为什么 fp32 而不是 bf16?

  • MMA 指令的 C 操作数是 fp32
  • 多步累加(Phase 1 的 8 步 K 循环)需要高精度防止误差积累
  • 只在需要时(Phase 2/3 结尾)才转成 bf16

复用逻辑

u_acc 在 3 个 Phase 中承担不同角色:

Phase u_acc 的角色
Phase 1 u_acc[i] += kd[:, k] @ s_acc[k, warp列i](8步累加,含义:k@s) ↓ Phase 2:值被提取到 u_bf16,u_acc 不再需要
Phase 3 clear(u_acc[i]) gemm(INV, u_B, u_acc[i])(1步,含义:INV@u = U) ↓ 值被提取到 u_bf16,u_acc 不再需要
Phase 6 clear(u_acc[bi]) gemm(kr_A, U_B, u_acc[bi])(每步m,含义:kr^T@U) ↓ 值被用于逐元素更新 s_acc

每次复用前都 clear() 清零。寄存器是宝贵资源,能复用就复用。

out_acc 在 2 个 Phase 中复用:

Phase out_acc 的角色
Phase 1 out_acc[i] += qd[:, k] @ s_acc[k, warp列i](8步累加,含义:q@s) ↓ Phase 2:值被提取到 out_bf16,out_acc 不再需要
Phase 4 clear(out_acc[i]) gemm(Mqk, U_B, out_acc[i])(1步,含义:Mqk@U) ↓ 值被提取到 gemm_bf16,加到 out_bf16 上

1.8、bf16 fragment(SFragT)的角色

SFragT = bf16 类型的 C 格式 fragment。每线程 8 个 bf16。

变量 用途
u_bf16[2] 承载逐元素运算的中间结果。 Phase 3:先存 bf16(k@s),再变成 (v-k@s)*β,再变成 bf16(INV@u) = U Phase 4:U 被 MOVM_T 转走后就不需要了
out_bf16[2] 最终输出的载体。 Phase 2:存 bf16(q@s) Phase 4:+= bf16(Mqk@U),变成最终 output Phase 5:STSM 写到 smem
v_bf16[2] 从 smem 加载的 v 值,Phase 3 做完 v-k@s 后就不需要了

为什么用 bf16 而不保持 fp32?

  1. 逐元素运算 (v-u)*β 对精度要求不高,bf16 够用
  2. 后面要做 MOVM_T 转成 B 格式,MMA 的 B 操作数只接受 bf16
  3. bf16 寄存器只占 fp32 的一半空间,省寄存器

1.9、Phase 1-6 寄存器使用的完整故事

一个 warp 的视角,追踪每个寄存器变量:

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
┌─ Phase 1 ──────────────────────────────────────────────────┐
│ │
│ [A格式] tCrA_k ← LDSM(kd[:, k]) 每步覆盖 │
│ [A格式] tCrA_q ← LDSM(qd[:, k]) 每步覆盖 │
│ [B格式] tCrB ← LDSM(s_acc[k, warp列]) 每步覆盖 │
│ [C格式] u_acc[0,1] += tCrA_k @ tCrB 8步累加 │
│ [C格式] out_acc[0,1] += tCrA_q @ tCrB 8步累加 │
│ │
│ tCrA_q, tCrB → Phase1 后不再需要,寄存器自然释放 │
│ tCrA_k → 继续在 Phase 2,3,4 作为 A 操作数容器 │
│ │
└────────────────────────────────────────────────────────────┘
│ u_acc 里是 k@s (fp32)
│ out_acc 里是 q@s (fp32)
v
┌─ Phase 2 ──────────────────────────────────────────────────┐
│ │
│ [C格式] u_bf16[0,1] = bf16(u_acc[0,1]) fp32→bf16 │
│ [C格式] out_bf16[0,1] = bf16(out_acc[0,1]) fp32→bf16 │
│ [C格式] v_bf16[0,1] ← LDSM_N(v_tile) 新加载 │
│ [A格式] tCrA_k ← LDSM_N(INV) 覆盖kd │
│ [标量] beta0, beta1 ← smem 标量读 │
│ │
│ u_acc, out_acc 的值已被提取,但寄存器还在 (Phase3,4复用) │
│ │
└────────────────────────────────────────────────────────────┘

v
┌─ Phase 3 (对 i=0,1) ──────────────────────────────────────┐
│ │
│ 逐元素 (全部 C 格式,每线程独立): │
│ [C格式] u_bf16[i] = (v_bf16[i] - u_bf16[i]) * beta0/1 │
│ ↑ v_bf16 用完释放,beta 用完释放 │
│ │
│ 格式转换: │
│ [C格式] u_bf16[i] ──MOVM_T──> [B格式] tCrB_u_tmp │
│ ↑ warp 内 32 线程 shuffle 寄存器 │
│ │
│ GEMM: │
│ clear(u_acc[i]) │
│ [A格式] tCrA_k (=INV) × [B格式] tCrB_u_tmp (=u) │
│ → [C格式] u_acc[i] = INV @ u │
│ │
│ [C格式] u_bf16[i] = bf16(u_acc[i]) 结果就是 U │
│ ↑ tCrB_u_tmp 释放 │
│ │
└────────────────────────────────────────────────────────────┘
│ u_bf16 里是 U (C 格式)
v
┌─ Phase 4 (对 i=0,1) ──────────────────────────────────────┐
│ │
│ [A格式] tCrA_k ← LDSM_N(Mqk) 覆盖 INV │
│ │
│ 格式转换: │
│ [C格式] u_bf16[i] ──MOVM_T──> [B格式] tCrB_u_arr[i] │
│ ↑ u_bf16 释放 │
│ ↑ tCrB_u_arr 保留到 Phase 6 !! │
│ │
│ GEMM: │
│ clear(out_acc[i]) │
│ [A格式] tCrA_k (=Mqk) × [B格式] tCrB_u_arr[i] (=U) │
│ → [C格式] out_acc[i] = Mqk @ U │
│ │
│ 累加到输出: │
│ [C格式] gemm_bf16 = bf16(out_acc[i]) │
│ [C格式] out_bf16[i] += gemm_bf16 │
│ = q@s + Mqk@U = 最终输出 │
│ ↑ out_acc 释放,tCrA_k 释放 │
│ │
└────────────────────────────────────────────────────────────┘
│ out_bf16 里是最终输出 (C 格式)
│ tCrB_u_arr 里是 U (B 格式,保留!)
v
┌─ Phase 5 ──────────────────────────────────────────────────┐
│ │
│ [C格式] out_bf16[i] ──STSM_N──> smem out_tile │
│ ↑ out_bf16 释放 │
│ │
└────────────────────────────────────────────────────────────┘

v
┌─ Phase 6 (m=0..7) ────────────────────────────────────────┐
│ │
│ [A格式] ring_A_kr ← LDSM_T(kr^T[m块]) 每步覆盖 │
│ [C格式] ring_S_acc[bi] ← LDSM_T(s_acc_T) 每步覆盖 │
│ [标量] ring_g0, g1 ← smem 标量读 每步覆盖 │
│ │
│ GEMM: │
│ clear(u_acc[bi]) │
│ [A格式] ring_A_kr (=kr^T[m]) × [B格式] tCrB_u_arr[bi] │
│ → [C格式] u_acc[bi] = kr^T @ U │
│ ↑ tCrB_u_arr 在这里│
│ 第三次被使用! │
│ │
│ 逐元素更新 (C 格式): │
│ ring_S_acc[bi] = bf16(f32(ring_S_acc[bi]) * g + u_acc[bi])│
│ │
│ [C格式] ring_S_acc[bi] ──STSM_T──> smem s_acc_T │
│ [C格式] ring_S_acc[bi] ← LDSM_T(s_acc_T[m+1]) 预取 │
│ │
└────────────────────────────────────────────────────────────┘

1.10、逐元素运算为什么必须在 C 格式下做

Phase 3:u = (v - k@s) * beta Phase 4:out = out_bf16 + gemm_bf16 Phase 6:s_new = s_old * g + gemm_result

这些逐元素运算都在 C 格式下进行。原因

  1. gemm 输出本来就是 C 格式,不需要额外转换
  2. 同一行的 beta/g 值相同,C 格式中同行元素在同一线程 → 只需一个标量
  3. 逐元素运算每线程独立,无需跨线程通信

如果在 A 或 B 格式下做:

  • 同一行的元素可能分散在不同线程 → 需要 shuffle 传 beta 值
  • 格式转换本身就有开销
  • 完全没必要

所以设计原则是:逐元素运算在 C 格式做,需要作为 operand 时再 MOVM_T 转格式


1.11、tCrA_k 的一生:理解寄存器名复用

tCrA_k 这个变量在 Phase 1→2→3→4 中被反复覆盖:

Phase tCrA_k 的内容 操作
Phase 1 循环中 kd[:, k*16:(k+1)*16] 每步覆盖,存 k_decayed 的列块
Phase 2 INV[16x16] 覆盖,存 INV
Phase 3 被 gemm 读取 使用(INV @ u)
Phase 4 Mqk[16x16] 覆盖,存 Mqk
Phase 4 被 gemm 读取 使用(Mqk @ U)

它始终是 A 格式,始终通过 LDSM_N 加载(覆盖旧值)。只是内容从 kd 变成 INV 再变成 Mqk。

为什么能这样?

因为 kd、INV、Mqk 的生命期不重叠。Phase 1 用完 kd 后再也不需要它,Phase 2 就可以覆盖成 INV。Phase 3 用完 INV 后,Phase 4 覆盖成 Mqk。


1.11、为什么 Phase 3 的 MOVM_T 和 Phase 4 的 MOVM_T 是两次

Phase 3

  • u_bf16(C格式)──MOVM_T──> tCrB_u_tmp(B格式,临时)
  • 用途:INV @ u(u 作为 B 操作数)
  • 结果:u_acc = INV @ uu_bf16 = bf16(u_acc) ← u_bf16 被新值覆盖了!

Phase 4

  • u_bf16(C格式,新值=U)──MOVM_T──> tCrB_u_arr[i](B格式,持久)
  • 用途:Mqk @ U(U 作为 B 操作数)
  • 额外:Phase 6 还要用 tCrB_u_arrkr^T @ U

为什么不能一次 MOVM_T 搞定?

因为 Phase 3 的 MOVM_T 转的是 u(还没乘 INV 之前),Phase 4 的 MOVM_T 转的是 U = INV @ u(乘完之后)。它们是不同的矩阵!

1
2
3
Phase 3: u → B格式 → INV@u → 得到 U (C格式)

Phase 4: U → B格式 → ...

中间必须经过一次 gemm(INV@u),gemm 的输出必然是 C 格式,所以必须再做一次 MOVM_T 把新的 C 格式转成 B 格式。


1.13、总结:格式转换路径图

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
smem kd  ──LDSM_N──> [A格式] ──┐
├─ gemm ──> [C格式] u_acc (k@s)
smem s_acc ──LDSM_N──> [B格式] ──┘ │
bf16 转换

smem v ──LDSM_N──> [C格式] v_bf16 ──┐ v
├─ 逐元素 ──> [C格式] u_bf16
smem beta ──标量读──> beta ──────────┘ │
MOVM_T (shuffle)

smem INV ──LDSM_N──> [A格式] tCrA ──┐ v
├─ gemm ──> [C格式] u_acc (U=INV@u)
[B格式] tCrB_u_tmp ──┘ │
bf16 转换

smem Mqk ──LDSM_N──> [A格式] tCrA ──┐ v
├─ gemm ──> [C格式] out_acc (Mqk@U)
[B格式] tCrB_u_arr ─┬──┘
(MOVM_T from U) │
│ (寄存器保留!)

smem kr^T ──LDSM_T──> [A格式] ring_A ──┐ │
├─ gemm ──> [C格式] u_acc (kr^T@U)
[B格式] tCrB_u_arr ┘ │
逐元素 + s_old

[C格式] ring_S_acc

STSM_T

v
smem s_acc (更新后)

2 K2 kernel 逐行讲解 Part 1: Layout + Shared Memory + 初始化**

文件: FlashKDA/csrc/smxx/fwd_kernel2.cuh

行 9-70: Layout 定义

1
D = 128, CHUNK = 16

两个关键 Atom

  • Layout_K_INTER_Atom**<bf16>**: 8行 x 16列 的 bf16 swizzle layout

  • 用于 MMA 的 A/B 操作数 (按 K 维度内部交错)

  • 避免 shared memory bank conflict

  • Layout_MN_INTER_Atom**<bf16>**: 转置版,16行 x 8列

  • 用于读写转置方向的数据

tile_to_shape

把小 atom 重复铺满整个 shape:

  • 例: tile_to_shape(8x16 atom, (16, 128)) = 把 8x16 块铺成 16x128

  • 行方向重复 16/8 = 2

  • 列方向重复 128/16 = 8

各 Layout 含义

Layout 尺寸 用途
MMALayout [16 x 128] k_decayed, q_decayed, k_restored, v, out 的 smem layout
TransposedMMALayout [128 x 16] k_restored^T 的 smem layout (转置视图)
StateSmemLayout [128 x 128] s_acc 的 smem layout (按列读取友好)
TransposedStateSmemLayout [128 x 128] s_acc 的转置视图 (按行读取友好)

StateSmemLayoutTransposedStateSmemLayout 指向同一片 smem,物理数据相同,只是访问模式不同。

Layout 尺寸 用途
BetaSmemLayout [32] beta 的 smem layout (1D,连续)
GTotalLayout [128] g_total 的 smem layout (1D,连续)
LMLayout [16 x 16] INV, Mqk 的 smem layout
TMAVOLayout / TMAStateSmemLayout / TMALMLayout [1, …] TMA descriptor 需要 3D tensor,第一维是 batch/tile 索引
FP32StateSmemLayout [128 x 128] fp32 版本的 state layout (fp32 state 输入输出时的临时 buffer)

为什么需要 Swizzle?

  • Shared memory 有 32 个 bank。如果多线程访问同一 bank,会串行化 (bank conflict)
  • Swizzle:对地址做 XOR 运算,把原本落在同一 bank 的访问分散到不同 bank
  • Layout_K_INTER_Atom 内部:8行x16列的 bf16 数据,用 swizzle 保证 MMA 的 LDSM 指令 32 线程同时读取时不会有 bank conflict

理解要点:不需要理解 swizzle 的具体 XOR 公式,只需要知道:用了 swizzle layout 的 smem,MMA 读取时不会有 bank conflict。


行 72-112: Shared Memory 结构体

1
2
template <class Layouts, int InputStages, int OutputStages>
struct SharedStorageK2 {
  • InputStages = 3 (输入流水线3级)
  • OutputStages = 2 (输出流水线2级)

state_acc — 隐状态

1
alignas(128) cute::ArrayEngine<BF16, cosize_v<StateSmemLayout>> state_acc;
  • 大小 = 128 * 128 * 2 = 32768 字节 = 32KB
  • alignas(128): 对齐到 128 字节 (TMA 要求)
  • 这是 K2 里最大的一块 smem,常驻不动,chunk 之间保留

InputStorage — K1 的输出

1
2
3
4
5
6
7
8
9
10
struct InputStorage {
alignas(128) ...v; // 16x128 bf16 = 4KB
alignas(128) ...beta; // 32 bf16 = 64B
alignas(128) ...k_decayed; // 16x128 bf16 = 4KB
alignas(128) ...q_decayed; // 16x128 bf16 = 4KB
alignas(128) ...k_restored; // 16x128 bf16 = 4KB
alignas(128) ...g_total; // 128 fp32 = 512B
alignas(128) ...INV; // 16x16 bf16 = 512B
alignas(128) ...Mqk; // 16x16 bf16 = 512B
};
  • 一个 InputStorage17.5KB (含 alignas padding)
  • 这些都是 K1 的输出 + v + beta

OutputStorage

1
2
3
struct OutputStorage {
alignas(128) ...out; // 16x128 bf16 = 4KB
};

Union: 复用内存

1
2
3
4
5
6
7
8
union {
struct {
InputStorage input[InputStages]; // input[0], input[1], input[2]
OutputStorage output[OutputStages]; // output[0], output[1]
};
alignas(128) char state_fp32_buf[cosize_v<StateSmemLayout> * sizeof(float)];
// 128*128*4 = 64KB,用于 fp32 state 转换
};

为什么能共用?

  • fp32 state 转换只发生在 主循环开始前 (加载 initial state) 和 主循环结束后 (存 final state)
  • pipeline buffer 只在 主循环中 使用
  • 两者不会同时需要,所以可以 union

Pipeline Barriers

1
2
3
typename cutlass::PipelineTmaAsync<InputStages>::SharedStorage load_pipeline;
typename cutlass::PipelineAsync<OutputStages>::SharedStorage store_pipeline;
alignas(16) cutlass::arch::ClusterTransactionBarrier state_acc_tma_barrier;
  • pipeline 的 barrier 存在 smem 中 (CUTLASS 要求)
  • state_acc_tma_barrier: 专门给 initial state 的 TMA 用的 barrier

Shared Memory 布局图

1
2
3
4
5
6
7
地址: 0                                                              ~82KB
+----------+---------------------------------------------------+----------+
| state_acc| union { }| barriers |
| [128x128]| input[0] | input[1] | input[2] | out[0] | out[1]| |
| 32KB | ~17KB | ~17KB | ~17KB | 4KB | 4KB | |
| | 或者: state_fp32_buf [128x128] fp32 = 64KB | |
+----------+---------------------------------------------------+----------+

行 114-151: Kernel 函数签名

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
__global__ void __launch_bounds__(NumThreads)
_flash_kda_fwd_recurrence(
CUTE_GRID_CONSTANT TmaLoadV const tma_load_v, // v 的 TMA descriptor
CUTE_GRID_CONSTANT TmaLoadBeta const tma_load_beta, // beta 的 TMA descriptor
CUTE_GRID_CONSTANT TmaLoadWsKD const tma_load_ws_kd, // k_decayed
CUTE_GRID_CONSTANT TmaLoadWsQD const tma_load_ws_qd, // q_decayed
CUTE_GRID_CONSTANT TmaLoadWsKR const tma_load_ws_kr, // k_restored
CUTE_GRID_CONSTANT TmaLoadWsGT const tma_load_ws_gt, // g_total
CUTE_GRID_CONSTANT TmaLoadWsINV const tma_load_ws_inv, // INV
CUTE_GRID_CONSTANT TmaLoadWsMqk const tma_load_ws_mqk, // Mqk
CUTE_GRID_CONSTANT TmaLoadState const tma_load_initial_state, // initial state
CUTE_GRID_CONSTANT TmaStoreState const tma_store_final_state, // final state
CUTE_GRID_CONSTANT TmaStoreOut const tma_store_out, // output
cutlass::bfloat16_t* out_raw_ptr, // output 的原始指针 (尾部 chunk 手动写用)
int T_total, // 总 token 数
int H, // head 数
int N, // 序列数
int64_t const* cu_seqlens, // 变长模式的累积长度
int total_tiles // 所有序列的 tile 总数
)

CUTE_GRID_CONSTANT

  • 告诉编译器这些 TMA descriptor 存在 constant memory
  • GPU 所有线程共享同一份,只读,有缓存

__launch_bounds__(NumThreads)

  • 告诉编译器这个 kernel 最多用 NumThreads 个线程
  • 编译器据此优化寄存器分配

TMA Descriptor 统计

  • 8 个 load + 1 个 load state + 1 个 store state + 1 个 store out = 11 个
  • 这些都在 host 端 (fwd_launch.cu) 创建好,通过 kernel 参数传入

行 152-197: 类型定义 + Warp 角色分配

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
using BF16 = cutlass::bfloat16_t;
using FP16 = cutlass::half_t;
using Layouts = K2Layouts<D, CHUNK>;
// ... 一堆 using 别名,把上面定义的 layout 拿过来用

constexpr int kWarpSize = 32;
constexpr int kComputeThreads = 128;

int warp_id = threadIdx.x / kWarpSize; // 0,1,2,3,4,5
WarpRole warp_role = WarpRole::NonParticipant;

if (warp_id < kComputeThreads / kWarpSize) { // warp_id < 4
warp_role = WarpRole::MMA;
} else if (warp_id < kComputeThreads / kWarpSize + 1) { // warp_id == 4
warp_role = WarpRole::LOAD_QKG;
} else if (warp_id < kComputeThreads / kWarpSize + 2) { // warp_id == 5
warp_role = WarpRole::STORE;
}

Warp 角色分配表

threadIdx.x 0-31 32-63 64-95 96-127 128-159 160-191
warp_id 0 1 2 3 4 5
role MMA MMA MMA MMA LOAD STORE
  • 4 个 warp (128 线程) 负责 MMA 计算
  • 1 个 warp (32 线程) 负责从 workspace 加载数据
  • 1 个 warp (32 线程) 负责存储输出

行 172-213: Transaction Bytes + Pipeline 创建

Transaction Bytes 计算

1
2
3
4
5
6
7
8
constexpr uint32_t kTmaTransactionBytes =
cosize_v<VOLayout> * sizeof(BF16) + // v: 16*128*2 = 4096
32 * sizeof(BF16) + // beta: 32*2 = 64
cosize_v<MMALayout> * sizeof(BF16) * 3 + // kd,qd,kr: 4096*3 = 12288
cosize_v<GTotalLayout> * sizeof(float) + // g_total: 128*4 = 512
cosize_v<LMLayout> * sizeof(BF16) * 2 + // INV,Mqk: 512*2 = 1024
0u;
// 总计 = 17984 字节
  • 这个数字告诉 pipeline barrier:一次要搬这么多字节
  • 当 TMA 搬完这么多字节后,barrier 会自动 “满足”,通知 consumer

Shared Memory 初始化

1
2
3
extern __shared__ __align__(128) unsigned char shared_mem[];
using SharedStorageT = SharedStorageK2<Layouts, InputStages, OutputStages>;
SharedStorageT& shared_storage = *reinterpret_cast<SharedStorageT*>(shared_mem);
  • extern __shared__: kernel launch 时由 host 指定大小的动态 shared memory
  • reinterpret_cast: 把 raw byte 指针当作 SharedStorageK2 结构体用

创建 Load Pipeline

1
2
3
4
5
6
7
LoadPipeline load_pipeline = make_load_pipeline<InputStages>(
shared_storage.load_pipeline, // barrier 存在 smem 中
kTmaTransactionBytes, // 每次搬运的字节数
warp_role, // 当前 warp 的角色
1, // producer 线程数 (LOAD warp 只用 1 个线程)
kComputeThreads // consumer 线程数 (128 个 MMA 线程)
);

make_load_pipeline 做了什么:

  1. 初始化 3 个 barrier (在 smem 中),每个对应一个 stage
  2. 设置每个 barrier 的 expected transaction bytes
  3. 根据 warp_role 决定这个线程是 producer 还是 consumer
角色 可调用的方法
producer (LOAD warp) acquire / commit / tail
consumer (MMA warps) wait / release

创建 Store Pipeline

1
2
3
4
5
6
StorePipeline store_pipeline = make_store_pipeline<OutputStages>(
shared_storage.store_pipeline,
warp_role,
kComputeThreads, // producer = MMA warps (128线程)
1 // consumer = STORE warp (1线程)
);

行 215-237: 序列信息

1
2
3
4
int seq_idx  = blockIdx.x;    // grid 的 x 维: 第几个序列
int head_idx = blockIdx.y; // grid 的 y 维: 第几个 head
int64_t bos, eos;
int tile_base;

变长模式 vs 固定长度模式

变长模式 (IsVarlen = true)

1
2
3
4
5
6
7
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;
}
// tile_base = 之前所有序列的 tile 总数

固定长度模式 (IsVarlen = false)

1
2
3
4
int T_seq = T_total / N;             // 每个序列长度相同
bos = seq_idx * T_seq;
eos = bos + T_seq;
tile_base = seq_idx * ((T_seq + CHUNK - 1) / CHUNK);

辅助变量

1
2
3
int seq_len  = int(eos - bos);
int t_tiles = (seq_len + CHUNK - 1) / CHUNK; // 这个序列有几个 chunk
bool lane_predicate = cute::elect_one_sync(); // warp 内选一个线程 (通常是 lane 0)
  • elect_one_sync(): warp 内的 32 个线程投票,只有一个返回 true
  • 用于 TMA:只需要一个线程发指令

tile_base 的作用

用于索引 workspace:K1 的输出按 (head, tile) 排列

1
ws_idx = head_idx * total_tiles + tile_base + t

示例: 3个序列,长度分别 48, 32, 64,CHUNK=16

序列 bos eos tile_base t_tiles
0 0 48 0 3
1 48 80 3 2
2 80 144 5 4

行 240-313: 加载 initial_state

情况 1: bf16 state (HasStateIn && !StateFP32)

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
if constexpr (HasStateIn && !StateFP32) {
if (warp_role == WarpRole::LOAD_QKG && lane_predicate) {
// 只有 LOAD warp 的 lane 0 进入

// 1. 初始化 barrier
shared_storage.state_acc_tma_barrier.init(1);

// 2. 设置期望的 transaction bytes
shared_storage.state_acc_tma_barrier.arrive_and_expect_tx(kStateTransactionBytes);
// kStateTransactionBytes = 128*128*2 = 32768

// 3. 创建 global tensor 视图
Tensor g_init = tma_load_initial_state.get_tma_tensor(
make_shape(N * H, D, D));

auto init_off = g_init.layout()(seq_idx * H + head_idx, 0, 0);

Tensor g_init_tile = make_tensor(g_init.data() + init_off,
make_layout(make_shape(Int<1>{}, Int<D>{}, Int<D>{}),
stride(g_init.layout())));

// 4. 创建 smem tensor 视图
Tensor s_state = make_tensor(
make_smem_ptr(shared_storage.state_acc.begin()),
TMAStateSmemLayout{});

// 5. 发起 TMA 拷贝
auto cta_tma_load_state = tma_load_initial_state.get_slice(Int<0>{});
cute::copy(
tma_load_initial_state.with(
reinterpret_cast<BarrierType&>(shared_storage.state_acc_tma_barrier)),
cta_tma_load_state.partition_S(g_init_tile),
cta_tma_load_state.partition_D(s_state)
);
}

__syncthreads(); // 确保 barrier 已初始化
shared_storage.state_acc_tma_barrier.wait(0); // 等待 TMA 完成
cutlass::arch::fence_view_async_shared(); // 内存栅栏
}

数据流图

1
2
3
4
5
global memory: initial_state[N*H][128][128] bf16
|
| TMA (一个线程发起,硬件异步搬运)

shared memory: state_acc[128][128] bf16 = 32KB

情况 2: fp32 state (HasStateIn && StateFP32)

和 bf16 类似,但多了转换步骤:

1
2
3
4
5
6
7
8
// 1. TMA 加载到 state_fp32_buf (64KB,union 空间)
// 2. 全 192 线程并行转换 fp32 → bf16

smem_cvt_fp32_to_bf16<FP32StateSmemLayout, StateSmemLayout, D, NumThreads>(
reinterpret_cast<float*>(shared_storage.state_fp32_buf),
shared_storage.state_acc.begin(),
threadIdx.x);
__syncthreads();

数据流图

1
2
3
4
5
6
7
8
9
global memory: initial_state[N*H][128][128] fp32 = 64KB
|
| TMA → state_fp32_buf (64KB,union空间)

state_fp32_buf [128x128] fp32
|
| 全192线程并行转换 fp32→bf16

state_acc [128x128] bf16 = 32KB

情况 3: 无 initial state

1
2
3
4
5
6
BF16* buf = shared_storage.state_acc.begin();
constexpr int kTotal = cosize_v<StateSmemLayout>; // 128*128 = 16384
for (int i = threadIdx.x; i < kTotal; i += NumThreads) {
buf[i] = BF16(0);
}
__syncthreads();

192 线程并行清零

线程 写入位置
线程 0 buf[0], buf[192], buf[384], …
线程 1 buf[1], buf[193], buf[385], …
  • 每线程写 16384/192 ≈ 85 个 bf16 零值

3 K2 kernel 逐行讲解 Part 2: LOAD warp + MMA 准备 + Phase 1

行 316-418: LOAD warp 主循环

1
2
3
4
__syncthreads();

if (warp_role == WarpRole::LOAD_QKG && lane_predicate) {
// 只有 warp4 的 lane0 进入

获取全局内存视图

1
2
3
4
5
6
7
8
9
10
11
12
Tensor g_v = tma_load_v.get_tma_tensor(make_shape(H, T_total, D));
// v 在全局内存的视图: [H, T_total, 128]
// 注意这里是 [H, T, D] 不是 [B, T, H, D]
// 因为 host 端已经 reshape 过了

Tensor g_beta = tma_load_beta.get_tma_tensor(make_shape(H * T_total));
// beta 在全局内存: [H * T_total], 1D (host 端已转置)

auto g_ws_kd = tma_load_ws_kd.get_tma_tensor(
make_shape(H * total_tiles, CHUNK, D));
// workspace 中的 k_decayed: [H*total_tiles, 16, 128]
// ... 其他 workspace tensor 类似 ...

初始化 Pipeline 状态

1
2
3
4
5
6
LoadPipelineState load_write = cutlass::make_producer_start_state<LoadPipeline>();
// load_write.index() = 0, 从 stage 0 开始

auto cta_tma_load_v = tma_load_v.get_slice(Int<0>{});
// TMA 的 slice: CTA 级别的 partition (TMA 是 CTA 级操作, 不分 thread)
// ... 其他 TMA slice 类似 ...

主循环结构

1
2
3
4
5
6
7
8
9
10
11
12
13
for (int t = 0; t < t_tiles; ++t) {
load_pipeline.producer_acquire(load_write);
// 等待: 这个 stage 的 smem buffer 是否可用?
// 如果 MMA warp 还在用上一轮的数据, 就等着

using LoadBarrierType = typename LoadPipeline::ProducerBarrierType;
LoadBarrierType* tma_barrier = load_pipeline.producer_get_barrier(load_write);
// 拿到这个 stage 对应的 barrier 指针
// 所有 TMA copy 都绑到这个 barrier 上

int stage = load_write.index(); // 0, 1, 2 循环
int ws_idx = head_idx * total_tiles + tile_base + t;
// workspace 中的全局索引

加载 v

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
auto v_off = g_v.layout()(head_idx, int(bos) + t * CHUNK, 0);
// v 在全局内存中的偏移: 第 head_idx 个 head, 第 (bos + t*16) 行

Tensor g_v_tile = make_tensor(g_v.data() + v_off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<D>{}),
stride(g_v.layout())));
// 切出 [1, 16, 128] 的子 tensor

Tensor s_v_tile = make_tensor(
make_smem_ptr(shared_storage.input[stage].v.begin()),
TMAVOLayout{});
// smem 端: input[stage].v

cute::copy(tma_load_v.with(*tma_barrier),
cta_tma_load_v.partition_S(g_v_tile),
cta_tma_load_v.partition_D(s_v_tile));
// 发起 TMA: global v -> smem input[stage].v
// 绑定到 tma_barrier, 搬完后 barrier 计数增加

加载 beta

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
int beta_linear = head_idx * T_total + (int(bos) + t * CHUNK);
int beta_aligned = beta_linear & ~7;
// TMA 要求地址对齐, &~7 = 向下对齐到 8 的倍数

auto beta_off = g_beta.layout()(beta_aligned);
Tensor g_beta_tile = make_tensor(g_beta.data() + beta_off,
BetaSmemLayout{});
// 从对齐地址开始取 32 个 bf16

Tensor s_beta_tile = make_tensor(
make_smem_ptr(shared_storage.input[stage].beta.begin()),
TMABetaSmemLayout{});

cute::copy(tma_load_beta.with(*tma_barrier),
cta_tma_load_beta.partition_S(g_beta_tile),
cta_tma_load_beta.partition_D(s_beta_tile));

为什么 load 32 个 beta 而不是 16 个?

  • TMA 对 1D tensor 有对齐要求,beta 可能不在 8 的倍数地址上
  • 向下对齐后多 load 一些,MMA warp 用 beta_smem_offset 跳过前面的

示例

1
2
beta_linear = 35, beta_aligned = 32
load beta[32..63],但只用 beta[35..50] (即 offset=3)

加载 Workspace (以 k_decayed 为例)

1
2
3
4
5
6
7
8
9
10
11
12
{
auto off = g_ws_kd.layout()(ws_idx, 0, 0);
Tensor g_tile = make_tensor(g_ws_kd.data() + off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<D>{}),
stride(g_ws_kd.layout())));
Tensor s_tile = make_tensor(
make_smem_ptr(shared_storage.input[stage].k_decayed.begin()),
TMAVOLayout{});
cute::copy(tma_load_ws_kd.with(*tma_barrier),
cta_ws_kd.partition_S(g_tile),
cta_ws_kd.partition_D(s_tile));
}

:q_decayed, k_restored, g_total, INV, Mqk 的加载代码结构完全相同,只是换了不同的 TMA descriptor 和 smem 目标。

Stage 推进

1
2
3
4
5
6
    ++load_write;
// stage 推进: 0->1->2->0->1->2->...
}
load_pipeline.producer_tail(load_write);
// 通知 pipeline: 没有更多数据了
// MMA warp 读完最后一个 stage 后就不会再 wait

每次循环加载的数据量

数据 大小
v 16×128×2 = 4096 B
beta 32×2 = 64 B
k_decayed 16×128×2 = 4096 B
q_decayed 16×128×2 = 4096 B
k_restored 16×128×2 = 4096 B
g_total 128×4 = 512 B
INV 16×16×2 = 512 B
Mqk 16×16×2 = 512 B
总计 17984 字节 ≈ 17.6 KB/chunk

LOAD warp 到此结束,不参与后续计算。


行 422-428: MMA warp 主循环开始

1
2
3
4
5
6
7
8
9
10
11
12
13
14
if (warp_role == WarpRole::MMA) {
// 线程 0-127 进入

cutlass::arch::NamedBarrier compute_barrier(kComputeThreads, 0);
// 128 线程的 barrier, 用于 Phase 6 结束后同步
// 保证 4 个 warp 都写完 s_acc 再进入下一个 chunk

LoadPipelineState load_read;
// load_read.index() = 0, 从 stage 0 开始消费

StorePipelineState out_write = cutlass::make_producer_start_state<StorePipeline>();
// out_write.index() = 0

int compute_tid = threadIdx.x; // 0-127

主循环

1
2
3
4
5
6
7
8
9
for (int t = 0; t < t_tiles; ++t) {
store_pipeline.producer_acquire(out_write);
// 等 output stage 可用 (STORE warp 用完了)

load_pipeline.consumer_wait(load_read);
// 等 input stage 数据就绪 (LOAD warp 搬完了)

int load_stage = load_read.index(); // 0, 1, 2
int out_stage = out_write.index(); // 0, 1

此时:smem 中的 input[load_stage] 已经有了当前 chunk 的所有数据。

时序

  • LOAD warp 已经搬完 chunk t 的数据到 input[load_stage]
  • MMA warp 开始用 input[load_stage] 计算
  • 同时 LOAD warp 可能在搬 chunk t+1 到另一个 stage

行 441-527: 取出 smem tensor + 准备 fragment

创建 Smem Tensor 视图

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
Tensor v_tile = make_tensor(
make_smem_ptr(shared_storage.input[load_stage].v.begin()),
VOLayout{});
// v_tile: smem 中的 v [16x128], 用 MMA 友好的 layout

Tensor beta_tile = make_tensor(
make_smem_ptr(shared_storage.input[load_stage].beta.begin()),
BetaSmemLayout{});
int beta_smem_offset = (head_idx * T_total + int(bos) + t * CHUNK) & 7;
// 跳过 TMA 对齐多 load 的那几个元素

Tensor out_tile = make_tensor(
make_smem_ptr(shared_storage.output[out_stage].out.begin()),
VOLayout{});
// out_tile: smem 中的 output buffer [16x128]

Tensor k_decayed = make_tensor(..., MMALayout{}); // [16x128]
Tensor q_decayed = make_tensor(..., MMALayout{}); // [16x128]
Tensor k_restored = make_tensor(..., MMALayout{}); // [16x128]
Tensor g_total = make_tensor(..., GTotalLayout{}); // [128]
Tensor INV = make_tensor(..., LMLayout{}); // [16x16]
Tensor Mqk = make_tensor(..., LMLayout{}); // [16x16]

Tensor s_acc = make_tensor(
make_smem_ptr(shared_storage.state_acc.begin()),
StateSmemLayout{});
// s_acc: 隐状态 [128x128], 常驻 smem

Tensor s_acc_T = make_tensor(
make_smem_ptr(shared_storage.state_acc.begin()),
TransposedStateSmemLayout{});
// s_acc_T: 同一片物理内存的转置视图
// Phase 1 用 s_acc (按列读), Phase 6 用 s_acc_T (按行读写)

Warp/线程信息

1
2
3
4
5
6
const int warp_id = compute_tid / 32;     // 0,1,2,3
const int lane_id = compute_tid % 32; // 0-31
const int group_id = (lane_id / 4) % 8; // 0-7

ThrMMA thr_mma = mma.get_slice(lane_id);
// 拿到这个 lane 的 MMA 视图: 告诉 CuTe 这个线程在 MMA 中的角色

128 个 MMA 线程的职责

Warp 线程范围 负责列 列块数
warp0 0-31 0-31 2 (0-15, 16-31)
warp1 32-63 32-63 2 (32-47, 48-63)
warp2 64-95 64-95 2 (64-79, 80-95)
warp3 96-127 96-127 2 (96-111, 112-127)

每个 Warp 内的 group_id 到行的映射 (MMA m16n8k16)

group_id lane 范围 行 (第一块) 行 (第二块)
0 0-3 0,1 8,9
1 4-7 2,3 10,11
2 8-11 4,5 12,13
3 12-15 6,7 14,15
4 16-19 8,9 0,1
5 20-23 10,11 2,3
6 24-27 12,13 4,5
7 28-31 14,15 6,7

Copy 对象

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
// A 操作数的 copy: smem (K_INTER layout) -> 寄存器 (LDSM_N)
auto smem_tiled_copy_A = make_tiled_copy_A(
Copy_Atom<SM75_U32x4_LDSM_N, BF16>{}, mma);
auto smem_thr_copy_A = smem_tiled_copy_A.get_thread_slice(lane_id);

// A 操作数转置版 (Phase 6 读 k_restored^T)
auto smem_tiled_copy_A_T = make_tiled_copy_A(
Copy_Atom<SM75_U16x8_LDSM_T, BF16>{}, mma);
auto smem_thr_copy_A_T = smem_tiled_copy_A_T.get_thread_slice(lane_id);

// B 操作数的 copy: smem -> 寄存器
auto smem_tiled_copy_B = make_tiled_copy_B(
Copy_Atom<SM75_U32x4_LDSM_N, BF16>{}, mma);
auto smem_thr_copy_B = smem_tiled_copy_B.get_thread_slice(lane_id);

// C 操作数的 load/store (读写 MMA 的输出)
auto smem_tiled_load_C = make_tiled_copy_C(
Copy_Atom<SM75_U32x4_LDSM_N, BF16>{}, mma);
auto smem_thr_load_C = smem_tiled_load_C.get_slice(lane_id);

auto smem_tiled_store_C = make_tiled_copy_C(
Copy_Atom<SM90_U32x4_STSM_N, BF16>{}, mma);
auto smem_thr_store_C = smem_tiled_store_C.get_slice(lane_id);

// C 转置版 (Phase 6 读写 s_acc_T)
auto smem_tiled_load_C_T = make_tiled_copy_C(
Copy_Atom<SM75_U16x8_LDSM_T, BF16>{}, mma);
auto smem_thr_load_C_T = smem_tiled_load_C_T.get_slice(lane_id);

auto smem_tiled_store_C_T = make_tiled_copy_C(
Copy_Atom<SM90_U16x8_STSM_T, BF16>{}, mma);
auto smem_thr_store_C_T = smem_tiled_store_C_T.get_slice(lane_id);

Copy 对象总结

用途 对象 指令 方向 使用场景
加载 A smem_thr_copy_A LDSM_N smem → reg Phase 1,2,3,4 (kd, qd, INV, Mqk)
加载 A 转置 smem_thr_copy_A_T LDSM_T smem → reg Phase 6 (k_restored^T)
加载 B smem_thr_copy_B LDSM_N smem → reg Phase 1 (s_acc)
读 C smem_thr_load_C LDSM_N smem → reg Phase 2 (v)
写 C smem_thr_store_C STSM_N reg → smem Phase 5 (out)
读 C 转置 smem_thr_load_C_T LDSM_T smem → reg Phase 6 (s_acc_T)
写 C 转置 smem_thr_store_C_T STSM_T reg → smem Phase 6 (s_acc_T)

创建参考 Tensor 和 Fragment

1
2
3
4
5
6
7
8
9
10
11
Tensor A_ref = local_tile(k_decayed, make_shape(Int<16>{}, Int<16>{}),
make_coord(0, 0));
// A_ref: k_decayed[16x128] 的左上角 16x16 块, 作为形状参考

Tensor B_ref = local_tile(s_acc, make_shape(Int<16>{}, Int<16>{}),
make_coord(0, 0));
// B_ref: s_acc[128x128] 的左上角 16x16 块

Tensor C_ref = local_tile(v_tile, make_shape(Int<16>{}, Int<16>{}),
make_coord(0, 0));
// C_ref: v_tile[16x128] 的左上角 16x16 块

:这三个 ref 只是用来推导 fragment 的形状,不会真的读这些数据。MMA 一次算 A[16×16] @ B[16×16] → C[16×16],所以参考都是 16×16。

1
2
3
4
5
// k_decayed 的 A fragment: 两份
Tensor tCrAi_k = make_fragment_like<BF16>(
thr_mma.partition_fragment_A(A_ref));
auto tCrAi_k_view = smem_thr_copy_A.retile_D(tCrAi_k);
auto tCrA_k = thr_mma.partition_fragment_A(A_ref);

thr_mma.partition_fragment_A(A_ref)

  • 告诉 CuTe:我这个线程在 MMA 中需要 A 操作数的哪些元素
  • 返回一个寄存器 tensor,形状由 MMA atom 决定
  • 对于 m16n8k16 + tile 16×16×16:每线程 4 个 uint32 = 8 个 bf16

命名规则

  • tCr = thread-level C-format register
  • Ai = A input (load 目标)
  • _k = 给 k_decayed 用
  • _view = retile 后的视图
1
2
3
4
5
6
7
8
9
// q_decayed 的 A fragment (同结构)
Tensor tCrAi_q = make_fragment_like<BF16>(...);
auto tCrAi_q_view = smem_thr_copy_A.retile_D(tCrAi_q);
auto tCrA_q = thr_mma.partition_fragment_A(A_ref);

// s_acc 的 B fragment
Tensor tCrBi = make_fragment_like<BF16>(...);
auto tCrBi_view = smem_thr_copy_B.retile_D(tCrBi);
auto tCrB = thr_mma.partition_fragment_B(B_ref);

为什么 A 有两份 (k 和 q) 但 B 只有一份?

Phase 1 中 k_decayed 和 q_decayed 用不同的 A,但共享同一个 B (s_acc 的列块)。

Fragment 结构(一个线程)

  • tCrA_k / tCrA_q:8 个 bf16 (4 个 uint32)
  • tCrB:8 个 bf16 (4 个 uint32)

初始化累加器

1
2
3
4
5
6
7
8
9
AccFragT u_acc[2], out_acc[2];
for (int i = 0; i < 2; ++i) {
u_acc[i] = thr_mma.make_fragment_C(tCrC_ref);
clear(u_acc[i]);
}
for (int i = 0; i < 2; ++i) {
out_acc[i] = thr_mma.make_fragment_C(tCrC_ref);
clear(out_acc[i]);
}
  • make_fragment_C:创建 MMA 累加器 fragment,fp32
  • 每线程持有 8 个 fp32 (16×16 矩阵中的 8 个元素)

为什么 [2]

每个 warp 负责 32 列 = 2 个 16×16 块:

  • [0] = 前 16 列
  • [1] = 后 16 列
Warp 负责列 [0] [1]
warp0 0-31 0-15 16-31
warp1 32-63 32-47 48-63
warp2 64-95 64-79 80-95
warp3 96-127 96-111 112-127

行 529-564: Phase 1 — 双 GEMM k@s 和 q@s

目标

  • u_acc[0..1] = k_decayed[16×128] @ s_acc[128×(warp的32列)]
  • out_acc[0..1] = q_decayed[16×128] @ s_acc[128×(warp的32列)]

K_BLOCKS = 128 / 16 = 8

预取第 0 步

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
constexpr int K_BLOCKS = decltype(cute::size<1>(k_decayed))::value / 16;
// K_BLOCKS = 128 / 16 = 8

// 预取 k_decayed[:, 0:16] 的 A fragment
copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(
local_tile(k_decayed, make_shape(Int<16>{}, Int<16>{}),
make_coord(0, 0))),
tCrAi_k_view);

// 预取 q_decayed[:, 0:16]
copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(
local_tile(q_decayed, make_shape(Int<16>{}, Int<16>{}),
make_coord(0, 0))),
tCrAi_q_view);

// 预取 s_acc 的 [warp_id*2, 0] 块
copy(smem_tiled_copy_B, smem_thr_copy_B.partition_S(
local_tile(s_acc, make_shape(Int<16>{}, Int<16>{}),
make_coord(warp_id * 2, 0))),
tCrBi_view);

local_tile(s_acc, (16,16), make_coord(warp_id*2, 0)) 的含义

  • s_acc[128×128] 按 16×16 分块
  • make_coord(warp_id*2, 0):warp 的第 0 列块,K 维度第 0 步
Warp 列块 s_acc 区域
warp0 0 s_acc[0:16, 0:16]
warp1 2 s_acc[0:16, 32:48]
warp2 4 s_acc[0:16, 64:80]
warp3 6 s_acc[0:16, 96:112]

预取图解

1
2
3
4
5
6
7
8
9
k_decayed [16 x 128]              s_acc [128 x 128]:
col: K0 K1 ... K7 col: w0列(0-31) w1列(32-63) ...
+---+---+---+---+ +---+---+---+---+---+---+---+---+
|K0 | | | | K0行 |B0 |B1 | | | | | | |
+---+---+---+---+ +---+---+---+---+---+---+---+---+
...

预取: kd[:, K0] = 16x16 预取: s_acc[warp*2, K0] = warp的第0列块, K维度第0步
qd[:, K0] = 16x16

主循环实现

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
#pragma unroll
for (int k = 0; k < K_BLOCKS; ++k) { // k = 0,1,...,7

// 格式转换: load layout -> compute layout
cute::transform(tCrAi_k, tCrA_k, cute::identity{});
cute::transform(tCrAi_q, tCrA_q, cute::identity{});
cute::transform(tCrBi, tCrB, cute::identity{});

// 预取 B 的下一个列块: s_acc 的 warp第1列块, K第k步
copy(smem_tiled_copy_B, smem_thr_copy_B.partition_S(
local_tile(s_acc, make_shape(Int<16>{}, Int<16>{}),
make_coord(warp_id * 2 + 1, k))),
tCrBi_view);

// 用第 0 列块做 gemm
gemm(thr_mma, tCrA_k(_,_,Int<0>{}), tCrB(_,_,Int<0>{}), u_acc[0]);
gemm(thr_mma, tCrA_q(_,_,Int<0>{}), tCrB(_,_,Int<0>{}), out_acc[0]);

// 等第 1 个列块加载完, 格式转换
cute::transform(tCrBi, tCrB, cute::identity{});

// 预取下一步的 A 和 B (如果不是最后一步)
if (k + 1 < K_BLOCKS) {
copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(
local_tile(k_decayed, make_shape(Int<16>{}, Int<16>{}),
make_coord(0, k + 1))),
tCrAi_k_view);
copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(
local_tile(q_decayed, make_shape(Int<16>{}, Int<16>{}),
make_coord(0, k + 1))),
tCrAi_q_view);
copy(smem_tiled_copy_B, smem_thr_copy_B.partition_S(
local_tile(s_acc, make_shape(Int<16>{}, Int<16>{}),
make_coord(warp_id * 2, k + 1))),
tCrBi_view);
}

// 用第 1 列块做 gemm
gemm(thr_mma, tCrA_k(_,_,Int<0>{}), tCrB(_,_,Int<0>{}), u_acc[1]);
gemm(thr_mma, tCrA_q(_,_,Int<0>{}), tCrB(_,_,Int<0>{}), out_acc[1]);
}

Phase 1 指令级并行

每步操作

  1. transform (零开销)
  2. 预取 B1,k (LDSM)
  3. gemm 使用 B0,k (2 条 mma)
  4. transform B1,k
  5. 预取 A_{k+1} 和 B0,k+1 (如需)
  6. gemm 使用 B1,k (2 条 mma)

每步:4 条 gemm = 8 条 mma 指令 8 步总计:64 条 mma 指令/warp

Phase 1 结束后的寄存器状态

累加器 内容
u_acc[0] k_decayed @ s_acc[:, warp第0列块] [16×16] fp32
u_acc[1] k_decayed @ s_acc[:, warp第1列块] [16×16] fp32
out_acc[0] q_decayed @ s_acc[:, warp第0列块] [16×16] fp32
out_acc[1] q_decayed @ s_acc[:, warp第1列块] [16×16] fp32
```

4 K2 kernel 逐行讲解 Part 3: Phase 2-6 + STORE warp

行 566-583: Phase 2 — 类型转换 + 加载 v/INV/beta

1
2
3
4
5
SFragT out_bf16[2];
#pragma unroll
for (int i = 0; i < 2; ++i)
cute::transform(out_acc[i], out_bf16[i],
[] __device__ (float x) { return BF16(x); });
  • out_acc[i] 是 fp32 (MMA 累加器)
  • 转成 bf16 存到 out_bf16[i]
  • 截断精度: fp32 的 23 位尾数 → bf16 的 7 位尾数

此时 out_bf16 = bf16(q_decayed @ s_acc),后面 Phase 4 会加上 Mqk @ U


加载 v

1
2
3
4
5
6
7
8
9
10
SFragT v_bf16[2];
#pragma unroll
for (int i = 0; i < 2; ++i) {
Tensor v_block = local_tile(v_tile,
make_shape(Int<16>{}, Int<16>{}),
make_coord(0, warp_id * 2 + i));
copy(smem_tiled_load_C,
smem_thr_load_C.partition_S(v_block),
smem_thr_load_C.retile_D(v_bf16[i]));
}

从 smem 加载 v 到寄存器,按 C 格式(因为后面要和 u 做逐元素减法)。

v_tile[16×128] 按 16×16 分块

Warp v_block[0] v_block[1]
warp0 v[:, 0:16] v[:, 16:32]
warp1 v[:, 32:48] v[:, 48:64]
warp2 v[:, 64:80] v[:, 80:96]
warp3 v[:, 96:112] v[:, 112:128]

加载 INV

1
2
3
copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(INV),
tCrAi_k_view);
cute::transform(tCrAi_k, tCrA_k, cute::identity{});
  • 加载 INV [16×16] 到 A fragment
  • 所有 4 个 warp 加载同一份 INV(它只有 16×16,不分列块)
  • Phase 3 用它做 U = INV @ u

加载 beta 并计算 sigmoid

1
2
3
4
BF16 beta0 = BF16(sigmoid_tanh_approx_f32(
float(beta_tile(beta_smem_offset + group_id))));
BF16 beta1 = BF16(sigmoid_tanh_approx_f32(
float(beta_tile(beta_smem_offset + group_id + 8))));
  • group_id = (lane_id / 4) % 8,范围 0-7

  • MMA 的 16 行分成两组:

  • 行 0-7:由 group_id 0-7 的线程负责,用 beta0 = sigmoid(beta[group_id])

  • 行 8-15:同样的线程负责,用 beta1 = sigmoid(beta[group_id+8])

  • 每个线程持有的 8 个 C 元素中:

  • 4 个属于行 0-7(用 beta0)

  • 4 个属于行 8-15(用 beta1)

beta_smem_offset 示例

1
2
3
beta_smem_offset=3, group_id=2
beta0 = sigmoid(beta_tile(3+2)) = sigmoid(beta[5]) → chunk 内第 5 个 token 的 beta
beta1 = sigmoid(beta_tile(3+2+8)) = sigmoid(beta[13]) → 第 13 个 token 的 beta

行 585-619: Phase 3 — u = (v - k@s) * beta; U = INV @ u

1
2
3
4
5
SFragT u_bf16[2];
uint32_t u_b_regs[4];

#pragma unroll
for (int i = 0; i < 2; ++i) { // 两个16x16列块

步骤 1: fp32 → bf16 转换

1
2
cute::transform(u_acc[i], u_bf16[i],
[] __device__ (float x) { return BF16(x); });

此时 u_bf16[i] = bf16(k_decayed @ s_acc),即 k@s 的列块 i。


步骤 2: u = (v - k@s) * beta

1
2
3
4
5
6
7
8
9
10
#pragma unroll
for (int a = 0; a < 2; ++a) {
#pragma unroll
for (int d = 0; d < 2; ++d) {
auto c0 = make_coord(make_coord(a, 0), 0, d);
auto c1 = make_coord(make_coord(a, 1), 0, d);
u_bf16[i](c0) = (v_bf16[i](c0) - u_bf16[i](c0)) * beta0;
u_bf16[i](c1) = (v_bf16[i](c1) - u_bf16[i](c1)) * beta1;
}
}

MMA C fragment 的坐标系统

  • make_coord(make_coord(a, row_half), 0, d)

  • a = 0,1:两个 “MMA 重复”(16行被分成 2×8)

  • row_half = 0:行 0-7(用 beta0)

  • row_half = 1:行 8-15(用 beta1)

  • d = 0,1:两个列(每线程在每个半行中持有 2 个值)

每线程持有2(a) × 2(d) × 2(row_half) = 8 个元素

逐元素含义

1
u[i][j] = (v[i][j] - k_decayed@s[i][j]) × sigmoid(beta[i])

新值 v 减去旧预测 k@s,乘以写入强度 beta。


步骤 3: MOVM_T 转置(C 格式 → B 格式)

1
2
3
4
5
uint32_t* u_c = reinterpret_cast<uint32_t*>(&u_bf16[i](0));
SM75_U32x1_MOVM_T::copy(u_c[0], u_b_regs[0]);
SM75_U32x1_MOVM_T::copy(u_c[1], u_b_regs[1]);
SM75_U32x1_MOVM_T::copy(u_c[2], u_b_regs[2]);
SM75_U32x1_MOVM_T::copy(u_c[3], u_b_regs[3]);
  • u_bf16 在寄存器中是 C 格式(MMA 输出格式)
  • 下面要做 INV @ u,u 需要作为 B 操作数(B 格式)

MOVM_T

  • 把 4 个 uint32 从 C 格式重排成 B 格式
  • 底层:warp 内 32 线程互相 shuffle 寄存器
  • 每条 MOVM_T 处理 1 个 uint32(= 2 个 bf16)
  • 4 条 = 8 个 bf16 = 一个线程在 16×16 中持有的全部 B 分量

传统做法 vs MOVM_T

  • 传统:C → smem (STSM) → smem → B (LDSM):两次 smem 访问
  • MOVM_T:C → warp shuffle → B:零 smem 访问!

步骤 4: 创建 B fragment

1
2
3
4
auto tCrB_u_tmp = thr_mma.partition_fragment_B(B_ref);
uint32_t* b_dst = reinterpret_cast<uint32_t*>(&tCrB_u_tmp(0));
b_dst[0] = u_b_regs[0]; b_dst[1] = u_b_regs[1];
b_dst[2] = u_b_regs[2]; b_dst[3] = u_b_regs[3];

创建一个 B fragment,把 MOVM_T 转置后的值填进去。tCrB_u_tmp 现在持有 u 的 B 格式。


步骤 5: U = INV @ u

1
2
3
clear(u_acc[i]);
gemm(thr_mma, tCrA_k(_,_,Int<0>{}), tCrB_u_tmp(_,_,Int<0>{}),
u_acc[i]);
  • tCrA_k = INV [16×16](Phase 2 加载的)
  • tCrB_u_tmp = u [16×16](B 格式)
  • 只需要 1 次 gemm(K=16,一步就够),内部 2 条 mma 指令

步骤 6: fp32 → bf16 转换

1
2
cute::transform(u_acc[i], u_bf16[i],
[] __device__ (float x) { return BF16(x); });

Phase 3 结束后

  • u_bf16[0]:U 的 warp 第 0 列块 [16×16] bf16
  • u_bf16[1]:U 的 warp 第 1 列块 [16×16] bf16

行 621-646: Phase 4 — out = q@s + Mqk@U

1
2
3
copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(Mqk),
tCrAi_k_view);
cute::transform(tCrAi_k, tCrA_k, cute::identity{});

加载 Mqk [16×16] 到 A fragment,覆盖之前的 INV。所有 warp 加载同一份 Mqk。


关键变量声明

1
BFragT_u tCrB_u_arr[2];   // 保留到 Phase 6!

tCrB_u_arr 会存 U 的 B 格式,Phase 6 更新 state 时还要用。


再次 MOVM_T:U → B 格式

1
2
3
4
5
6
7
8
9
10
11
12
13
#pragma unroll
for (int i = 0; i < 2; ++i) {
uint32_t* u_c = reinterpret_cast<uint32_t*>(&u_bf16[i](0));
SM75_U32x1_MOVM_T::copy(u_c[0], u_b_regs[0]);
SM75_U32x1_MOVM_T::copy(u_c[1], u_b_regs[1]);
SM75_U32x1_MOVM_T::copy(u_c[2], u_b_regs[2]);
SM75_U32x1_MOVM_T::copy(u_c[3], u_b_regs[3]);

tCrB_u_arr[i] = thr_mma.partition_fragment_B(B_ref);
uint32_t* b_dst = reinterpret_cast<uint32_t*>(&tCrB_u_arr[i](0));
b_dst[0] = u_b_regs[0]; b_dst[1] = u_b_regs[1];
b_dst[2] = u_b_regs[2]; b_dst[3] = u_b_regs[3];
}

为什么又做一次 MOVM_T?

Phase 3 结束时 u_bf16[i] 被更新成了 INV@u 的 C 格式结果。现在要把这个新的 U(C 格式)再次转成 B 格式

注意:这次存到 tCrB_u_arr[i] 而不是临时变量,因为 Phase 6 还要用。


out += Mqk @ U

1
2
3
4
5
6
7
8
9
10
    clear(out_acc[i]);
gemm(thr_mma, tCrA_k(_,_,Int<0>{}), tCrB_u_arr[i](_,_,Int<0>{}),
out_acc[i]);

SFragT gemm_bf16;
cute::transform(out_acc[i], gemm_bf16,
[] __device__ (float x) { return BF16(x); });
cute::transform(out_bf16[i], gemm_bf16, out_bf16[i],
[] __device__ (BF16 c, BF16 a) { return c + a; });
}
  • out_acc[i] = Mqk @ U
  • tCrA_k = Mqk [16×16],tCrB_u_arr[i] = U [16×16] B 格式
  • gemm_bf16 = bf16(Mqk @ U)
  • out_bf16[i] = out_bf16[i] + gemm_bf16 = q_decayed@s + Mqk@U = 最终输出

行 648-653: Phase 5 — 存储 output

1
2
3
4
5
6
7
8
9
#pragma unroll
for (int i = 0; i < 2; ++i) {
Tensor out_block = local_tile(out_tile,
make_shape(Int<16>{}, Int<16>{}),
make_coord(0, warp_id * 2 + i));
copy(smem_tiled_store_C,
smem_thr_store_C.retile_S(out_bf16[i]),
smem_thr_store_C.partition_D(out_block));
}

STSM:寄存器 out_bf16[i] → smem out_tile 的对应 16×16 块

4 个 warp 并行写 out_tile[16×128]

Warp 列范围
warp0 0-31
warp1 32-63
warp2 64-95
warp3 96-127

写完后,STORE warp 会通过 TMA 搬到 global memory。


行 655-727: Phase 6 — 隐状态更新

1
2
3
4
5
6
7
8
9
// s_acc[D, D] = s_acc * g_total + k_restored_t[D, 16] @ U[16, D]
constexpr int S_M_BLOCKS = decltype(cute::size<0>(k_restored_t))::value / 16;
// = 128 / 16 = 8 (s_acc 的 128 行分成 8 个行块)

Tensor k_restored_t = make_tensor(
make_smem_ptr(shared_storage.input[load_stage].k_restored.begin()),
TransposedMMALayout{});
// k_restored 原始 [16x128], TransposedMMALayout 让它逻辑上变成 [128x16]
// 物理数据不变, 只是读取时自动转置

创建 Phase 6 的 fragment

1
2
3
4
Tensor tCrAi_kr = make_fragment_like<BF16>(
thr_mma.partition_fragment_A(A_ref));
auto tCrAi_kr_view = smem_thr_copy_A_T.retile_D(tCrAi_kr);
// 注意用的是 copy_A_T (LDSM_T,转置加载)

Ring buffer 结构(PREFETCH=1,就是单个 buffer):

Buffer 内容
ring_A_kr[0] k_restored^T 的当前行块
ring_S_acc[bi][0] s_acc^T 的当前行块(bi=0:warp第0列块,bi=1:第1列块)
ring_g0[0] g_total 的行 0-7 对应值
ring_g1[0] g_total 的行 8-15 对应值

预取第 0 行块

1
2
3
4
5
6
7
8
9
#pragma unroll
for (int i = 0; i < PREFETCH; ++i) { // i = 0
// k_restored^T 的第 0 个 [16x16] 行块
Tensor kr_block = local_tile(k_restored_t,
make_shape(Int<16>{}, Int<16>{}), make_coord(0, 0));
copy(smem_tiled_copy_A_T,
smem_thr_copy_A_T.partition_S(kr_block),
tCrAi_kr_view);
cute::transform(tCrAi_kr, ring_A_kr[0], cute::identity{});

k_restored_t [128×16] 按 16×16 分块,第 0 行块 = [0:16, 0:16]LDSM_T:从 smem 转置加载到寄存器。

1
2
3
4
5
6
7
8
9
#pragma unroll
for (int bi = 0; bi < 2; ++bi) {
Tensor s_block = local_tile(s_acc_T,
make_shape(Int<16>{}, Int<16>{}),
make_coord(0, warp_id * 2 + bi));
copy(smem_tiled_load_C_T,
smem_thr_load_C_T.partition_S(s_block),
smem_thr_load_C_T.retile_D(ring_S_acc[bi][0]));
}

s_acc_T [128×128] 按 16×16 分块:make_coord(0, warp_id*2+bi) = 第 0 行块,warp 的第 bi 个列块。LDSM_T:转置加载,让数据在寄存器中按 C 格式排列。

1
2
3
    ring_g0[0] = g_total(0 * 16 + group_id);
ring_g1[0] = g_total(0 * 16 + group_id + 8);
}
  • g_total[128] 是 fp32 标量数组
  • ring_g0[0] = g_total[group_id],对应行 0-7(每个线程不同的行)
  • ring_g1[0] = g_total[group_id+8],对应行 8-15

这些 g_total 值已经是 exp2(cumsum) 形式,直接当衰减因子用。


为什么用 s_acc_T(转置视图)?

  • s_accStateSmemLayout 是为 Phase 1 的 B 操作数设计的(按列读取)
  • Phase 6 需要按行更新,用转置视图 + LDSM_T/STSM_T 实现
  • 物理内存不变,只是读写模式变了

Phase 6 主循环

1
2
3
4
5
6
#pragma unroll
for (int m = 0; m < S_M_BLOCKS; ++m) { // m = 0,1,...,7
const int slot = m % PREFETCH; // = 0 (PREFETCH=1)

float g0 = ring_g0[slot];
float g1 = ring_g1[slot];

步骤 1: k_restored^T @ U

1
2
3
4
5
6
#pragma unroll
for (int bi = 0; bi < 2; ++bi) {
clear(u_acc[bi]);
gemm(thr_mma, ring_A_kr[slot](_,_,Int<0>{}),
tCrB_u_arr[bi](_,_,Int<0>{}), u_acc[bi]);
}
  • ring_A_kr = k_restored^T 的第 m 行块 [16×16],A 格式
  • tCrB_u_arr[bi] = U 的第 bi 列块 [16×16],B 格式(从 Phase 4 保留!)
  • u_acc[bi] = k_restored^T[m行块] @ U[warp列块bi] [16×16] fp32
  • 2 次 gemm = 4 条 mma 指令

MOVM_T 的价值体现:U 一直在寄存器里,被 3 次复用:

  • Phase 3:INV @ u(作为 B)
  • Phase 4:Mqk @ U(作为 B)
  • Phase 6:k_restored^T @ U(作为 B)

步骤 2: 预取下一行块(与 gemm 重叠)

1
2
3
4
5
6
7
8
9
10
11
12
if (m + PREFETCH < S_M_BLOCKS) {
Tensor kr_next = local_tile(k_restored_t,
make_shape(Int<16>{}, Int<16>{}),
make_coord(m + PREFETCH, 0));
copy(smem_tiled_copy_A_T,
smem_thr_copy_A_T.partition_S(kr_next),
tCrAi_kr_view);
cute::transform(tCrAi_kr, ring_A_kr[slot], cute::identity{});

ring_g0[slot] = g_total((m + PREFETCH) * 16 + group_id);
ring_g1[slot] = g_total((m + PREFETCH) * 16 + group_id + 8);
}

gemm 在算当前行块时,LDSM 在加载下一行块(指令级并行)。

步骤 3: 逐元素更新 s_acc

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
#pragma unroll
for (int bi = 0; bi < 2; ++bi) {
#pragma unroll
for (int a = 0; a < 2; ++a) {
#pragma unroll
for (int d = 0; d < 2; ++d) {
auto c0 = make_coord(make_coord(a, 0), 0, d);
auto c1 = make_coord(make_coord(a, 1), 0, d);
ring_S_acc[bi][slot](c0) = BF16(
bf16_to_f32(ring_S_acc[bi][slot](c0)) * g0
+ u_acc[bi](c0));
ring_S_acc[bi][slot](c1) = BF16(
bf16_to_f32(ring_S_acc[bi][slot](c1)) * g1
+ u_acc[bi](c1));
}
}

每个元素的更新公式

1
s_new = bf16( f32(s_old) × g + gemm_result )
  • s_old:bf16,从 ring_S_acc 读取
  • bf16_to_f32:转 fp32(无精度损失,bf16 是 fp32 的子集)
  • × g:fp32 乘法,g 是 exp2(g_total[对应行])
  • + u_acc:fp32 加法,gemm 结果
  • bf16():转回 bf16 存储(截断精度)

步骤 4: 写回 s_acc + 预取下一块

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
        Tensor s_block = local_tile(s_acc_T,
make_shape(Int<16>{}, Int<16>{}),
make_coord(m, warp_id * 2 + bi));
copy(smem_tiled_store_C_T,
smem_thr_store_C_T.retile_S(ring_S_acc[bi][slot]),
smem_thr_store_C_T.partition_D(s_block));

if (m + PREFETCH < S_M_BLOCKS) {
Tensor s_next = local_tile(s_acc_T,
make_shape(Int<16>{}, Int<16>{}),
make_coord(m + PREFETCH, warp_id * 2 + bi));
copy(smem_tiled_load_C_T,
smem_thr_load_C_T.partition_S(s_next),
smem_thr_load_C_T.retile_D(ring_S_acc[bi][slot]));
}
}
}
  • STSM_T:转置写回 smem,ring_S_acc(C 格式寄存器)→ s_acc_T[m, warp列块bi]
  • LDSM_T:预取 s_acc_T 的下一个行块(与 STSM_T 写回不冲突,写的是 m,读的是 m+1)

Phase 6 结束s_acc[128×128] 已全部更新。


Phase 6 数据流图

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
k_restored^T   U (寄存器)    s_acc (smem)    g_total
[128x16] [16x128] [128x128] [128]
| | | |
第m行块 warp列块 第m行块,warp列 第m行
[16x16] [16x32] [16x32]
| | | |
+---- gemm ----+ | |
| [16x32] | |
v v v
gemm_result s_old * exp2(g_total)
\ /
\ /
+------ + -------+
|
v
s_new [16x32] → 写回 smem

4 个 Warp 的分工(Phase 6)

m 所有 Warp 更新 warp0 更新 warp1 更新 warp2 更新 warp3 更新
0 s_acc[0:16, :] [0:16, 0:32] [0:16, 32:64] [0:16, 64:96] [0:16, 96:128]
1 s_acc[16:32, :] [16:32, 0:32] [16:32, 32:64] [16:32, 64:96] [16:32, 96:128]
7 s_acc[112:128, :] [112:128, 0:32] [112:128, 32:64] [112:128, 64:96] [112:128, 96:128]

行 729-737: 同步 + Pipeline 推进

1
2
}
compute_barrier.arrive_and_wait();

128 个 MMA 线程的 barrier,确保 4 个 warp 都写完了:

  1. out_tile(Phase 5)
  2. s_acc(Phase 6)

才能进入下一个 chunk。

1
2
3
4
5
cutlass::arch::fence_view_async_shared();
store_pipeline.producer_commit(out_write);
load_pipeline.consumer_release(load_read);
++load_read;
++out_write;
  • fence_view_async_shared:内存栅栏,STSM 是异步写,fence 确保写入对其他 warp 可见
  • producer_commit(out_write):通知 STORE warp “output[out_stage] 写好了,你可以存了”
  • consumer_release(load_read):通知 LOAD warp “input[load_stage] 我用完了,你可以覆盖了”
  • ++load_read:下一个 chunk 用下一个 input stage
  • ++out_write:下一个 chunk 写下一个 output stage

行 742-798: STORE warp 主循环

1
2
3
4
5
6
7
8
9
10
11
12
13
if (warp_role == WarpRole::STORE && lane_predicate) {
// warp5 的 lane0

Tensor g_out = tma_store_out.get_tma_tensor(make_shape(H, T_total, D));
auto cta_tma_store = tma_store_out.get_slice(Int<0>{});
StorePipelineState out_read;

for (int t = 0; t < t_tiles; ++t) {
store_pipeline.consumer_wait(out_read);
// 等 MMA warp commit

int stage = out_read.index();
int actual_len = min(CHUNK, seq_len - t * CHUNK);

尾部 chunk 处理

1
2
3
4
5
6
7
8
9
10
11
BF16* out_stage_ptr = shared_storage.output[stage].out.begin();

if (actual_len < CHUNK) {
// 尾部 chunk: 手动逐元素写
Tensor s_out = make_tensor(make_smem_ptr(out_stage_ptr), VOLayout{});
for (int row = 0; row < actual_len; ++row) {
int64_t global_base = (bos + t * CHUNK + row) * H * D + head_idx * D;
for (int col = 0; col < D; ++col) {
out_raw_ptr[global_base + col] = s_out(row, col);
}
}

为什么尾部不用 TMA?

  • TMA 会写满整个 16×128 = 2048 个 bf16
  • 如果序列只剩 5 个 token,TMA 会越界写 11 行,覆盖下一个序列的数据
  • 所以用 raw pointer 手动写,只写 actual_len

完整 chunk 处理

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
        } else {
// 完整 chunk: TMA store
auto out_off = g_out.layout()(head_idx, int(bos) + t * CHUNK, 0);
Tensor g_out_tile = make_tensor(g_out.data() + out_off,
make_layout(make_shape(Int<1>{}, Int<CHUNK>{}, Int<D>{}),
stride(g_out.layout())));
Tensor s_out_tile = make_tensor(make_smem_ptr(out_stage_ptr),
TMAVOLayout{});
cute::copy(tma_store_out,
cta_tma_store.partition_S(s_out_tile),
cta_tma_store.partition_D(g_out_tile));
tma_store_arrive();
}

tma_store_wait<0>();
store_pipeline.consumer_release(out_read);
++out_read;
}
}

TMA store:smem → global memory,一条指令

  • tma_store_arrive():发起 TMA store 后标记
  • tma_store_wait<0>():等 TMA store 完成(0 表示等所有 pending store)
  • consumer_release:通知 MMA warp “output[stage] 我存完了,你可以覆盖了”

序列长度 50,CHUNK=16 示例

chunk actual_len 存储方式 写入行
0 16 TMA store 0-15
1 16 TMA store 16-31
2 16 TMA store 32-47
3 2 手动写 48-49

行 782-833: 存储 final_state

所有 chunk 处理完后。

bf16 state(HasStateOut && !StateFP32

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
if constexpr (HasStateOut && !StateFP32) {
Tensor g_final = tma_store_final_state.get_tma_tensor(
make_shape(N * H, D, D));
auto state_off = g_final.layout()(seq_idx * H + head_idx, 0, 0);
Tensor g_final_tile = make_tensor(g_final.data() + state_off,
make_layout(make_shape(Int<1>{}, Int<D>{}, Int<D>{}),
stride(g_final.layout())));
Tensor s_state = make_tensor(
make_smem_ptr(shared_storage.state_acc.begin()),
TMAStateSmemLayout{});

auto cta_tma_store_state = tma_store_final_state.get_slice(Int<0>{});
cute::copy(
tma_store_final_state,
cta_tma_store_state.partition_S(s_state),
cta_tma_store_state.partition_D(g_final_tile)
);
tma_store_arrive();
}

STORE warp:TMA 直接从 state_acc 存到 global memory。state_acc[128×128] bf16 = 32KB,一次 TMA store。


fp32 state(HasStateOut && StateFP32

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
if constexpr (HasStateOut && StateFP32) {
__syncthreads(); // 等所有 warp (pipeline 结束)

// 全 192 线程: bf16 -> fp32 转换
smem_cvt_bf16_to_fp32<StateSmemLayout, FP32StateSmemLayout, D, NumThreads>(
shared_storage.state_acc.begin(), // 源: bf16 state
reinterpret_cast<float*>(shared_storage.state_fp32_buf), // 目标: fp32 buffer
threadIdx.x);
__syncthreads();

if (warp_role == WarpRole::STORE && lane_predicate) {
// TMA store fp32 state
cute::copy(tma_store_final_state, s_fp32, g_final_tile);
tma_store_arrive();
}
}

__syncthreads();

fp32 state 输出流程

  1. Pipeline 结束,此时 union 空间的 input/output buffer 不再使用
  2. 全线程把 state_acc(bf16,32KB)转成 state_fp32_buf(fp32,64KB)
  • fp32 buffer 占 union 空间,正好够用
  1. STORE warp 通过 TMA 存到 global memory

最后 __syncthreads:确保 TMA store 发起后再退出 kernel(TMA 是异步的,但 kernel 退出前必须保证完成)。


附录(一)

FlashKDA K2 流水线机制详解

一、为什么需要流水线

K2 kernel 的工作: 串行遍历一个序列的所有 chunk,做 delta rule 递推。

朴素做法 (无流水线):

1
2
3
4
for 每个 chunk:
从 global memory 加载数据 → shared memory (~500 cycle)
MMA 计算 Phase 1-6 (~2000 cycle)
把 output 从 shared memory 存回 global (~500 cycle)

总延迟 = t_tiles × (500 + 2000 + 500) = t_tiles × 3000 cycle 其中加载和存储时,MMA 完全空闲

流水线做法: 让加载、计算、存储同时进行:

1
2
3
4
时间步:     0      1      2      3      4
LOAD: [ L0 ][ L1 ][ L2 ][ L3 ][ L4 ]
MMA: [ C0 ][ C1 ][ C2 ][ C3 ][ C4 ]
STORE: [ S0 ][ S1 ][ S2 ][ S3 ][ S4 ]

总延迟 ≈ t_tiles × max(L, C, S) + 启动/排空开销 如果 C 是瓶颈: 总延迟 ≈ t_tiles × 2000 (加载/存储被完全隐藏)

二、三种 Warp 角色

K2 用 warp specialization: 不同 warp 永久承担不同角色。

1
2
3
4
5
192 个线程 = 6 个 warp:

线程号: 0-31 32-63 64-95 96-127 128-159 160-191
warp_id: 0 1 2 3 4 5
角色: MMA MMA MMA MMA LOAD STORE

代码 (utils.cuh 行 79-84):

1
enum class WarpRole { MMA, LOAD_QKG, STORE, NonParticipant };

分配逻辑 (fwd_kernel2.cuh 行 189-197):

1
2
3
warp_id < 4  →  MMA
warp_id == 4 → LOAD_QKG
warp_id == 5 → STORE

为什么 warp specialization 而不是全部线程一起干?

  1. TMA 指令只需 1 个线程发起,硬件自动搬运,多线程没意义
  2. MMA 计算需要 4 个 warp (128 线程) 才能充分利用 Tensor Core
  3. 分开后三者可以真正并行,不需要 __syncthreads() 全局同步

三、两条 Pipeline 的类型

K2 使用 CUTLASS 提供的两种 pipeline:

3.1 Input Pipeline: PipelineTmaAsync<3>

  • 类型: cutlass::PipelineTmaAsync<InputStages> (InputStages=3)
  • 本质: 基于 TMA barrier 的异步 pipeline
  • 特点: producer 发 TMA 后不用管,硬件搬完自动通知 consumer
  • stages: 3 (三缓冲)

角色分配:

  • LOAD warp → Producer (发 TMA,绑 barrier)
  • MMA warps → Consumer (等 barrier,读数据)

创建 (utils.cuh 行 88-116):

1
2
3
4
5
6
7
make_load_pipeline<3>(
shared_storage.load_pipeline, // barrier 在 smem 中
kTmaTransactionBytes, // 17984 字节 / chunk
warp_role, // 决定 producer or consumer
1, // 1 个 producer 线程
128 // 128 个 consumer 线程
)

关键参数: transaction_bytes = 17984 pipeline 内部的 barrier 会计数: TMA 搬了多少字节? 搬够 17984 字节 → barrier 自动 arrive → consumer 被唤醒

3.2 Output Pipeline: PipelineAsync<2>

  • 类型: cutlass::PipelineAsync<OutputStages> (OutputStages=2)
  • 本质: 基于 arrive/wait 的软件 pipeline (非 TMA)
  • 特点: producer (MMA) 手动 commit,consumer (STORE) 手动 wait
  • stages: 2 (双缓冲)

角色分配:

  • MMA warps → Producer (STSM 写 output,然后 commit)
  • STORE warp → Consumer (wait,然后 TMA store)

创建 (utils.cuh 行 118-143):

1
2
3
4
5
6
make_store_pipeline<2>(
shared_storage.store_pipeline,
warp_role,
128, // 128 个 producer 线程
1 // 1 个 consumer 线程
)

注意: producer_arv_count=128 128 个 MMA 线程都要 arrive,barrier 才算满足。 这保证 4 个 warp 全部写完 output 后 STORE 才开始搬。

3.3 两种 pipeline 的区别

PipelineTmaAsync PipelineAsync
同步机制 TMA barrier (硬件) arrive/wait barrier (软件)
producer 通知 TMA 硬件自动 arrive 手动调 producer_commit
consumer 通知 自动 (搬完字节数达标) 手动等 arrive 计数达标
用途 global→smem (TMA load) smem→global (TMA store)
为什么不同 TMA load 有硬件字节计数 STSM 是软件写,没有字节计数

四、PipelineState: 状态机

每个 pipeline 的 producer 和 consumer 各有一个 PipelineState:

1
2
3
4
PipelineState<Stages> {
int index_; // 当前 stage 编号 (0, 1, ..., Stages-1)
int phase_; // 当前 phase (用于 barrier 的奇偶翻转)
}
  • index() 方法: 返回 stage 编号,即 smem buffer 的下标
  • ++操作: index 循环递增 (0→1→2→0→1→2→…),phase 在溢出时翻转

Load pipeline (3-stage):

  • load_write: LOAD warp 的写指针,指向下一个要写入的 stage
  • load_read: MMA warps 的读指针,指向下一个要消费的 stage

Store pipeline (2-stage):

  • out_write: MMA warps 的写指针
  • out_read: STORE warp 的读指针

初始化:

1
2
3
4
5
6
7
8
load_write = make_producer_start_state<LoadPipeline>()
// → index=0,但 phase 经过特殊初始化,让 acquire 直接通过

load_read = LoadPipelineState()
// → index=0, phase=0, consumer 从 stage 0 开始等

out_write = make_producer_start_state<StorePipeline>()
out_read = StorePipelineState()

五、Pipeline API: 每个调用的含义

5.1 LOAD warp 的 API

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
// ┌────────────────────────────────────────────────────────┐
// │ producer_acquire(load_write) │
// │ │
// │ "我要写 stage X, 它空了吗?" │
// │ │
// │ 检查 stage X 的 barrier: │
// │ 如果 MMA 还在用 (没 release) → 阻塞等待 │
// │ 如果 MMA 已 release (或从未使用) → 立刻通过 │
// │ │
// │ 底层: barrier.try_wait(phase) │
// └────────────────────────────────────────────────────────┘

// ┌────────────────────────────────────────────────────────┐
// │ producer_get_barrier(load_write) │
// │ │
// │ 返回 stage X 对应的 barrier 指针 │
// │ 后续 TMA copy 绑定到这个 barrier: │
// │ cute::copy(tma_load.with(*barrier), src, dst) │
// │ │
// │ TMA 每搬完一块数据, barrier 的字节计数自动增加 │
// │ 搬够 kTmaTransactionBytes 字节 → barrier arrive │
// └────────────────────────────────────────────────────────┘

// ┌────────────────────────────────────────────────────────┐
// │ ++load_write │
// │ │
// │ 推进到下一个 stage: 0→1→2→0→1→2→... │
// │ 注意: 没有显式的 "commit"! │
// │ PipelineTmaAsync 不需要手动 commit, TMA 搬完后 │
// │ barrier 自动满足, consumer 自动被唤醒 │
// └────────────────────────────────────────────────────────┘

// ┌────────────────────────────────────────────────────────┐
// │ producer_tail(load_write) │
// │ │
// │ "没有更多数据了" │
// │ 对剩余未使用的 stage 的 barrier 做特殊标记 │
// │ 这样 consumer 等到最后一个 stage 后不会永远阻塞 │
// └────────────────────────────────────────────────────────┘

5.2 MMA warps 的 API (consumer of input, producer of output)

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
// ── 作为 input consumer ──

// ┌────────────────────────────────────────────────────────┐
// │ consumer_wait(load_read) │
// │ │
// │ "stage X 的数据到了吗?" │
// │ │
// │ 检查 stage X 的 TMA barrier: │
// │ 如果 TMA 还没搬完 → 阻塞等待 │
// │ 如果字节计数已达标 → 立刻通过 │
// │ │
// │ 128 个 MMA 线程全部在这里等,然后同时通过 │
// │ 底层: barrier.wait(phase) │
// └────────────────────────────────────────────────────────┘

// ┌────────────────────────────────────────────────────────┐
// │ consumer_release(load_read) │
// │ │
// │ "stage X 的数据我用完了, LOAD 可以覆盖了" │
// │ │
// │ 重置 stage X 的 barrier,翻转 phase │
// │ LOAD warp 的 producer_acquire 检测到翻转 → 通过 │
// │ │
// │ 注意: 放在 Phase 6 和 compute_barrier 之后 │
// │ 确保 4 个 warp 都算完才释放 │
// └────────────────────────────────────────────────────────┘

// ── 作为 output producer ──

// ┌────────────────────────────────────────────────────────┐
// │ producer_acquire(out_write) │
// │ │
// │ "output stage Y 空了吗?" │
// │ │
// │ 如果 STORE 还在读 stage Y → 阻塞等待 │
// │ 如果 STORE 已 release → 通过 │
// └────────────────────────────────────────────────────────┘

// ┌────────────────────────────────────────────────────────┐
// │ producer_commit(out_write) │
// │ │
// │ "output stage Y 写好了, STORE 可以存了" │
// │ │
// │ 128 个 MMA 线程都调用 arrive │
// │ arrive 计数达到 128 → barrier 满足 → STORE 被唤醒 │
// │ │
// │ 前面有 compute_barrier.arrive_and_wait(): │
// │ 确保 4 个 warp 都写完 out_tile 和 s_acc 才 commit │
// └────────────────────────────────────────────────────────┘

5.3 STORE warp 的 API

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
// ┌────────────────────────────────────────────────────────┐
// │ consumer_wait(out_read) │
// │ │
// │ "output stage Y 的数据写好了吗?" │
// │ │
// │ 等 MMA 的 128 个线程全部 arrive → 通过 │
// └────────────────────────────────────────────────────────┘

// ┌────────────────────────────────────────────────────────┐
// │ consumer_release(out_read) │
// │ │
// │ "output stage Y 我存完了, MMA 可以覆盖了" │
// │ │
// │ 发在 tma_store_wait<0>() 之后 (确保 TMA store 完成) │
// └────────────────────────────────────────────────────────┘

六、三个 warp 角色的完整代码结构

6.1 LOAD warp (fwd_kernel2.cuh 行 319-417)

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
if (warp_role == LOAD_QKG && lane_predicate) {
// 只有 warp4 的 1 个线程

LoadPipelineState load_write = make_producer_start_state();

for (int t = 0; t < t_tiles; ++t) {

load_pipeline.producer_acquire(load_write); ← 等 stage 可写
auto* barrier = load_pipeline.producer_get_barrier(load_write);
int stage = load_write.index();

// 8 个 TMA load,全部绑定到同一个 barrier:
copy(tma_load_v.with(*barrier), ..., input[stage].v);
copy(tma_load_beta.with(*barrier), ..., input[stage].beta);
copy(tma_load_ws_kd.with(*barrier), ..., input[stage].k_decayed);
copy(tma_load_ws_qd.with(*barrier), ..., input[stage].q_decayed);
copy(tma_load_ws_kr.with(*barrier), ..., input[stage].k_restored);
copy(tma_load_ws_gt.with(*barrier), ..., input[stage].g_total);
copy(tma_load_ws_inv.with(*barrier), ..., input[stage].INV);
copy(tma_load_ws_mqk.with(*barrier), ..., input[stage].Mqk);
// 8 个 TMA 是异步的,发出就返回
// 全部搬完后 barrier 自动 arrive (字节数 = 17984)

++load_write; // 推进 stage
}
load_pipeline.producer_tail(load_write); // 没有更多数据
}

6.2 MMA warps (fwd_kernel2.cuh 行 421-738)

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
if (warp_role == MMA) {
// 128 个线程 (warp 0-3)

NamedBarrier compute_barrier(1280); ← 4 个 warp 的内部同步
LoadPipelineState load_read;
StorePipelineState out_write = make_producer_start_state();

for (int t = 0; t < t_tiles; ++t) {

store_pipeline.producer_acquire(out_write); ← 等 output stage 可写
load_pipeline.consumer_wait(load_read); ← 等 input stage 数据就绪

int load_stage = load_read.index(); // 从哪个 input buffer 读
int out_stage = out_write.index(); // 写到哪个 output buffer

// ====== Phase 1-6 计算 (用 input[load_stage] 的数据) ======
// Phase 1: kd @ s_acc, qd @ s_acc
// Phase 2: fp32→bf16 + 加载 v/INV/beta
// Phase 3: u=(v-k@s)*β, U=INV@u
// Phase 4: out=q@s+Mqk@U
// Phase 5: STSM out_bf16 → output[out_stage].out
// Phase 6: s_acc 更新
// =========================================================

compute_barrier.arrive_and_wait(); ← 4 个 MMA warp 同步

fence_view_async_shared(); ← smem 写入可见性
store_pipeline.producer_commit(out_write); ← 通知 STORE
load_pipeline.consumer_release(load_read); ← 通知 LOAD
++load_read;
++out_write;
}
}

6.3 STORE warp (fwd_kernel2.cuh 行 742-780)

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
if (warp_role == STORE && lane_predicate) {
// 只有 warp5 的 1 个线程

StorePipelineState out_read;

for (int t = 0; t < t_tiles; ++t) {

store_pipeline.consumer_wait(out_read); ← 等 MMA 写好 output
int stage = out_read.index();
int actual_len = min(CHUNK, seq_len - t * CHUNK);

if (actual_len < CHUNK) {
// 尾部 chunk: 手动逐元素写 (避免 TMA 越界)
for row, col: out_raw_ptr[...] = smem[row][col];
} else {
// 完整 chunk: TMA store
copy(tma_store_out, smem → global);
tma_store_arrive();
}

tma_store_wait<0>(); ← 等 TMA store 完成
store_pipeline.consumer_release(out_read); ← 通知 MMA: stage 可覆盖
++out_read;
}
}

七、Pipeline 时序: 详细展开 (t_tiles=5)

假设: LOAD 约 1 单位时间,MMA 约 2 单位,STORE 约 1 单位

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
时间:  0   1   2   3   4   5   6   7   8   9  10  11  12
├───┼───┼───┼───┼───┼───┼───┼───┼───┼───┼───┼───┤

LOAD: L0 L1 L2 · L3 · L4
stg0 stg1 stg2 wait stg0 wait stg1
↑等C0 ↑等C1
release release

MMA: C0────── C1────── C2────── C3────── C4──────
stg0 stg1 stg2 stg0 stg1
↑wait L0 ↑wait L1

STORE: · S0 · S1 · S2 · S3 · S4
stg0 stg1 stg0 stg1 stg0
↑wait C0

stage 使用情况 (input pipeline):

1
2
3
4
5
6
t=0: L→stg0                    stg0: 被写入
t=1: L→stg1, M←stg0 stg0: 被读取 stg1: 被写入
t=2: L→stg2, M←stg1 stg1: 被读取 stg2: 被写入
t=3: L等待, M←stg2, release0 stg0: 被释放 stg2: 被读取
t=4: L→stg0, M←..., release1 stg0: 被写入 stg1: 被释放
...

7.1 启动阶段 (pipeline filling)

  • t=0: LOAD 写 stg0,MMA 在等 (consumer_wait 阻塞)
  • t=1: LOAD 写 stg1,stg0 barrier arrive → MMA 开始 C0
  • t=2: LOAD 写 stg2,MMA 还在算 C0
  • t=3: 3 个 stage 全满! LOAD 的 acquire 阻塞,等 MMA release

这就是三缓冲的价值: LOAD 连发 3 个 TMA 不停顿。 如果 LOAD 比 MMA 快,这 3 步的加载延迟被完全隐藏。

7.2 稳态阶段 (pipeline steady state)

每当 MMA release 一个 stage:

  • LOAD acquire 通过 → 写新数据 → ++stage
  • MMA consumer_wait 下一个 stage → 数据已就绪 (提前搬好的)

稳态下: MMA 是瓶颈,LOAD 和 STORE 的延迟被完全隐藏。 总延迟 ≈ t_tiles × C (MMA时间) + 启动排空开销

7.3 排空阶段 (pipeline draining)

  • LOAD 搬完最后一个 chunk → producer_tail()
  • MMA 算完最后一个 chunk → 不再 consumer_wait
  • STORE 存完最后一个 chunk → 循环结束

八、两条 Pipeline 的交互

MMA warps 同时是 input pipeline 的 consumer 和 output pipeline 的 producer。

每个 chunk 的 MMA 循环体开头要 两个等待:

1
2
store_pipeline.producer_acquire(out_write);  ← 等 output stage 可写
load_pipeline.consumer_wait(load_read); ← 等 input stage 就绪

为什么 store acquire 在 load wait 前面? 先确认 output buffer 可写,再等输入数据。 如果反过来: 数据到了但 output buffer 没空,还是得等。 实际上两个 wait 可以任意顺序,但先检查更快满足的那个可以微微减少等待。

循环体结尾:

1
2
3
4
5
6
compute_barrier.arrive_and_wait();           ← 4 个 warp 同步
fence_view_async_shared(); ← smem 写入可见
store_pipeline.producer_commit(out_write); ← 通知 STORE
load_pipeline.consumer_release(load_read); ← 通知 LOAD
++load_read;
++out_write;

两个通知同时发出,STORE 和 LOAD 同时被唤醒。

九、Barrier 的底层机制

9.1 TMA Barrier (Input Pipeline)

Hopper (SM90) 新增的硬件 barrier,存在 shared memory 中。

核心能力: 字节计数

  • arrive_and_expect_tx(N): 告诉 barrier “我期望 N 字节”
  • TMA 硬件每搬完一块: barrier.bytes_arrived += 搬运字节数
  • 当 bytes_arrived >= expected → barrier 自动满足
1
2
3
4
5
6
7
8
9
10
11
// LOAD warp:
barrier.arrive_and_expect_tx(17984) // 设定期望字节数
TMA copy × 8 (绑定到 barrier) // 发出 8 个异步搬运
// ... LOAD warp 可以去忙别的了 ...
// 8 个 TMA 陆续完成,字节数累加
// 字节数达到 17984 → barrier 自动 arrive

// MMA warps:
barrier.wait(phase) // 等 barrier arrive
// 阻塞,直到 TMA 搬够了字节
// 一旦通过,smem 中的数据保证完整

这就是为什么 LOAD warp 不需要显式 commit: TMA 硬件自己在计数,搬完就通知。

9.2 Software Barrier (Output Pipeline)

Output pipeline 用的是软件 arrive barrier:

1
2
3
4
producer_commit: 128 个 MMA 线程各调一次 arrive()
arrive 计数达到 128 → barrier 满足

consumer_wait: STORE warp 等 arrive 计数达标

为什么不用 TMA barrier? 因为 output 是 MMA warp 用 STSM 写入 smem 的 (不是 TMA)。 STSM 是普通的 smem 写指令,没有硬件字节计数能力。 所以用软件 arrive 计数代替。

十、compute_barrier vs pipeline barrier

两种不同的 barrier,容易混淆:

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
// ┌─────────────────────────────────────────────────────────────────┐
// │ compute_barrier = NamedBarrier(128, 0) │
// │ │
// │ 用途: 4 个 MMA warp 之间的同步 │
// │ 时机: Phase 6 结束后 │
// │ │
// │ 为什么需要: │
// │ Phase 5: 每 warp STSM 写 out_tile 的不同列 │
// │ Phase 6: 每 warp STSM_T 写 s_acc 的不同列 │
// │ 必须确保 4 个 warp 全部写完,才能: │
// │ (1) 通知 STORE warp 存 output (需要完整的 out_tile) │
// │ (2) 进入下一个 chunk (需要完整更新的 s_acc) │
// │ │
// │ compute_barrier.arrive_and_wait() = 所有 128 线程互相等待 │
// └─────────────────────────────────────────────────────────────────┘

// ┌─────────────────────────────────────────────────────────────────┐
// │ pipeline barrier (load / store) │
// │ │
// │ 用途: 不同角色 warp 之间的同步 (LOAD↔MMA, MMA↔STORE) │
// │ 时机: 每个 chunk 开头 (wait) 和结尾 (release/commit) │
// │ │
// │ 和 compute_barrier 的区别: │
// │ compute_barrier: warp 0-3 之间同步 (同角色) │
// │ pipeline barrier: 不同角色之间同步 (LOAD vs MMA vs STORE) │
// └─────────────────────────────────────────────────────────────────┘

时间线:

1
2
3
4
5
consumer_wait(input)       ← 等 LOAD (跨角色)
Phase 1-6 ← 4 warp 各自计算
compute_barrier ← 等 4 warp 都完成 (同角色)
producer_commit(output) ← 通知 STORE (跨角色)
consumer_release(input) ← 通知 LOAD (跨角色)

十一、fence_view_async_shared 的作用

代码 (行 732):

1
cutlass::arch::fence_view_async_shared();

放在 compute_barrier 之后,producer_commit 之前。

为什么需要: Phase 5 用 STSM 写 out_tile (寄存器→smem) Phase 6 用 STSM_T 写 s_acc (寄存器→smem)

STSM 是异步指令: 发出后不保证立刻写到 smem 的全局可见视图。 compute_barrier 只保证本 warp 的 STSM 已发出,不保证其他 warp 可见。

fence_view_async_shared: 确保当前线程之前的所有 STSM 对其他 warp/线程可见。 STORE warp 读 smem 时一定能看到最新的数据。

顺序:

1
2
3
4
5
STSM (Phase 5, 6)
→ compute_barrier (4 warp 同步,确保都写完)
→ fence (确保写入全局可见)
→ producer_commit (通知 STORE)
→ STORE warp 读 smem (安全)

十二、Shared Memory 的生命周期管理

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
┌─────────────────────────────────────────────────────────────┐
│ state_acc [128×128] bf16 = 32KB │
│ │
│ 生命周期: 整个 kernel,从头到尾常驻 │
│ 不受 pipeline 管理 │
│ Phase 1 读, Phase 6 写, chunk 之间保留 │
└─────────────────────────────────────────────────────────────┘

┌─────────────────────────────────────────────────────────────┐
│ input[0], input[1], input[2] 各 ~14KB │
│ │
│ 由 input pipeline 管理 (3-stage 轮换) │
│ LOAD 写入 → MMA 消费 → release → LOAD 覆盖 │
│ │
│ 每个 stage 包含: │
│ v[16×128], beta[32], kd[16×128], qd[16×128], │
│ kr[16×128], g_total[128], INV[16×16], Mqk[16×16] │
└─────────────────────────────────────────────────────────────┘

┌─────────────────────────────────────────────────────────────┐
│ output[0], output[1] 各 4KB │
│ │
│ 由 output pipeline 管理 (2-stage 轮换) │
│ MMA 写入 → STORE 消费 → release → MMA 覆盖 │
│ │
│ 每个 stage 包含: │
│ out[16×128] bf16 │
└─────────────────────────────────────────────────────────────┘

┌─────────────────────────────────────────────────────────────┐
│ union: input + output 与 state_fp32_buf 共用物理空间 │
│ │
│ 主循环前: state_fp32_buf 用于 fp32→bf16 转换 (如果需要) │
│ 主循环中: input[3] + output[2] 用于 pipeline │
│ 主循环后: state_fp32_buf 用于 bf16→fp32 转换 (如果需要) │
│ 不会同时使用,所以安全 │
└─────────────────────────────────────────────────────────────┘

十三、完整 API 调用时序图

以一个 chunk (t=2) 为例,展示所有 pipeline API 调用:

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
LOAD warp (warp4)           MMA warps (warp0-3)          STORE warp (warp5)
═══════════════ ═════════════════ ═══════════════════
│ │ │
│ acquire(stg2) ✓ │ │
│ get_barrier(stg2) │ │
│ │ │
│ TMA×8 → input[2] │ │ wait(stg0) ← 等 C0 的 output
│ (异步,绑 barrier) │ │
│ │ │ ... TMA store C0 的 output ...
│ ++load_write (→stg0) │ acquire(out_stg) ✓ │
│ │ wait(stg2) ← 等 TMA 完 │ release(stg0)
│ acquire(stg0) │ ↓ 阻塞 │ ++out_read
│ ← 等 MMA release stg0 │ │
│ ↓ 阻塞 │ ↓ TMA 完成! │
│ │ ↓ 通过 │
│ │ │
│ │ ── Phase 1-6 计算 ── │ wait(stg1) ← 等 C1 的 output
│ │ (读 input[2]) │
│ │ (写 output[out_stg]) │
│ │ │
│ │ compute_barrier ← 4warp同步│
│ │ fence │
│ │ commit(out_stg) ───────────→ wait 通过!
│ │ release(stg2) ─────→ │
│ │ ++load_read, ++out_write │
│ ← acquire 通过! │ │
│ TMA×8 → input[0] │ │ TMA store C2 的 output
│ ... │ │ ...

附录(二)

一、为什么需要 Prefetch

GPU 上两类操作的延迟差异:

  • LDSM (smem → 寄存器): ~20 cycle
  • LDSM_T (smem → 寄存器): ~20 cycle
  • MMA (矩阵乘): ~16 cycle
  • ALU (逐元素运算): ~5-10 cycle

如果串行执行: LDSM 等 20 cycle → MMA 等 16 cycle → LDSM 等 20 cycle → … 大量时间浪费在等待上

Prefetch 的做法: 先发 LDSM 加载"下一步"的数据,然后立刻执行"当前步"的 MMA/ALU。 LDSM 和 MMA/ALU 在硬件上使用不同的功能单元,可以并行执行。 等 MMA 做完,LDSM 的结果也差不多到了。

关键前提: LDSM 用的是 Load/Store 单元,MMA 用的是 Tensor Core, ALU 用的是 CUDA Core,三者互不干扰。

二、K2 中的四层 Prefetch / 缓冲

从离计算最远 (global memory) 到最近 (寄存器),有四层缓冲:

级别 缓冲深度 producer consumer 每份大小 隐藏的延迟
TMA input 3-stage LOAD warp MMA warps ~14 KB(smem) global→smem (~500c)
TMA output 2-stage MMA warps STORE warp 4 KB(smem) smem→global (~500c)
Phase1 B0/B1 2 份 LDSM gemm 8 bf16(reg) smem→reg (~20c)
Phase6 ring 1 份 LDSM_T gemm+ALU 8-16 bf16 smem→reg (~20c)

越靠近 global memory 延迟越大,需要的缓冲深度越深。 越靠近寄存器延迟越小,单缓冲 + 代码排序就够了。

三、TMA 3-stage Input Pipeline —— 三缓冲

3.1 结构

1
2
3
4
5
6
smem input[0]:  buffer A  ─┐
smem input[1]: buffer B ├─ 三份物理空间,每份 ~14KB
smem input[2]: buffer C ─┘

LOAD warp (warp4): producer,TMA 写入 smem
MMA warps (warp0-3): consumer,从 smem 读取计算

3.2 为什么三缓冲而不是双缓冲

双缓冲只能领先 1 步:

1
2
3
4
5
双缓冲:
LOAD: | L0→buf0 | L1→buf1 | 等release | L2→buf0 | ...
MMA: | 等L0 | C0←buf0 | C1←buf1 | 等L1完 | ...
↑ MMA算C1时LOAD想写,但buf0还没释放
如果MMA比LOAD慢,LOAD会频繁等待

三缓冲多一步余量:

1
2
3
4
5
三缓冲:
LOAD: | L0→stg0 | L1→stg1 | L2→stg2 | 等release | L3→stg0 |
MMA: | 等L0 | C0←stg0 | C1←stg1 | C2←stg2 | C3←stg0 |
↑ LOAD写stg2时,MMA在读stg0,不冲突
LOAD连发3个TMA不停顿

核心收益: LOAD 能连续发 3 个 TMA 不停顿,前几个 chunk 的加载延迟被完全隐藏。

3.3 同步机制

1
2
3
4
5
6
7
8
9
10
LOAD warp:
producer_acquire(load_write) // 等 MMA release → stage 可写
TMA_load(... → input[stage]) // 异步搬运,绑 barrier
++load_write // 推进 stage: 0→1→2→0→...

MMA warps:
consumer_wait(load_read) // 等 LOAD 的 barrier → 数据就绪
... Phase 1-6 计算 ...
consumer_release(load_read) // 通知 LOAD: 这个 stage 用完了
++load_read

stage 0 的生命周期: L0 写入 → C0 消费 → release → L3 写入 → C3 消费 → release → L6 写入 → …

3.4 Output Pipeline (2-stage) 也是多缓冲

1
2
3
4
smem output[0], output[1]: 双缓冲

MMA warps: producer,STSM 写 output
STORE warp: consumer,TMA 存到 global

只需双缓冲: MMA 写一个 stage 时 STORE 读另一个,足够隐藏 TMA store 延迟。

四、Phase 1 的 Prefetch —— 双列块交替 (寄存器级双缓冲)

4.1 背景

Phase 1: u_acc = kd[16×128] @ s_acc[128×128] (每 warp 取 32 列)

沿 K 维度分 8 步 (K_BLOCKS=8),每步:

  • A = kd[:, k*16:(k+1)*16] 16×16
  • B = s_acc 的两个列块 (B0, B1) 各 16×16

每步 4 条 gemm (kd@B0, qd@B0, kd@B1, qd@B1)

问题: B fragment 只有一份寄存器 (tCrBi/tCrB),怎么处理 B0 和 B1?

4.2 双缓冲设计

1
2
tCrBi = load buffer   (LDSM 写入目标)
tCrB = compute buffer (gemm 读取源)

两个变量交替: 一个在被 LDSM 写入,一个在被 gemm 读取。 transform(tCrBi → tCrB) 做"交接"。

4.3 完整展开 (以 k=0 为例)

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
═══ 循环前: 预取第 0 步的 A 和 B0 ═══

LDSM: kd[:, 0:16] → tCrAi_k 发出,等待
LDSM: qd[:, 0:16] → tCrAi_q 发出,等待
LDSM: s_acc[warp列0, K=0] → tCrBi 发出,等待
(循环前没有 gemm 可重叠,只能等)

═══ k=0 进入循环 ═══

① transform: tCrAi_k → tCrA_k A(kd) 就绪
tCrAi_q → tCrA_q A(qd) 就绪
tCrBi → tCrB B0 就绪 (上面预取的)

② LDSM: s_acc[warp列1, K=0] → tCrBi 发出! 加载 B1
此时 tCrBi 正在被写入,但 tCrB 已经是 B0 的副本,不冲突

③ gemm: u_acc[0] += tCrA_k @ tCrB 用 B0 算列块0 ──┐
gemm: out_acc[0] += tCrA_q @ tCrB 用 B0 算列块0 ├ 4条mma执行时
(4 条 mma 指令) │ ②的LDSM在并行完成!

④ transform: tCrBi → tCrB B1 就绪 (②加载完了)

⑤ LDSM: kd[:, 16:32] → tCrAi_k 发出! 加载下一步的 A(kd)
LDSM: qd[:, 16:32] → tCrAi_q 发出! 加载下一步的 A(qd)
LDSM: s_acc[warp列0, K=1] → tCrBi 发出! 加载下一步的 B0

⑥ gemm: u_acc[1] += tCrA_k @ tCrB 用 B1 算列块1 ──┐
gemm: out_acc[1] += tCrA_q @ tCrB 用 B1 算列块1 ├ ⑤的LDSM在并行完成!
(4 条 mma 指令) ┘

═══ k=1 进入循环,⑤预取的数据已就绪 ═══

① transform ... (⑤的结果直接可用)
...

4.4 寄存器状态变化表

tCrAi_k (load) tCrA_k (计算) tCrAi_q (load) tCrA_q (计算) tCrBi (load) tCrB (计算)
循环前: kd[K=0] qd[K=0] B0[K=0]
↑ 预取 ↑ 预取 ↑ 预取
① trans: ────→ kd[K=0] ────→ qd[K=0] ────→ B0[K=0]
可用gemm 可用gemm 可用gemm
② LDSM: B1[K=0]
↑加载中 B0仍可用
③ gemm: kd @ B0 qd @ B0 B0 被读取
④ trans: ────→ B1[K=0]
可用gemm
⑤ LDSM: kd[K=1] qd[K=1] B0[K=1]
↑加载中 ↑加载中 ↑加载中 B1仍可用
⑥ gemm: kd @ B1 qd @ B1 B1 被读取
← ⑤加载完成 下步①可用

4.5 为什么 B 有 B0/B1 交替而 A 没有

每步 k:

  • A (kd, qd): 同一个 16×16 块,被 B0 和 B1 两次 gemm 共享
  • B (s_acc): 两个不同的 16×16 块 (列块0 和 列块1)

A 在一步内不变,只在 k→k+1 时更新 B 在一步内要换两次 (B0→B1)

所以 B 需要 “用 B0 时预取 B1,用 B1 时预取下步 B0” 而 A 只需要 “用完后预取下步”

4.6 指令级并行的时间线

1
2
3
4
5
6
7
8
9
时间 →
┌──────────────┬──────────────┬───────────────┬──────────────┐
│ ② prefetch B1│ ③ gemm u[0] │ ⑤ prefetch │ ⑥ gemm u[1] │
│ (LDSM异步) │ gemm o[0] │ A, B0 (LDSM) │ gemm o[1] │
│ │ (用B0) │ (异步) │ (用B1) │
└──────────────┴──────────────┴───────────────┴──────────────┘
Load/Store单元: ▓▓▓▓▓▓▓▓▓▓▓▓ ▓▓▓▓▓▓▓▓▓▓▓▓▓▓
Tensor Core: ▓▓▓▓▓▓▓▓▓▓▓▓ ▓▓▓▓▓▓▓▓▓▓▓▓
← 并行 → ← 并行 → ← 并行 → ← 并行 →

五、Phase 6 的 Prefetch —— 单缓冲 + 精确排序

5.1 背景

Phase 6: s_acc = s_acc * g + kr^T @ U

8 步循环 (m=0…7),每步更新 s_acc 的 16 行。 三种数据需要 prefetch:

  • ring_A_kr: kr^T 的行块 [16×16] LDSM_T 从 smem
  • ring_S_acc: s_acc 的行块 [16×16] LDSM_T 从 smem
  • ring_g0/g1: g_total 标量 从 smem 读

5.2 PREFETCH=1 意味着什么

PREFETCH=1 → ring buffer 深度为 1 ring_A_kr[0], ring_S_acc[bi][0], ring_g0[0], ring_g1[0]

slot = m % 1 = 0 (永远是 0)

没有 “环形” 效果,就是单个 buffer,靠代码顺序保证 “先读后写”。 如果 PREFETCH=2,就会有 ring[0] 和 ring[1] 交替,变成真正的双缓冲。

5.3 完整展开 (m=0, m=1 为例)

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
═══ 循环前: 预取 m=0 的数据 ═══

LDSM_T: kr^T[0:16, 0:16] → tCrAi_kr → ring_A_kr[0]
LDSM_T: s_acc_T[0:16, warp列块0] → ring_S_acc[0][0]
LDSM_T: s_acc_T[0:16, warp列块1] → ring_S_acc[1][0]
标量读: g_total[0..7] → ring_g0[0]
g_total[8..15] → ring_g1[0]

═══ m=0 ═══

① 读 g:
g0 = ring_g0[0], g1 = ring_g1[0] m=0 的数据,预取好了
(拷贝到局部变量, ring_g 随时可覆盖)

② gemm (用 m=0 预取的 ring_A_kr):
clear(u_acc[0])
gemm(ring_A_kr[0], tCrB_u_arr[0], u_acc[0]) kr^T[行块0] @ U[列块0]
clear(u_acc[1])
gemm(ring_A_kr[0], tCrB_u_arr[1], u_acc[1]) kr^T[行块0] @ U[列块1]
↑ U 从 Phase4 保留,不需预取!
ring_A_kr 用完了 ✓ 可以覆盖
ring_g 已拷贝 ✓ 可以覆盖
ring_S_acc 还没用 ✗ 不能覆盖

③ 预取 A (m+1 的 kr^T 和 g):
LDSM_T: kr^T[16:32, 0:16] → ring_A_kr[0] 覆盖 m=0 的,安全
标量读: g_total[16..23] → ring_g0[0] 覆盖,安全
g_total[24..31] → ring_g1[0] 覆盖,安全
← 这些 LDSM 和下面 ④ 的 ALU 并行!

④ 逐元素更新 (用 m=0 预取的 ring_S_acc):
对 bi=0,1:
ring_S_acc[bi][0] = bf16(f32(ring_S_acc[bi][0]) * g0 + u_acc[bi])
ring_S_acc 用完了,值已更新

⑤ 写回 s_acc:
对 bi=0,1:
STSM_T: ring_S_acc[bi][0] → s_acc_T[行块0, warp列bi] 写回 m=0
ring_S_acc 写出了 ✓ 可以覆盖

⑥ 预取 C (m+1 的 s_acc):
对 bi=0,1:
LDSM_T: s_acc_T[行块1, warp列bi] → ring_S_acc[bi][0] 预取 m=1
覆盖,安全 (写的是m=0,读的是m=1,不同行)

═══ m=1 (③⑥预取的数据已就绪) ═══

① g0 = ring_g0[0], g1 = ring_g1[0] ③预取的 m=1 的 g
② gemm(ring_A_kr[0], ...) ③预取的 m=1 的 kr^T
...
④ 逐元素 ring_S_acc[bi][0] * g + gemm ⑥预取的 m=1 的 s_acc
...

5.4 为什么 ③ 和 ⑥ 不能合并到一起

如果把 ⑥ 移到 ③ 旁边:

1
2
3
4
② gemm 完成
③ 预取 kr^T[m+1]
⑥ 预取 s_acc[m+1] → ring_S_acc[bi][0] 覆盖了!
④ ring_S_acc[bi][0] * g + u_acc[bi] 读到的是 m+1 的数据,错了!

ring_S_acc 还没被 ④ 消费,就被 ⑥ 覆盖了。

每个变量的 “可覆盖” 时机:

变量 消费完毕的时刻 最早可覆盖 实际预取位置
ring_A_kr ② gemm 读完后 ③ (紧接②后) ③ ✓
ring_g0/g1 ① 已拷贝到局部变量 ③ (随时) ③ ✓
ring_S_acc ⑤ STSM_T 写回后 ⑥ (紧接⑤后) ⑥ ✓

5.5 ③ 也不能移到 ⑥ 旁边

正确性上可以 (ring_A_kr 在 ② 后就不需要了),但丢失重叠机会:

当前 (③在④前面):

1
2
3
4
5
6
7
8
时间 →
MMA: ┃ ② gemm ┃ ┃
LDSM: ┃ ┃③ kr^T ┃ ┃
ALU: ┃ ┃④ 逐元素 ┃ ┃
STSM: ┃ ┃ ┃⑤ 写回 ┃
LDSM: ┃ ┃ ┃⑥ s_acc ┃
^^^^^^^^^^
③和④并行! LDSM和ALU用不同硬件单元

如果 ③ 挪到 ⑥ 旁边:

1
2
3
4
5
6
7
8
时间 →
MMA: ┃ ② gemm ┃ ┃ ┃
ALU: ┃ ┃④ 逐元素 ┃ ┃
STSM: ┃ ┃ ┃⑤ 写回 ┃
LDSM: ┃ ┃ 空闲! ┃③ kr^T ┃⑥ s_acc ┃
^^^^^^^^ ^^^^^^^^^^^^^^^^^^^^^^^^
④期间LDSM ③⑥串行排队,下一步②要等更久
白白空闲

量化: LDSM_T ~20 cycle,④逐元素 ~15-20 cycle

  • 当前: ③ 的 20c 被 ④ 完全隐藏,免费
  • 挪后: 多出 ~20c 裸等,8 步共多 ~160 cycle / chunk

5.6 Phase 6 一步内的完整时间线

1
2
3
4
5
6
7
8
9
10
11
12
13
时间 →
┌──────────┬──────────────────────────────┬───────────────────┐
│ ② gemm │ ③ LDSM_T: kr^T[m+1] │ │
│ (Tensor │ ③ 标量读: g[m+1] │ │
│ Core) │ (与④的ALU并行) │ │
│ ├──────────────────────────────┤ │
│ │ ④ 逐元素: s = s*g + gemm │ │
│ │ (CUDA Core ALU) │ │
├──────────┴──────────────────────────────┤ │
│ ⑤ STSM_T: 写回 s_acc[m] │ │
│ ⑥ LDSM_T: 预取 s_acc[m+1] │ │
│ (写m行,读m+1行,不冲突) │ │
└─────────────────────────────────────────┴───────────────────┘

六、Prefetch 与双缓冲/三缓冲的关系

核心思想完全一样: 加载和计算重叠,隐藏延迟。 只是实现方式因场景而异。

6.1 经典双缓冲

1
2
3
4
5
6
buffer A, buffer B

步骤0: 加载 → buf A
步骤1: 加载 → buf B | 计算 ← buf A 交替
步骤2: 加载 → buf A | 计算 ← buf B 交替
步骤3: 加载 → buf B | 计算 ← buf A 交替

关键特征: 2 份存储空间,角色轮换,不同 buffer 永远不冲突。

6.2 K2 中各处的对应

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
┌─────────────────────────────────────────────────────────────────┐
│ TMA 3-stage input pipeline │
│ │
│ = 三缓冲 (双缓冲的扩展) │
│ 3 份 smem buffer, stage 0/1/2 轮换 │
│ LOAD 最多领先 MMA 2 步 │
│ 完全标准的多缓冲设计 │
└─────────────────────────────────────────────────────────────────┘

┌─────────────────────────────────────────────────────────────────┐
│ TMA 2-stage output pipeline │
│ │
│ = 双缓冲 │
│ 2 份 smem buffer, stage 0/1 轮换 │
│ MMA 最多领先 STORE 1 步 │
└─────────────────────────────────────────────────────────────────┘

┌─────────────────────────────────────────────────────────────────┐
│ Phase 1 的 tCrBi / tCrB │
│ │
│ = 寄存器级双缓冲 │
│ 2 份寄存器 fragment, load/compute 交替 │
│ transform 做交接 │
└─────────────────────────────────────────────────────────────────┘

┌─────────────────────────────────────────────────────────────────┐
│ Phase 6 的 ring buffer (PREFETCH=1) │
│ │
│ = 退化的双缓冲 / 单缓冲 + 精确排序 │
│ 只有 1 个 slot,没有交替,靠 "先读后写" 保证正确 │
│ 如果 PREFETCH=2 就变成真正的双缓冲 │
│ │
│ 为什么选 PREFETCH=1: │
│ 寄存器压力。Phase 6 同时存活的寄存器变量已经很多: │
│ tCrB_u_arr[2]: 16 bf16 (U的B格式, Phase4保留) │
│ ring_A_kr[P]: 8 bf16 × P │
│ ring_S_acc[2][P]:16 bf16 × P │
│ ring_g0/g1[P]: 2 fp32 × P │
│ u_acc[2]: 8 fp32 │
│ │
│ PREFETCH=2 会让 ring 变量翻倍,可能导致 register spill │
│ (溢出到 local memory),反而更慢。 │
│ PREFETCH=1 靠精确排序也能隐藏大部分延迟, │
│ 是寄存器压力和延迟隐藏之间的平衡点。 │
└─────────────────────────────────────────────────────────────────┘

6.3 对比总结

经典双缓冲 Phase1 tCrBi/tCrB Phase6 PREFETCH=1
缓冲区数量 2 2 1
交替方式 A↔B 轮换 load↔compute 轮换 先读后写,原地覆盖
正确性保证 不同buffer不冲突 transform做交接 代码顺序保证
能隐藏的延迟 完整一步加载延迟 完整LDSM延迟 部分 (靠与ALU重叠)
额外存储开销 2×(两份fragment) 1× (无额外开销)
```
  • 标题: KDA源码剖析之FlashKDA(下)
  • 作者: 鱿鱼圈
  • 创建于 : 2026-06-05 23:59:00
  • 更新于 : 2026-06-14 16:01:41
  • 链接: https://yuyanqi.com/2026/06/05/KDA源码剖析之FlashKDA(下)/
  • 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。
评论
目录
KDA源码剖析之FlashKDA(下)