v4.4 tag release update. (#3032)
This commit is contained in:
@@ -33,9 +33,6 @@ import argparse
|
||||
from typing import List, Type, Tuple, Optional
|
||||
import cuda.bindings.driver as cuda
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
import cutlass
|
||||
import cutlass.cute as cute
|
||||
import cutlass.cute.testing as testing
|
||||
@@ -43,7 +40,6 @@ import cutlass.utils as utils
|
||||
import cutlass.pipeline as pipeline
|
||||
from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait
|
||||
from cutlass.cute.nvgpu import cpasync, tcgen05
|
||||
import cutlass.torch as cutlass_torch
|
||||
import cutlass.utils.blackwell_helpers as sm100_utils
|
||||
from cutlass.cute.runtime import from_dlpack
|
||||
|
||||
@@ -702,7 +698,7 @@ class SSDKernel:
|
||||
G = cute.size(tma_tensor_b, mode=[3])
|
||||
NGROUP_RATIO = EH // G
|
||||
|
||||
# Make TiledMma
|
||||
# Make tiledMma
|
||||
(
|
||||
tiled_mma_intra1,
|
||||
tiled_mma_intra2,
|
||||
@@ -1670,7 +1666,10 @@ class SSDKernel:
|
||||
cute.copy(tiled_r2s_p, tRS_rP, tRS_sP[inter2_p_coord])
|
||||
|
||||
# Fence for shared memory
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
cute.arch.fence_proxy(
|
||||
"async.shared",
|
||||
space="cta",
|
||||
)
|
||||
# Async arrive INTER2_P buffer full
|
||||
inter2_p_pipeline.producer_commit(inter2_p_producer_state)
|
||||
# Advance INTER2_P producer state
|
||||
@@ -1700,7 +1699,10 @@ class SSDKernel:
|
||||
]
|
||||
|
||||
# Fence for shared memory
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
cute.arch.fence_proxy(
|
||||
"async.shared",
|
||||
space="cta",
|
||||
)
|
||||
|
||||
# Combine B/Delta/DeltaA/last_column
|
||||
tScaledB = self.pre_inter_scale_bt_with_delta(
|
||||
@@ -1716,7 +1718,10 @@ class SSDKernel:
|
||||
cute.copy(tiled_r2s_b, tBrB_r2s, tBsB_r2s[inter1_b_coord])
|
||||
|
||||
# Fence for shared memory
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
cute.arch.fence_proxy(
|
||||
"async.shared",
|
||||
space="cta",
|
||||
)
|
||||
|
||||
# Async arrive B/Delta/B_TMEM buffer empty/empty/full
|
||||
b_pipeline.consumer_release(
|
||||
@@ -1743,14 +1748,9 @@ class SSDKernel:
|
||||
|
||||
# Combine INTER1_ACC/last_column/State
|
||||
exp_last_column = cute.math.exp(last_column, fastmath=True)
|
||||
for reg_idx in range(0, cute.size(tTR_rP), 2):
|
||||
(
|
||||
tTR_rP[reg_idx],
|
||||
tTR_rP[reg_idx + 1],
|
||||
) = cute.arch.fma_packed_f32x2(
|
||||
(exp_last_column, exp_last_column),
|
||||
(tState[reg_idx], tState[reg_idx + 1]),
|
||||
(tTR_rP[reg_idx], tTR_rP[reg_idx + 1]),
|
||||
for reg_idx in cutlass.range(cute.size(tTR_rP), vectorize=True):
|
||||
tTR_rP[reg_idx] = (
|
||||
exp_last_column * tState[reg_idx] + tTR_rP[reg_idx]
|
||||
)
|
||||
|
||||
# Store scaled P to tRS_rP
|
||||
@@ -1765,7 +1765,10 @@ class SSDKernel:
|
||||
cute.copy(tiled_r2s_p, tRS_rP, tRS_sP[inter2_p_coord])
|
||||
|
||||
# Fence for shared memory
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
cute.arch.fence_proxy(
|
||||
"async.shared",
|
||||
space="cta",
|
||||
)
|
||||
|
||||
# Async arrive INTER1_ACC/INTER2_P buffer empty/full
|
||||
inter1_acc_pipeline.consumer_release(inter1_acc_consumer_state)
|
||||
@@ -1798,8 +1801,11 @@ class SSDKernel:
|
||||
# END of for chunk_idx in cutlass.range(C, unroll=1)
|
||||
|
||||
# Store last INTER2_P (State) from smem to gmem
|
||||
# Wait for all previous stores to smem to be done
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
# Wait for all previous store to smem done
|
||||
cute.arch.fence_proxy(
|
||||
"async.shared",
|
||||
space="cta",
|
||||
)
|
||||
self.pre_inter_sync_barrier.arrive_and_wait()
|
||||
|
||||
if local_warp_idx == 0:
|
||||
@@ -2245,58 +2251,26 @@ class SSDKernel:
|
||||
cute.copy(s2r_atom_d, tRS_sD[d_coord], tRS_rD)
|
||||
|
||||
# Combine INTRA2_ACC/INTER2_ACC/Delta/X/D
|
||||
for reg_idx in range(0, cute.size(tRS_rCompute), 2):
|
||||
(
|
||||
tRS_rCompute[reg_idx],
|
||||
tRS_rCompute[reg_idx + 1],
|
||||
) = cute.arch.fma_packed_f32x2(
|
||||
(tTR_rInter[reg_idx], tTR_rInter[reg_idx + 1]),
|
||||
(
|
||||
cute.math.exp(
|
||||
tTR_rDeltaA[reg_idx], fastmath=True
|
||||
),
|
||||
cute.math.exp(
|
||||
tTR_rDeltaA[reg_idx + 1], fastmath=True
|
||||
),
|
||||
),
|
||||
(tTR_rIntra[reg_idx], tTR_rIntra[reg_idx + 1]),
|
||||
for reg_idx in cutlass.range(
|
||||
cute.size(tRS_rCompute), vectorize=True
|
||||
):
|
||||
tRS_rCompute[reg_idx] = (
|
||||
tTR_rInter[reg_idx]
|
||||
* cute.math.exp(tTR_rDeltaA[reg_idx], fastmath=True)
|
||||
+ tTR_rIntra[reg_idx]
|
||||
)
|
||||
# Fuse Y += X * D
|
||||
if cutlass.const_expr(self.d_has_hdim):
|
||||
(
|
||||
tRS_rCompute[reg_idx],
|
||||
tRS_rCompute[reg_idx + 1],
|
||||
) = cute.arch.fma_packed_f32x2(
|
||||
(
|
||||
tRS_rD[reg_idx].to(self.acc_dtype),
|
||||
tRS_rD[reg_idx + 1].to(self.acc_dtype),
|
||||
),
|
||||
(
|
||||
tSR_rX[reg_idx].to(self.acc_dtype),
|
||||
tSR_rX[reg_idx + 1].to(self.acc_dtype),
|
||||
),
|
||||
(
|
||||
tRS_rCompute[reg_idx],
|
||||
tRS_rCompute[reg_idx + 1],
|
||||
),
|
||||
tRS_rCompute[reg_idx] = (
|
||||
tRS_rD[reg_idx].to(self.acc_dtype)
|
||||
* tSR_rX[reg_idx].to(self.acc_dtype)
|
||||
+ tRS_rCompute[reg_idx]
|
||||
)
|
||||
elif cutlass.const_expr(self.has_d):
|
||||
(
|
||||
tRS_rCompute[reg_idx],
|
||||
tRS_rCompute[reg_idx + 1],
|
||||
) = cute.arch.fma_packed_f32x2(
|
||||
(
|
||||
tRS_rD.to(self.acc_dtype),
|
||||
tRS_rD.to(self.acc_dtype),
|
||||
),
|
||||
(
|
||||
tSR_rX[reg_idx].to(self.acc_dtype),
|
||||
tSR_rX[reg_idx + 1].to(self.acc_dtype),
|
||||
),
|
||||
(
|
||||
tRS_rCompute[reg_idx],
|
||||
tRS_rCompute[reg_idx + 1],
|
||||
),
|
||||
tRS_rCompute[reg_idx] = (
|
||||
tRS_rD.to(self.acc_dtype)
|
||||
* tSR_rX[reg_idx].to(self.acc_dtype)
|
||||
+ tRS_rCompute[reg_idx]
|
||||
)
|
||||
|
||||
tRS_rY.store(tRS_rCompute.load().to(self.io_dtype))
|
||||
@@ -2309,7 +2283,10 @@ class SSDKernel:
|
||||
)
|
||||
|
||||
# Fence for R2S store
|
||||
cute.arch.fence_proxy("async.shared", space="cta")
|
||||
cute.arch.fence_proxy(
|
||||
"async.shared",
|
||||
space="cta",
|
||||
)
|
||||
# Sync before TMA store
|
||||
self.epilog_sync_barrier.arrive_and_wait()
|
||||
|
||||
@@ -2426,7 +2403,6 @@ class SSDKernel:
|
||||
internal_stages,
|
||||
intra1_acc_stages,
|
||||
):
|
||||
SM100_TMEM_CAPACITY_COLUMNS = 512
|
||||
BITS_PER_TMEM_COL = 32
|
||||
# (MMA, MMA_M, MMA_N)
|
||||
acc_shape_intra1 = tiled_mma_intra1.partition_shape_C(tile_shape_mnk_intra1[:2])
|
||||
@@ -2483,7 +2459,7 @@ class SSDKernel:
|
||||
num_tmem_cols_total = 1
|
||||
while num_tmem_cols_total < num_tmem_cols_total_tmp:
|
||||
num_tmem_cols_total *= 2
|
||||
assert num_tmem_cols_total <= SM100_TMEM_CAPACITY_COLUMNS
|
||||
assert num_tmem_cols_total <= cute.arch.get_max_tmem_alloc_cols("sm_100")
|
||||
|
||||
return (
|
||||
tmem_intra1_acc_offset,
|
||||
@@ -3036,41 +3012,26 @@ class SSDKernel:
|
||||
|
||||
# SegSum
|
||||
# fadd2 + fsel + fmul2/mufu + fmul2
|
||||
for subtile_idx in cutlass.range(0, cute.size(tTR_rQ), 2, unroll_full=True):
|
||||
(
|
||||
tCompute[subtile_idx],
|
||||
tCompute[subtile_idx + 1],
|
||||
) = cute.arch.add_packed_f32x2(
|
||||
(tCrDeltaA_Col[subtile_idx], tCrDeltaA_Col[subtile_idx + 1]),
|
||||
(-tCrDeltaA_Row[subtile_idx], -tCrDeltaA_Row[subtile_idx + 1]),
|
||||
for subtile_idx in cutlass.range(
|
||||
cute.size(tTR_rQ), unroll_full=True, vectorize=True
|
||||
):
|
||||
tCompute[subtile_idx] = tCrDeltaA_Col[subtile_idx] + (
|
||||
-tCrDeltaA_Row[subtile_idx]
|
||||
)
|
||||
for subtile_idx in cutlass.range(cute.size(tTR_rQ), unroll_full=True):
|
||||
m, n = tCoord[subtile_idx]
|
||||
if m < n:
|
||||
tCompute[subtile_idx] = cutlass.Float32(-float("inf"))
|
||||
LOG2_E = cutlass.Float32(1.4426950408889634)
|
||||
for subtile_idx in cutlass.range(0, cute.size(tTR_rQ), 2, unroll_full=True):
|
||||
for subtile_idx in cutlass.range(
|
||||
cute.size(tTR_rQ), unroll_full=True, vectorize=True
|
||||
):
|
||||
# TODO: use math.exp directly
|
||||
tCompute_log2e = cute.arch.mul_packed_f32x2(
|
||||
(tCompute[subtile_idx], tCompute[subtile_idx + 1]), (LOG2_E, LOG2_E)
|
||||
)
|
||||
(
|
||||
tCompute[subtile_idx],
|
||||
tCompute[subtile_idx + 1],
|
||||
) = cute.arch.mul_packed_f32x2(
|
||||
(
|
||||
cute.math.exp2(tCompute_log2e[0], fastmath=True),
|
||||
cute.math.exp2(tCompute_log2e[1], fastmath=True),
|
||||
),
|
||||
(tCrDelta[subtile_idx], tCrDelta[subtile_idx + 1]),
|
||||
)
|
||||
(
|
||||
tCompute[subtile_idx],
|
||||
tCompute[subtile_idx + 1],
|
||||
) = cute.arch.mul_packed_f32x2(
|
||||
(tCompute[subtile_idx], tCompute[subtile_idx + 1]),
|
||||
(tTR_rQ[subtile_idx], tTR_rQ[subtile_idx + 1]),
|
||||
tCompute_log2e = tCompute[subtile_idx] * LOG2_E
|
||||
tCompute[subtile_idx] = (
|
||||
cute.math.exp2(tCompute_log2e, fastmath=True) * tCrDelta[subtile_idx]
|
||||
)
|
||||
tCompute[subtile_idx] = tCompute[subtile_idx] * tTR_rQ[subtile_idx]
|
||||
|
||||
tRT_rQ.store(tCompute.load().to(self.io_dtype))
|
||||
return tRT_rQ
|
||||
@@ -3211,6 +3172,7 @@ class SSDKernel:
|
||||
)
|
||||
return sDeltaA
|
||||
|
||||
@cute.jit
|
||||
def pre_inter_scale_bt_with_delta(
|
||||
self, tBrB_s2r, tBrDelta_s2r, tBrDeltaA_s2r, last_column
|
||||
):
|
||||
@@ -3223,22 +3185,15 @@ class SSDKernel:
|
||||
tBrDelta_Compute.store(tBrDelta_s2r.load().to(self.acc_dtype))
|
||||
tBrDeltaA_Compute.store(tBrDeltaA_s2r.load().to(self.acc_dtype))
|
||||
|
||||
for reg_idx in range(0, cute.size(tBrB_Compute), 2):
|
||||
tCompute[reg_idx], tCompute[reg_idx + 1] = cute.arch.mul_packed_f32x2(
|
||||
(
|
||||
cute.math.exp(
|
||||
(last_column - tBrDeltaA_Compute[reg_idx]), fastmath=True
|
||||
),
|
||||
cute.math.exp(
|
||||
(last_column - tBrDeltaA_Compute[reg_idx + 1]), fastmath=True
|
||||
),
|
||||
),
|
||||
(tBrDelta_Compute[reg_idx], tBrDelta_Compute[reg_idx + 1]),
|
||||
)
|
||||
tCompute[reg_idx], tCompute[reg_idx + 1] = cute.arch.mul_packed_f32x2(
|
||||
(tCompute[reg_idx], tCompute[reg_idx + 1]),
|
||||
(tBrB_Compute[reg_idx], tBrB_Compute[reg_idx + 1]),
|
||||
for reg_idx in cutlass.range(
|
||||
cute.size(tBrB_Compute), vectorize=True, unroll_full=True
|
||||
):
|
||||
tCompute[reg_idx] = (
|
||||
cute.math.exp((last_column - tBrDeltaA_Compute[reg_idx]), fastmath=True)
|
||||
* tBrDelta_Compute[reg_idx]
|
||||
)
|
||||
|
||||
tCompute[reg_idx] = tCompute[reg_idx] * tBrB_Compute[reg_idx]
|
||||
return tCompute
|
||||
|
||||
def epilog_make_delta(self, smem_cumsum_delta):
|
||||
@@ -3349,6 +3304,10 @@ def run(
|
||||
print(f"Skip reference checking: {skip_ref_check}")
|
||||
print(f"Use cold L2: {'True' if use_cold_l2 else 'False'}")
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import cutlass.torch as cutlass_torch
|
||||
|
||||
# Unpack parameters
|
||||
G, B, E, H, C, D, L, N = gbehcdln
|
||||
EH = E * H
|
||||
|
||||
Reference in New Issue
Block a user