v4.3.4 update. (#2892)

This commit is contained in:
Junkai-Wu
2025-12-22 00:49:12 +08:00
committed by GitHub
parent 331e2f451c
commit 7f5fe3edf1
31 changed files with 839 additions and 240 deletions

View File

@@ -839,7 +839,7 @@ class HopperFusedMultiHeadAttentionForward:
)
s_max_layout = cute.make_layout(
cute.size(layout_acc_mn(pv_tiled_mma, acc_pv.layout), mode=[0])
cute.size(self.layout_acc_mn(pv_tiled_mma, acc_pv.layout), mode=[0])
)
s_max = cute.make_rmem_tensor_like(s_max_layout, self.qk_acc_dtype)
a_sum = cute.make_rmem_tensor_like(s_max, cutlass.Float32)
@@ -888,7 +888,7 @@ class HopperFusedMultiHeadAttentionForward:
# MMA QK
cute.nvgpu.warpgroup.fence()
gemm_zero_acc(
self.gemm_zero_acc(
qk_tiled_mma,
tSrQ[(None, None, None, q_handle.index)],
tSrK[(None, None, None, k_handle.index)],
@@ -901,7 +901,7 @@ class HopperFusedMultiHeadAttentionForward:
# Wait for the pipeline MMAs to drain
cute.nvgpu.warpgroup.wait_group(0)
s_max, a_sum = softmax_step(
s_max, a_sum = self.softmax_step(
True,
self.mask_type,
acc_qk,
@@ -919,7 +919,7 @@ class HopperFusedMultiHeadAttentionForward:
True,
)
acc_qk_fixed = make_acc_into_op(
acc_qk_fixed = self.make_acc_into_op(
acc_qk, pv_tiled_mma.tv_layout_A, self.q_dtype
)
@@ -928,7 +928,7 @@ class HopperFusedMultiHeadAttentionForward:
# MMA PV
cute.nvgpu.warpgroup.fence()
gemm_zero_acc(
self.gemm_zero_acc(
pv_tiled_mma,
acc_qk_fixed,
tOrV[(None, None, None, v_handle.index)],
@@ -1040,7 +1040,7 @@ class HopperFusedMultiHeadAttentionForward:
cute.nvgpu.warpgroup.wait_group(0)
# acc_pv updated
lse = tail(
lse = self.tail(
s_max, a_sum, acc_pv, pv_tiled_mma, scale_softmax, scale_output
)
@@ -1077,10 +1077,10 @@ class HopperFusedMultiHeadAttentionForward:
if tOcO[0][1] == 0:
tOgLSE_mn = cute.make_tensor(
tOgLSE.iterator, layout_acc_mn(pv_tiled_mma, tOgLSE.layout)
tOgLSE.iterator, self.layout_acc_mn(pv_tiled_mma, tOgLSE.layout)
)
tOcO_mn = cute.make_tensor(
tOcO.iterator, layout_acc_mn(pv_tiled_mma, tOcO.layout)
tOcO.iterator, self.layout_acc_mn(pv_tiled_mma, tOcO.layout)
)
for i in cutlass.range_constexpr(cute.size(tOgLSE_mn, mode=[0])):
if (
@@ -1241,7 +1241,7 @@ class HopperFusedMultiHeadAttentionForward:
# MMA QK
cute.nvgpu.warpgroup.fence()
gemm_zero_acc(
self.gemm_zero_acc(
qk_tiled_mma,
tSrQ[(None, None, None, q_handle.index)],
tSrK[(None, None, None, k_handle.index)],
@@ -1255,7 +1255,7 @@ class HopperFusedMultiHeadAttentionForward:
# Wait for the pipeline MMAs to drain
cute.nvgpu.warpgroup.wait_group(0)
s_max, a_sum = softmax_step(
s_max, a_sum = self.softmax_step(
fusion,
self.mask_type,
acc_qk,
@@ -1272,7 +1272,7 @@ class HopperFusedMultiHeadAttentionForward:
window_size_right,
)
acc_qk_fixed = make_acc_into_op(
acc_qk_fixed = self.make_acc_into_op(
acc_qk, pv_tiled_mma.tv_layout_A, self.q_dtype
)
@@ -1300,6 +1300,7 @@ class HopperFusedMultiHeadAttentionForward:
@cute.jit
def softmax_step(
self,
fusion: bool,
mask_type: fmha_utils.MaskEnum,
acc_qk: cute.ThrMma,
@@ -1328,10 +1329,10 @@ class HopperFusedMultiHeadAttentionForward:
)
acc_qk_mn = cute.make_tensor(
acc_qk.iterator, layout_acc_mn(tiled_mma_qk, acc_qk.layout)
acc_qk.iterator, self.layout_acc_mn(tiled_mma_qk, acc_qk.layout)
)
reduction_target_qk = reduction_target_n(tiled_mma_qk)
reduction_target_qk = self.reduction_target_n(tiled_mma_qk)
red_rank = cute.rank(reduction_target_qk)
s_max_prev = None
@@ -1346,7 +1347,7 @@ class HopperFusedMultiHeadAttentionForward:
s_max[i] = cute.arch.fmax(s_max[i], acc_qk_mn[i, j])
else:
acc_pv_mn = cute.make_tensor(
acc_pv.iterator, layout_acc_mn(tiled_mma_pv, acc_pv.layout)
acc_pv.iterator, self.layout_acc_mn(tiled_mma_pv, acc_pv.layout)
)
s_max_prev = cute.make_rmem_tensor_like(s_max, s_max._dtype)
@@ -1396,15 +1397,15 @@ class HopperFusedMultiHeadAttentionForward:
return s_max, a_sum
@cute.jit
def reduction_target_n(tiled_mma):
separated = layout_separate(
def reduction_target_n(self, tiled_mma):
separated = self.layout_separate(
tiled_mma.shape_mnk[0],
cute.make_layout(tiled_mma.tv_layout_C.shape[0]),
tiled_mma.tv_layout_C.stride[0],
)
return separated[1]
@cute.jit
@staticmethod
def convert_c_layout_to_a_layout(c, a):
return cute.make_layout(
(a, c.shape[1], (c.shape[2], cute.size(c, mode=[0]) // cute.size(a))),
@@ -1416,9 +1417,9 @@ class HopperFusedMultiHeadAttentionForward:
)
@cute.jit
def make_acc_into_op(acc, operand_layout_tv, Element):
def make_acc_into_op(self, acc, operand_layout_tv, Element):
operand = cute.make_rmem_tensor_like(
convert_c_layout_to_a_layout(acc.layout, operand_layout_tv.shape[1]),
self.convert_c_layout_to_a_layout(acc.layout, operand_layout_tv.shape[1]),
Element,
)
operand_as_acc = cute.make_tensor(operand.iterator, acc.layout)
@@ -1499,7 +1500,7 @@ class HopperFusedMultiHeadAttentionForward:
return operand
@cute.jit
def tail(s_max, a_sum, acc_pv, tiled_mma_pv, scale_softmax, scale_output):
def tail(self, s_max, a_sum, acc_pv, tiled_mma_pv, scale_softmax, scale_output):
"""
Final processing step for FMHA that computes log-sum-exp (LSE) and scales the output.
@@ -1527,9 +1528,9 @@ class HopperFusedMultiHeadAttentionForward:
"""
# Create tensor view of accumulated P*V values with M*N layout
acc_pv_mn = cute.make_tensor(
acc_pv.iterator, layout_acc_mn(tiled_mma_pv, acc_pv.layout)
acc_pv.iterator, self.layout_acc_mn(tiled_mma_pv, acc_pv.layout)
)
reduction_target = reduction_target_n(tiled_mma_pv)
reduction_target = self.reduction_target_n(tiled_mma_pv)
red_rank = cute.rank(reduction_target)
for r in cutlass.range_constexpr(red_rank):
for i in cutlass.range_constexpr(cute.size(acc_pv_mn, mode=[0])):
@@ -1538,7 +1539,7 @@ class HopperFusedMultiHeadAttentionForward:
)
acc_mn = cute.make_tensor(
acc_pv.iterator, layout_acc_mn(tiled_mma_pv, acc_pv.layout)
acc_pv.iterator, self.layout_acc_mn(tiled_mma_pv, acc_pv.layout)
)
lse = cute.make_rmem_tensor_like(a_sum, a_sum._dtype)
@@ -1559,7 +1560,7 @@ class HopperFusedMultiHeadAttentionForward:
return lse
@cute.jit
@staticmethod
def layout_separate(thr, src, ref):
lt = cute.make_layout(())
ge = cute.make_layout(())
@@ -1577,6 +1578,7 @@ class HopperFusedMultiHeadAttentionForward:
r = cute.append(cute.append(cute.make_layout(()), lt), ge)
return r
@staticmethod
@cute.jit
def gemm_zero_acc(tiled_mma, A, B, C):
rA = cute.rank(A)
@@ -1606,8 +1608,8 @@ class HopperFusedMultiHeadAttentionForward:
assert 0
@cute.jit
def layout_acc_mn(tiled_mma, acc):
separated = layout_separate(
def layout_acc_mn(self, tiled_mma, acc):
separated = self.layout_separate(
tiled_mma.shape_mnk[0], acc[0], tiled_mma.tv_layout_C.stride[1]
)