Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
54 changes: 19 additions & 35 deletions cula/ops/kda/sm90/cp/pre_scan.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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:
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand All @@ -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))
Expand Down Expand Up @@ -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,
Expand Down
83 changes: 25 additions & 58 deletions cula/ops/kda/sm90/k1.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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():
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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(
Expand All @@ -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))

Expand All @@ -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)

Expand All @@ -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,
Expand Down
75 changes: 26 additions & 49 deletions cula/ops/kda/sm90/k2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand All @@ -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)

Expand Down Expand Up @@ -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,
Expand Down