diff --git a/cula/ops/kda/sm90/cp/pre_scan.py b/cula/ops/kda/sm90/cp/pre_scan.py index 900b333..0e233e3 100644 --- a/cula/ops/kda/sm90/cp/pre_scan.py +++ b/cula/ops/kda/sm90/cp/pre_scan.py @@ -48,16 +48,14 @@ def pre_scan_kernel( tma_atom_v: cute.CopyAtom, tma_tensor_v: cute.Tensor, - tma_atom_kd: cute.CopyAtom, - tma_tensor_kd: cute.Tensor, - tma_atom_kr: cute.CopyAtom, - tma_tensor_kr: cute.Tensor, tma_atom_inv: cute.CopyAtom, tma_tensor_inv: cute.Tensor, tma_atom_gt: cute.CopyAtom, tma_tensor_gt: cute.Tensor, tma_atom_beta: cute.CopyAtom, tma_tensor_beta: cute.Tensor, + ws_kd: cute.Tensor, + ws_kr: cute.Tensor, H: cutlass.Constexpr[int], total_tiles: cutlass.Int32, T_total: cutlass.Int32, @@ -118,22 +116,6 @@ def pre_scan_kernel( cute.group_modes(sV, 0, 2), cute.group_modes(gSrc_v, 0, 2), ) - gSrc_kd = cute.local_tile(tma_tensor_kd, (CHUNK, D), (None, None, None)) - tKDs, tKDg = cpasync.tma_partition( - tma_atom_kd, - 0, - cute.make_layout(1), - cute.group_modes(sKd, 0, 2), - cute.group_modes(gSrc_kd, 0, 2), - ) - gSrc_kr = cute.local_tile(tma_tensor_kr, (CHUNK, D), (None, None, None)) - tKRs, tKRg = cpasync.tma_partition( - tma_atom_kr, - 0, - cute.make_layout(1), - cute.group_modes(sKr, 0, 2), - cute.group_modes(gSrc_kr, 0, 2), - ) gSrc_inv = cute.local_tile(tma_tensor_inv, (CHUNK, CHUNK), (None, None, None)) tIs, tIg = cpasync.tma_partition( tma_atom_inv, @@ -158,6 +140,19 @@ def pre_scan_kernel( cute.group_modes(sBeta, 0, 2), cute.group_modes(gSrc_beta, 0, 2), ) + raw_copy_atom = cute.make_copy_atom(cpasync.CopyBulkG2SOp(), cutlass.BFloat16) + raw_stage_layout = cute.make_layout( + (CHUNK * D, STAGES), + stride=(1, CHUNK * D), + ) + raw_gmem_layout = cute.make_layout( + (CHUNK * D, total_tiles * H), + stride=(1, CHUNK * D), + ) + sKD_raw = cute.make_tensor(sKd.iterator, raw_stage_layout) + sKR_raw = cute.make_tensor(sKr.iterator, raw_stage_layout) + gKD_raw = cute.make_tensor(ws_kd.iterator, raw_gmem_layout) + gKR_raw = cute.make_tensor(ws_kr.iterator, raw_gmem_layout) # sState=0, sM=I if tidx < D: @@ -273,8 +268,8 @@ def pre_scan_kernel( cute.copy(tma_atom_v, tVg_seq[(None, t, 0, head_idx)], tVs_seq[(None, s_dyn_l)], tma_bar_ptr=bar_l) else: cute.copy(tma_atom_v, tVg[(None, tg_l, 0, head_idx)], tVs[(None, s_dyn_l)], tma_bar_ptr=bar_l) - cute.copy(tma_atom_kd, tKDg[(None, 0, 0, wt_l)], tKDs[(None, s_dyn_l)], tma_bar_ptr=bar_l) - cute.copy(tma_atom_kr, tKRg[(None, 0, 0, wt_l)], tKRs[(None, s_dyn_l)], tma_bar_ptr=bar_l) + cute.copy(raw_copy_atom, gKD_raw[(None, wt_l)], sKD_raw[(None, s_dyn_l)], mbar_ptr=bar_l) + cute.copy(raw_copy_atom, gKR_raw[(None, wt_l)], sKR_raw[(None, s_dyn_l)], mbar_ptr=bar_l) cute.copy(tma_atom_inv, tIg[(None, 0, 0, wt_l)], tIs[(None, s_dyn_l)], tma_bar_ptr=bar_l) cute.copy(tma_atom_gt, tGTg[(None, 0, 0, wt_l)], tGTs[(None, s_dyn_l)], tma_bar_ptr=bar_l) cute.copy(tma_atom_beta, tBg[(None, 0, 0, wt_l)], tBs[(None, s_dyn_l)], tma_bar_ptr=bar_l) @@ -487,13 +482,6 @@ def make_thd_atom(t, op): ) return cpasync.make_tiled_tma_atom(op, view, kinter_smem, (CHUNK, D)) - def make_ws_qkd_atom(t): - view = cute.make_tensor( - t.iterator, - cute.make_layout((CHUNK, D, total_tiles * H), stride=(D, 1, CHUNK * D)), - ) - return cpasync.make_tiled_tma_atom(cpasync.CopyBulkTensorTileG2SOp(), view, kinter_smem, (CHUNK, D)) - def make_ws_cc_atom(t): view = cute.make_tensor( t.iterator, @@ -502,8 +490,6 @@ def make_ws_cc_atom(t): return cpasync.make_tiled_tma_atom(cpasync.CopyBulkTensorTileG2SOp(), view, cc_smem, (CHUNK, CHUNK)) tma_atom_v, tma_tensor_v = make_thd_atom(v, cpasync.CopyBulkTensorTileG2SOp()) - tma_atom_kd, tma_tensor_kd = make_ws_qkd_atom(ws_kd) - tma_atom_kr, tma_tensor_kr = make_ws_qkd_atom(ws_kr) tma_atom_inv, tma_tensor_inv = make_ws_cc_atom(ws_inv) gt_smem = cute.make_layout((D, 1), stride=(1, D)) @@ -551,16 +537,14 @@ def make_beta_atom(t): pre_scan_kernel( tma_atom_v, tma_tensor_v, - tma_atom_kd, - tma_tensor_kd, - tma_atom_kr, - tma_tensor_kr, tma_atom_inv, tma_tensor_inv, tma_atom_gt, tma_tensor_gt, tma_atom_beta, tma_tensor_beta, + ws_kd, + ws_kr, H, total_tiles, T_total, diff --git a/cula/ops/kda/sm90/k1.py b/cula/ops/kda/sm90/k1.py index 8b2fdd1..56f2ca9 100644 --- a/cula/ops/kda/sm90/k1.py +++ b/cula/ops/kda/sm90/k1.py @@ -45,16 +45,13 @@ def k1_kernel( tma_tensor_k: cute.Tensor, tma_atom_g: cute.CopyAtom, tma_tensor_g: cute.Tensor, - tma_atom_ws_qd: cute.CopyAtom, - tma_tensor_ws_qd: cute.Tensor, - tma_atom_ws_kd: cute.CopyAtom, - tma_tensor_ws_kd: cute.Tensor, - tma_atom_ws_kr: cute.CopyAtom, - tma_tensor_ws_kr: cute.Tensor, tma_atom_ws_inv: cute.CopyAtom, tma_tensor_ws_inv: cute.Tensor, tma_atom_ws_mqk: cute.CopyAtom, tma_tensor_ws_mqk: cute.Tensor, + ws_qd: cute.Tensor, + ws_kd: cute.Tensor, + ws_kr: cute.Tensor, a_log: cute.Tensor, dt_bias: cute.Tensor, beta: cute.Tensor, @@ -148,30 +145,6 @@ def k1_kernel( cute.group_modes(gSrc_g, 0, 2), ) - gDst_qd = cute.local_tile(tma_tensor_ws_qd, (CHUNK, D), (None, None, None)) - tQDws_s, tQDws_g = cpasync.tma_partition( - tma_atom_ws_qd, - 0, - cute.make_layout(1), - cute.group_modes(s_q_decayed, 0, 2), - cute.group_modes(gDst_qd, 0, 2), - ) - gDst_kd = cute.local_tile(tma_tensor_ws_kd, (CHUNK, D), (None, None, None)) - tKDws_s, tKDws_g = cpasync.tma_partition( - tma_atom_ws_kd, - 0, - cute.make_layout(1), - cute.group_modes(s_k_decayed, 0, 2), - cute.group_modes(gDst_kd, 0, 2), - ) - gDst_kr = cute.local_tile(tma_tensor_ws_kr, (CHUNK, D), (None, None, None)) - tKRws_s, tKRws_g = cpasync.tma_partition( - tma_atom_ws_kr, - 0, - cute.make_layout(1), - cute.group_modes(s_k_restored, 0, 2), - cute.group_modes(gDst_kr, 0, 2), - ) gDst_inv = cute.local_tile(tma_tensor_ws_inv, (CHUNK, CHUNK), (None, None, None)) tINVws_s, tINVws_g = cpasync.tma_partition( tma_atom_ws_inv, @@ -189,6 +162,18 @@ def k1_kernel( cute.group_modes(gDst_mqk, 0, 2), ) ws_slot = head_idx * total_tiles + tile_idx + raw_copy_atom = cute.make_copy_atom(cpasync.CopyBulkS2GOp(), cutlass.BFloat16) + raw_smem_layout = cute.make_layout((CHUNK * D,), stride=(1,)) + raw_gmem_layout = cute.make_layout( + (CHUNK * D, total_tiles * H), + stride=(1, CHUNK * D), + ) + sQDws_raw = cute.make_tensor(s_q_decayed.iterator, raw_smem_layout) + sKDws_raw = cute.make_tensor(s_k_decayed.iterator, raw_smem_layout) + sKRws_raw = cute.make_tensor(s_k_restored.iterator, raw_smem_layout) + gQDws_raw = cute.make_tensor(ws_qd.iterator, raw_gmem_layout) + gKDws_raw = cute.make_tensor(ws_kd.iterator, raw_gmem_layout) + gKRws_raw = cute.make_tensor(ws_kr.iterator, raw_gmem_layout) if warp_idx == 0: with cute.arch.elect_one(): @@ -533,12 +518,15 @@ def k1_kernel( cute.arch.fence_view_async_shared() cute.arch.barrier() - # TMA bulk store all 5 workspace tensors (one elect_one, one thread). + # Preserve the physical K_INTER byte image for qd/kd/kr; inv/mqk remain + # layout-aware tensor TMA stores. This raw-workspace transport idea comes + # from Flash-Flash-KDA: https://github.com/Itssshikhar/Flash-Flash-KDA if warp_idx == 0: + # CuTeDSL 4.6 elects the issuing lane inside direct bulk-copy atoms. + cute.copy(raw_copy_atom, sQDws_raw, gQDws_raw[(None, ws_slot)]) + cute.copy(raw_copy_atom, sKDws_raw, gKDws_raw[(None, ws_slot)]) + cute.copy(raw_copy_atom, sKRws_raw, gKRws_raw[(None, ws_slot)]) with cute.arch.elect_one(): - cute.copy(tma_atom_ws_qd, tQDws_s[(None,)], tQDws_g[(None, 0, 0, ws_slot)]) - cute.copy(tma_atom_ws_kd, tKDws_s[(None,)], tKDws_g[(None, 0, 0, ws_slot)]) - cute.copy(tma_atom_ws_kr, tKRws_s[(None,)], tKRws_g[(None, 0, 0, ws_slot)]) cute.copy(tma_atom_ws_inv, tINVws_s[(None,)], tINVws_g[(None, 0, 0, ws_slot)]) cute.copy(tma_atom_ws_mqk, tMQKws_s[(None,)], tMQKws_g[(None, 0, 0, ws_slot)]) cute.arch.cp_async_bulk_commit_group() @@ -571,9 +559,6 @@ def run_k1( stream: cuda_drv.CUstream, ): smem_layout_qk = cute.make_layout((CHUNK, D), stride=(D, 1)) - # K_INTER swizzled layout — must match kernel SMEM layout for TMA stores. - kinter_atom = make_smem_layout_atom(SmemLayoutAtomKind.K_INTER, cutlass.BFloat16) - smem_layout_qk_kinter = cute.tile_to_shape(kinter_atom, (CHUNK, D), order=(0, 1)) def make_atom(t): view = cute.make_tensor( @@ -587,18 +572,6 @@ def make_atom(t): (CHUNK, D), ) - def make_ws_store_atom(t): - view = cute.make_tensor( - t.iterator, - cute.make_layout((CHUNK, D, total_tiles * H), stride=(D, 1, CHUNK * D)), - ) - return cpasync.make_tiled_tma_atom( - cpasync.CopyBulkTensorTileS2GOp(), - view, - smem_layout_qk_kinter, - (CHUNK, D), - ) - # (CHUNK, CHUNK) bf16 plain layout for ws_inv / ws_mqk TMA bulk store. smem_layout_cc = cute.make_layout((CHUNK, CHUNK), stride=(CHUNK, 1)) @@ -620,9 +593,6 @@ def make_ws_cc_store_atom(t): tma_atom_q, tma_tensor_q = make_atom(q) tma_atom_k, tma_tensor_k = make_atom(k) tma_atom_g, tma_tensor_g = make_atom(g) - tma_atom_ws_qd, tma_tensor_ws_qd = make_ws_store_atom(ws_qd) - tma_atom_ws_kd, tma_tensor_ws_kd = make_ws_store_atom(ws_kd) - tma_atom_ws_kr, tma_tensor_ws_kr = make_ws_store_atom(ws_kr) tma_atom_ws_inv, tma_tensor_ws_inv = make_ws_cc_store_atom(ws_inv) tma_atom_ws_mqk, tma_tensor_ws_mqk = make_ws_cc_store_atom(ws_mqk) @@ -635,16 +605,13 @@ def make_ws_cc_store_atom(t): tma_tensor_k, tma_atom_g, tma_tensor_g, - tma_atom_ws_qd, - tma_tensor_ws_qd, - tma_atom_ws_kd, - tma_tensor_ws_kd, - tma_atom_ws_kr, - tma_tensor_ws_kr, tma_atom_ws_inv, tma_tensor_ws_inv, tma_atom_ws_mqk, tma_tensor_ws_mqk, + ws_qd, + ws_kd, + ws_kr, a_log, dt_bias, beta, diff --git a/cula/ops/kda/sm90/k2.py b/cula/ops/kda/sm90/k2.py index 12ccbe3..7395984 100644 --- a/cula/ops/kda/sm90/k2.py +++ b/cula/ops/kda/sm90/k2.py @@ -63,16 +63,13 @@ def _make_out_kinter_one_stage(): def k2_kernel( tma_atom_v: cute.CopyAtom, tma_tensor_v: cute.Tensor, - tma_atom_kd: cute.CopyAtom, - tma_tensor_kd: cute.Tensor, - tma_atom_qd: cute.CopyAtom, - tma_tensor_qd: cute.Tensor, - tma_atom_kr: cute.CopyAtom, - tma_tensor_kr: cute.Tensor, tma_atom_inv: cute.CopyAtom, tma_tensor_inv: cute.Tensor, tma_atom_mqk: cute.CopyAtom, tma_tensor_mqk: cute.Tensor, + ws_qd: cute.Tensor, + ws_kd: cute.Tensor, + ws_kr: cute.Tensor, tma_atom_out: cute.CopyAtom, tma_tensor_out: cute.Tensor, out_gmem: cute.Tensor, @@ -156,30 +153,6 @@ def k2_kernel( cute.group_modes(sV, 0, 2), cute.group_modes(gSrc_v, 0, 2), ) - gSrc_kd = cute.local_tile(tma_tensor_kd, (CHUNK, D), (None, None, None)) - tKDs, tKDg = cpasync.tma_partition( - tma_atom_kd, - 0, - cute.make_layout(1), - cute.group_modes(sKd, 0, 2), - cute.group_modes(gSrc_kd, 0, 2), - ) - gSrc_qd = cute.local_tile(tma_tensor_qd, (CHUNK, D), (None, None, None)) - tQDs, tQDg = cpasync.tma_partition( - tma_atom_qd, - 0, - cute.make_layout(1), - cute.group_modes(sQd, 0, 2), - cute.group_modes(gSrc_qd, 0, 2), - ) - gSrc_kr = cute.local_tile(tma_tensor_kr, (CHUNK, D), (None, None, None)) - tKRs, tKRg = cpasync.tma_partition( - tma_atom_kr, - 0, - cute.make_layout(1), - cute.group_modes(sKr, 0, 2), - cute.group_modes(gSrc_kr, 0, 2), - ) gSrc_inv = cute.local_tile(tma_tensor_inv, (CHUNK, CHUNK), (None, None, None)) tIs, tIg = cpasync.tma_partition( tma_atom_inv, @@ -220,6 +193,21 @@ def k2_kernel( cute.group_modes(sBeta, 0, 2), cute.group_modes(gSrc_beta, 0, 2), ) + raw_copy_atom = cute.make_copy_atom(cpasync.CopyBulkG2SOp(), cutlass.BFloat16) + raw_stage_layout = cute.make_layout( + (CHUNK * D, STAGES), + stride=(1, CHUNK * D), + ) + raw_gmem_layout = cute.make_layout( + (CHUNK * D, total_tiles * H), + stride=(1, CHUNK * D), + ) + sKD_raw = cute.make_tensor(sKd.iterator, raw_stage_layout) + sQD_raw = cute.make_tensor(sQd.iterator, raw_stage_layout) + sKR_raw = cute.make_tensor(sKr.iterator, raw_stage_layout) + gKD_raw = cute.make_tensor(ws_kd.iterator, raw_gmem_layout) + gQD_raw = cute.make_tensor(ws_qd.iterator, raw_gmem_layout) + gKR_raw = cute.make_tensor(ws_kr.iterator, raw_gmem_layout) # Load initial_state -> sState[K_in, D_out] if has_initial_state: @@ -360,9 +348,11 @@ def k2_kernel( cute.copy(tma_atom_v, tVg_seq[(None, t, 0, head_idx)], tVs_seq[(None, s_dyn_l)], tma_bar_ptr=bar_l) else: cute.copy(tma_atom_v, tVg[(None, tg_l, 0, head_idx)], tVs[(None, s_dyn_l)], tma_bar_ptr=bar_l) - cute.copy(tma_atom_kd, tKDg[(None, 0, 0, wt_l)], tKDs[(None, s_dyn_l)], tma_bar_ptr=bar_l) - cute.copy(tma_atom_qd, tQDg[(None, 0, 0, wt_l)], tQDs[(None, s_dyn_l)], tma_bar_ptr=bar_l) - cute.copy(tma_atom_kr, tKRg[(None, 0, 0, wt_l)], tKRs[(None, s_dyn_l)], tma_bar_ptr=bar_l) + # Restore the byte-identical K_INTER images produced by K1. The + # raw-workspace transport idea is credited there to Flash-Flash-KDA. + cute.copy(raw_copy_atom, gKD_raw[(None, wt_l)], sKD_raw[(None, s_dyn_l)], mbar_ptr=bar_l) + cute.copy(raw_copy_atom, gQD_raw[(None, wt_l)], sQD_raw[(None, s_dyn_l)], mbar_ptr=bar_l) + cute.copy(raw_copy_atom, gKR_raw[(None, wt_l)], sKR_raw[(None, s_dyn_l)], mbar_ptr=bar_l) cute.copy(tma_atom_inv, tIg[(None, 0, 0, wt_l)], tIs[(None, s_dyn_l)], tma_bar_ptr=bar_l) cute.copy(tma_atom_mqk, tMg[(None, 0, 0, wt_l)], tMs[(None, s_dyn_l)], tma_bar_ptr=bar_l) cute.copy(tma_atom_gt, tGTg[(None, 0, 0, wt_l)], tGTs[(None, s_dyn_l)], tma_bar_ptr=bar_l) @@ -633,13 +623,6 @@ def make_thd_atom(t, op, t_total: cutlass.Int32): ) return cpasync.make_tiled_tma_atom(op, view, kinter_smem, (CHUNK, D)) - def make_ws_qkd_atom(t): - view = cute.make_tensor( - t.iterator, - cute.make_layout((CHUNK, D, total_tiles * H), stride=(D, 1, CHUNK * D)), - ) - return cpasync.make_tiled_tma_atom(cpasync.CopyBulkTensorTileG2SOp(), view, kinter_smem, (CHUNK, D)) - def make_ws_cc_atom(t): view = cute.make_tensor( t.iterator, @@ -649,9 +632,6 @@ def make_ws_cc_atom(t): tma_atom_v, tma_tensor_v = make_thd_atom(v, cpasync.CopyBulkTensorTileG2SOp(), V_T_total) tma_atom_out, tma_tensor_out = make_thd_atom(out, cpasync.CopyBulkTensorTileS2GOp(), O_T_total) - tma_atom_kd, tma_tensor_kd = make_ws_qkd_atom(ws_kd) - tma_atom_qd, tma_tensor_qd = make_ws_qkd_atom(ws_qd) - tma_atom_kr, tma_tensor_kr = make_ws_qkd_atom(ws_kr) tma_atom_inv, tma_tensor_inv = make_ws_cc_atom(ws_inv) tma_atom_mqk, tma_tensor_mqk = make_ws_cc_atom(ws_mqk) @@ -701,16 +681,13 @@ def make_beta_atom(t): k2_kernel( tma_atom_v, tma_tensor_v, - tma_atom_kd, - tma_tensor_kd, - tma_atom_qd, - tma_tensor_qd, - tma_atom_kr, - tma_tensor_kr, tma_atom_inv, tma_tensor_inv, tma_atom_mqk, tma_tensor_mqk, + ws_qd, + ws_kd, + ws_kr, tma_atom_out, tma_tensor_out, out,