cutlass 3.9 update (#2255)
* cutlass 3.9 update * rebase * fixes out of shared memory for blockwise Blackwell * doc format * fix issue 2253 * disable host ref by default * fix sm120 smem capacity --------- Co-authored-by: yuzhai <yuzhai@nvidia.com> Co-authored-by: Haicheng Wu <haichengw@nvidia.com>
This commit is contained in:
@@ -324,8 +324,12 @@ def is_complex(data_type):
|
||||
def is_block_scaled(gemm_kind):
|
||||
return gemm_kind in (GemmKind.BlockScaledUniversal3x, GemmKind.GroupedBlockScaledUniversal3x)
|
||||
|
||||
def is_blockwise(gemm_kind):
|
||||
return gemm_kind in (GemmKind.BlockwiseUniversal3x, GemmKind.GroupedBlockwiseUniversal3x)
|
||||
|
||||
def is_grouped(gemm_kind):
|
||||
return gemm_kind in (GemmKind.GroupedUniversal3x, GemmKind.GroupedBlockScaledUniversal3x)
|
||||
return gemm_kind in (GemmKind.GroupedUniversal3x,
|
||||
GemmKind.GroupedBlockScaledUniversal3x, GemmKind.GroupedBlockwiseUniversal3x)
|
||||
|
||||
#
|
||||
def get_complex_from_real(real_type):
|
||||
@@ -493,6 +497,9 @@ class KernelScheduleType(enum.Enum):
|
||||
PtrArrayTmaWarpSpecializedPingpong = enum_auto()
|
||||
PtrArrayTmaWarpSpecializedPingpongFP8FastAccum = enum_auto()
|
||||
|
||||
BlockwiseTmaWarpSpecializedCooperative = enum_auto()
|
||||
PtrArrayBlockwiseTmaWarpSpecializedCooperative = enum_auto()
|
||||
|
||||
TmaWarpSpecialized1SmSm100 = enum_auto()
|
||||
TmaWarpSpecialized2SmSm100 = enum_auto()
|
||||
ImplicitTmaWarpSpecialized1SmSm100 = enum_auto()
|
||||
@@ -518,6 +525,13 @@ class KernelScheduleType(enum.Enum):
|
||||
Mxf8f6f4TmaWarpSpecialized1SmSm100 = enum_auto()
|
||||
Mxf8f6f4TmaWarpSpecialized2SmSm100 = enum_auto()
|
||||
|
||||
BlockwiseTmaWarpSpecialized1SmSm100 = enum_auto()
|
||||
BlockwiseTmaWarpSpecialized2SmSm100 = enum_auto()
|
||||
|
||||
PtrArrayBlockwiseTmaWarpSpecialized1SmSm100 = enum_auto()
|
||||
PtrArrayBlockwiseTmaWarpSpecialized2SmSm100 = enum_auto()
|
||||
|
||||
|
||||
Mxf4TmaWarpSpecialized1SmSm100 = enum_auto()
|
||||
Mxf4TmaWarpSpecialized2SmSm100 = enum_auto()
|
||||
Nvf4TmaWarpSpecialized1SmSm100 = enum_auto()
|
||||
@@ -547,6 +561,8 @@ KernelScheduleTag = {
|
||||
KernelScheduleType.TmaWarpSpecializedPingpongFP8FastAccum: 'cutlass::gemm::KernelTmaWarpSpecializedPingpongFP8FastAccum',
|
||||
KernelScheduleType.ImplicitTmaWarpSpecializedSm90: 'cutlass::conv::KernelImplicitTmaWarpSpecializedSm90',
|
||||
|
||||
KernelScheduleType.BlockwiseTmaWarpSpecializedCooperative: 'cutlass::gemm::KernelTmaWarpSpecializedCooperativeFP8BlockScaledAccum',
|
||||
|
||||
KernelScheduleType.TmaWarpSpecialized1SmSm100: 'cutlass::gemm::KernelTmaWarpSpecialized1SmSm100',
|
||||
KernelScheduleType.TmaWarpSpecialized2SmSm100: 'cutlass::gemm::KernelTmaWarpSpecialized2SmSm100',
|
||||
|
||||
@@ -564,6 +580,12 @@ KernelScheduleTag = {
|
||||
KernelScheduleType.Mxf8f6f4TmaWarpSpecialized1SmSm100: 'cutlass::gemm::KernelTmaWarpSpecialized1SmMxf8f6f4Sm100',
|
||||
KernelScheduleType.Mxf8f6f4TmaWarpSpecialized2SmSm100: 'cutlass::gemm::KernelTmaWarpSpecialized2SmMxf8f6f4Sm100',
|
||||
|
||||
KernelScheduleType.BlockwiseTmaWarpSpecialized1SmSm100: 'cutlass::gemm::KernelTmaWarpSpecializedBlockwise1SmSm100',
|
||||
KernelScheduleType.BlockwiseTmaWarpSpecialized2SmSm100: 'cutlass::gemm::KernelTmaWarpSpecializedBlockwise2SmSm100',
|
||||
|
||||
KernelScheduleType.PtrArrayBlockwiseTmaWarpSpecialized1SmSm100: 'cutlass::gemm::KernelPtrArrayTmaWarpSpecializedBlockwise1SmSm100',
|
||||
KernelScheduleType.PtrArrayBlockwiseTmaWarpSpecialized2SmSm100: 'cutlass::gemm::KernelPtrArrayTmaWarpSpecializedBlockwise2SmSm100',
|
||||
|
||||
KernelScheduleType.Mxf4TmaWarpSpecialized1SmSm100: 'cutlass::gemm::KernelTmaWarpSpecialized1SmMxf4Sm100',
|
||||
KernelScheduleType.Mxf4TmaWarpSpecialized2SmSm100: 'cutlass::gemm::KernelTmaWarpSpecialized2SmMxf4Sm100',
|
||||
KernelScheduleType.Nvf4TmaWarpSpecialized1SmSm100: 'cutlass::gemm::KernelTmaWarpSpecialized1SmNvf4Sm100',
|
||||
@@ -574,6 +596,8 @@ KernelScheduleTag = {
|
||||
KernelScheduleType.PtrArrayTmaWarpSpecializedPingpong: 'cutlass::gemm::KernelPtrArrayTmaWarpSpecializedPingpong',
|
||||
KernelScheduleType.PtrArrayTmaWarpSpecializedPingpongFP8FastAccum: 'cutlass::gemm::KernelPtrArrayTmaWarpSpecializedPingpongFP8FastAccum',
|
||||
|
||||
KernelScheduleType.PtrArrayBlockwiseTmaWarpSpecializedCooperative: 'cutlass::gemm::KernelPtrArrayTmaWarpSpecializedCooperativeFP8BlockScaledAccum',
|
||||
|
||||
KernelScheduleType.PtrArrayTmaWarpSpecialized1SmBlockScaledSm100: "cutlass::gemm::KernelPtrArrayTmaWarpSpecialized1SmBlockScaledSm100",
|
||||
KernelScheduleType.PtrArrayTmaWarpSpecialized2SmBlockScaledSm100: "cutlass::gemm::KernelPtrArrayTmaWarpSpecialized2SmBlockScaledSm100",
|
||||
KernelScheduleType.PtrArrayNvf4TmaWarpSpecialized1SmSm100: "cutlass::gemm::KernelPtrArrayTmaWarpSpecialized1SmNvf4Sm100",
|
||||
@@ -608,6 +632,8 @@ KernelScheduleSuffixes = {
|
||||
KernelScheduleType.TmaWarpSpecializedCooperativeFP8FastAccum: '_warpspecialized_cooperative_fp8_fastaccum',
|
||||
KernelScheduleType.TmaWarpSpecializedPingpongFP8FastAccum: '_warpspecialized_pingpong_fp8_fastaccum',
|
||||
KernelScheduleType.ImplicitTmaWarpSpecializedSm90: '_warpspecialized',
|
||||
|
||||
KernelScheduleType.BlockwiseTmaWarpSpecializedCooperative: '_warpspecialized_cooperative',
|
||||
|
||||
KernelScheduleType.TmaWarpSpecialized1SmSm100: '_1sm',
|
||||
KernelScheduleType.TmaWarpSpecialized2SmSm100: '_2sm',
|
||||
@@ -626,6 +652,11 @@ KernelScheduleSuffixes = {
|
||||
KernelScheduleType.Mxf8f6f4TmaWarpSpecialized1SmSm100: '_q_1sm',
|
||||
KernelScheduleType.Mxf8f6f4TmaWarpSpecialized2SmSm100: '_q_2sm',
|
||||
|
||||
KernelScheduleType.BlockwiseTmaWarpSpecialized1SmSm100: '_1sm',
|
||||
KernelScheduleType.BlockwiseTmaWarpSpecialized2SmSm100: '_2sm',
|
||||
KernelScheduleType.PtrArrayBlockwiseTmaWarpSpecialized1SmSm100: '_1sm',
|
||||
KernelScheduleType.PtrArrayBlockwiseTmaWarpSpecialized2SmSm100: '_2sm',
|
||||
|
||||
KernelScheduleType.Mxf4TmaWarpSpecialized1SmSm100: '_o_vs32_1sm',
|
||||
KernelScheduleType.Mxf4TmaWarpSpecialized2SmSm100: '_o_vs32_2sm',
|
||||
KernelScheduleType.Nvf4TmaWarpSpecialized1SmSm100: '_o_vs16_1sm',
|
||||
@@ -636,6 +667,8 @@ KernelScheduleSuffixes = {
|
||||
KernelScheduleType.PtrArrayTmaWarpSpecializedPingpong: '_warpspecialized_pingpong',
|
||||
KernelScheduleType.PtrArrayTmaWarpSpecializedPingpongFP8FastAccum: '_warpspecialized_pingpong_fp8_fastaccum',
|
||||
|
||||
KernelScheduleType.PtrArrayBlockwiseTmaWarpSpecializedCooperative: '_warpspecialized_cooperative',
|
||||
|
||||
KernelScheduleType.PtrArrayTmaWarpSpecialized1SmBlockScaledSm100: '_1sm',
|
||||
KernelScheduleType.PtrArrayTmaWarpSpecialized2SmBlockScaledSm100: '_2sm',
|
||||
KernelScheduleType.PtrArrayNvf4TmaWarpSpecialized1SmSm100: '_o_vs16_1sm',
|
||||
@@ -730,6 +763,7 @@ def to_grouped_schedule(schedule, grouped):
|
||||
group_schedule_map = {
|
||||
# SM90
|
||||
KernelScheduleType.TmaWarpSpecializedCooperative : KernelScheduleType.PtrArrayTmaWarpSpecializedCooperative,
|
||||
KernelScheduleType.BlockwiseTmaWarpSpecializedCooperative : KernelScheduleType.PtrArrayBlockwiseTmaWarpSpecializedCooperative,
|
||||
KernelScheduleType.TmaWarpSpecializedPingpong : KernelScheduleType.PtrArrayTmaWarpSpecializedPingpong,
|
||||
KernelScheduleType.TmaWarpSpecializedCooperativeFP8FastAccum : KernelScheduleType.PtrArrayTmaWarpSpecializedCooperativeFP8FastAccum,
|
||||
KernelScheduleType.TmaWarpSpecializedPingpongFP8FastAccum : KernelScheduleType.PtrArrayTmaWarpSpecializedPingpongFP8FastAccum,
|
||||
@@ -745,6 +779,9 @@ def to_grouped_schedule(schedule, grouped):
|
||||
KernelScheduleType.Mxf8f6f4TmaWarpSpecialized2SmSm100 : KernelScheduleType.PtrArrayMxf8f6f4TmaWarpSpecialized2SmSm100,
|
||||
EpilogueScheduleType.TmaWarpSpecialized1Sm: EpilogueScheduleType.PtrArrayTmaWarpSpecialized1Sm,
|
||||
EpilogueScheduleType.TmaWarpSpecialized2Sm: EpilogueScheduleType.PtrArrayTmaWarpSpecialized2Sm,
|
||||
KernelScheduleType.BlockwiseTmaWarpSpecialized1SmSm100 : KernelScheduleType.PtrArrayBlockwiseTmaWarpSpecialized1SmSm100,
|
||||
KernelScheduleType.BlockwiseTmaWarpSpecialized2SmSm100 : KernelScheduleType.PtrArrayBlockwiseTmaWarpSpecialized2SmSm100,
|
||||
|
||||
}
|
||||
|
||||
return group_schedule_map[schedule]
|
||||
@@ -932,6 +969,8 @@ class GemmKind(enum.Enum):
|
||||
BlockScaledUniversal3x = enum_auto()
|
||||
GroupedUniversal3x = enum_auto()
|
||||
GroupedBlockScaledUniversal3x = enum_auto()
|
||||
BlockwiseUniversal3x = enum_auto()
|
||||
GroupedBlockwiseUniversal3x = enum_auto()
|
||||
|
||||
#
|
||||
GemmKindNames = {
|
||||
@@ -945,7 +984,9 @@ GemmKindNames = {
|
||||
GemmKind.Grouped: "gemm_grouped",
|
||||
GemmKind.BlockScaledUniversal3x: "gemm",
|
||||
GemmKind.GroupedUniversal3x: "gemm_grouped",
|
||||
GemmKind.GroupedBlockScaledUniversal3x: "gemm_grouped"
|
||||
GemmKind.GroupedBlockScaledUniversal3x: "gemm_grouped",
|
||||
GemmKind.BlockwiseUniversal3x: "gemm",
|
||||
GemmKind.GroupedBlockwiseUniversal3x: "gemm_grouped"
|
||||
}
|
||||
|
||||
#
|
||||
@@ -1149,7 +1190,7 @@ class MathInstruction:
|
||||
#
|
||||
class TileDescription:
|
||||
|
||||
def __init__(self, threadblock_shape, stages, warp_count, math_instruction, min_compute, max_compute, cluster_shape = [1,1,1]):
|
||||
def __init__(self, threadblock_shape, stages, warp_count, math_instruction, min_compute, max_compute, cluster_shape = [1,1,1], explicit_vector_sizes = None):
|
||||
self.threadblock_shape = threadblock_shape
|
||||
self.tile_shape = threadblock_shape
|
||||
self.stages = stages
|
||||
@@ -1158,6 +1199,7 @@ class TileDescription:
|
||||
self.minimum_compute_capability = min_compute
|
||||
self.maximum_compute_capability = max_compute
|
||||
self.cluster_shape = cluster_shape
|
||||
self.explicit_vector_sizes = explicit_vector_sizes
|
||||
|
||||
def procedural_name(self):
|
||||
if self.minimum_compute_capability >= 90:
|
||||
|
||||
Reference in New Issue
Block a user