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:
Yujia Zhai
2025-04-24 12:42:40 -07:00
committed by GitHub
parent 8e345c5c5b
commit 331a1f5b3f
143 changed files with 18089 additions and 5935 deletions

View File

@@ -219,6 +219,15 @@ def CreateGemmUniversal3xOperator(
gemm_op_extra_args["ScaleFactorD"] = { "tensor": TensorDescription(data_type["sfd_type"]["type"], data_type["sfd_type"]["layout"]),
"vector_size" : data_type["sfd_type"]["vector_size"]}
assert is_block_scaled(gemm_kind)
if tile_description.explicit_vector_sizes != None:
assert len(tile_description.explicit_vector_sizes) == 3
gemm_op_extra_args["ScaleFactorMVecSize"] = tile_description.explicit_vector_sizes[0]
gemm_op_extra_args["ScaleFactorNVecSize"] = tile_description.explicit_vector_sizes[1]
gemm_op_extra_args["ScaleFactorKVecSize"] = tile_description.explicit_vector_sizes[2]
assert is_blockwise(gemm_kind)
else:
assert not is_blockwise(gemm_kind)
A_dtype = data_type["a_type"]
B_dtype = data_type["b_type"]
@@ -5811,6 +5820,87 @@ def GenerateSM90_TensorOp_fp8_WGMMA_gemm(manifest, cuda_version, gemm_kind=GemmK
stream_k_schedules,
tile_schedulers=[TileSchedulerType.StreamK])
def GenerateSM90_TensorOp_fp8_WGMMA_gemm_with_blockwise(manifest, cuda_version, gemm_kind=GemmKind.BlockwiseUniversal3x):
if not CudaToolkitVersionSatisfies(cuda_version, 12, 3 if is_grouped(gemm_kind) else 0):
return
instantiation_level = manifest.get_sm90_instantiation_level(pruned_level=20, default_level=121, exhaustive_level=9992)
is_aligned = True
# layouts for ABC and their alignments
layouts = [
[[LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 1]], # TN Layout
]
math_instructions = generate_fp8_math_instructions_sm90(instantiation_level)
tile_descriptions_ = generate_tile_descriptions_sm90(
math_instructions=math_instructions,
is_aligned=is_aligned,
level=instantiation_level)
tile_descriptions = list()
for desc in tile_descriptions_:
desc.explicit_vector_sizes = [1, desc.tile_shape[1], desc.tile_shape[2]]
tile_descriptions.append(copy.deepcopy(desc))
desc.explicit_vector_sizes = [desc.tile_shape[0], desc.tile_shape[1], desc.tile_shape[2]]
tile_descriptions.append(copy.deepcopy(desc))
desc.explicit_vector_sizes = [desc.tile_shape[0], desc.tile_shape[1], desc.tile_shape[2]]
tile_descriptions.append(copy.deepcopy(desc))
desc.explicit_vector_sizes = [1, 1, desc.tile_shape[2]]
tile_descriptions.append(copy.deepcopy(desc))
for tile_desc in tile_descriptions:
math_inst = tile_desc.math_instruction
data_types = []
fp8_types = [DataType.e4m3, DataType.e5m2]
valid_types_for_d = [DataType.f32, DataType.bf16, DataType.f16, DataType.e4m3, DataType.e5m2]
valid_types_for_c = copy.deepcopy(valid_types_for_d)
valid_types_for_c.append(DataType.void)
for c_type, d_type in product(valid_types_for_c, valid_types_for_d):
data_types.append(
generate_data_types_from_math_instruction(
math_inst,
element_source=c_type,
element_dest=d_type,
)
)
else:
for d_type in valid_types_for_d:
data_types.append(
generate_data_types_from_math_instruction(
math_inst,
element_source=DataType.void,
element_dest=d_type,
)
)
for layout in layouts:
for data_type in data_types:
# Inconsistency: alignments aren't fixed in FP8
# layout = fix_alignments(data_type, layout, alignment_bits=128)
schedules, stream_k_schedules = get_valid_schedules(
tile_description=tile_desc,
cuda_version=cuda_version,
is_aligned=is_aligned,
data_types=data_type,
instantiation_level=instantiation_level,
layout=layout,
gemm_kind=gemm_kind,
enable_fp8_fast_acc=False,
)
if len(schedules):
CreateGemmUniversal3xOperator(manifest, [layout], [tile_desc], data_type, schedules, gemm_kind=gemm_kind)
if len(stream_k_schedules):
assert CudaToolkitVersionSatisfies(cuda_version, 12, 1)
CreateGemmUniversal3xOperator(manifest, [layout], [tile_desc], data_type,
stream_k_schedules,
tile_schedulers=[TileSchedulerType.StreamK],
gemm_kind=gemm_kind)
def GenerateSM90_TensorOp_fp8_WGMMA_alignx_gemm(manifest, cuda_version):
if not CudaToolkitVersionSatisfies(cuda_version, 12, 0):
@@ -7499,6 +7589,245 @@ def GenerateSM100_TensorOp_fp8_UMMA_gemm(manifest, cuda_version, gemm_kind=GemmK
CreateGemmUniversal3xOperator(manifest, layouts, tile_descriptions, data_type,
[[kernel_schedule, epi_schedule]], tile_schedulers=tile_schedulers, gemm_kind=gemm_kind)
def GenerateSM100_TensorOp_fp8_UMMA_gemm_with_blockwise(manifest, cuda_version, gemm_kind=GemmKind.BlockwiseUniversal3x):
if not CudaToolkitVersionSatisfies(cuda_version, 12, 8):
return
grouped = is_grouped(gemm_kind)
# layouts for ABC and their alignments.
layouts = [
[[LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 0]],
[[LayoutType.ColumnMajor, 16], [LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 0]],
[[LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 0]],
[[LayoutType.RowMajor, 16], [LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 0]],
[[LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.RowMajor, 0]],
[[LayoutType.ColumnMajor, 16], [LayoutType.RowMajor, 16], [LayoutType.RowMajor, 0]],
[[LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.RowMajor, 0]],
[[LayoutType.RowMajor, 16], [LayoutType.RowMajor, 16], [LayoutType.RowMajor, 0]],
]
min_cc = 100
max_cc = 100
epi_type = DataType.f32
math_instructions_1sm = [
# inst 64x128
MathInstruction(
[64, 128, 32],
DataType.f8, DataType.f8, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[64, 128, 32],
DataType.e4m3, DataType.e4m3, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[64, 128, 32],
DataType.e4m3, DataType.e5m2, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[64, 128, 32],
DataType.e5m2, DataType.e4m3, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
# inst 128x32
MathInstruction(
[128, 32, 32],
DataType.f8, DataType.f8, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[128, 32, 32],
DataType.e4m3, DataType.e4m3, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[128, 32, 32],
DataType.e4m3, DataType.e5m2, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[128, 32, 32],
DataType.e5m2, DataType.e4m3, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
# inst 128x64
MathInstruction(
[128, 64, 32],
DataType.f8, DataType.f8, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[128, 64, 32],
DataType.e4m3, DataType.e4m3, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[128, 64, 32],
DataType.e4m3, DataType.e5m2, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[128, 64, 32],
DataType.e5m2, DataType.e4m3, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
# inst 128x128
MathInstruction(
[128, 128, 32],
DataType.f8, DataType.f8, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[128, 128, 32],
DataType.e4m3, DataType.e4m3, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[128, 128, 32],
DataType.e4m3, DataType.e5m2, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[128, 128, 32],
DataType.e5m2, DataType.e4m3, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
# inst 128x256
MathInstruction(
[128, 256, 32],
DataType.e4m3, DataType.e4m3, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[128, 256, 32],
DataType.e4m3, DataType.e5m2, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[128, 256, 32],
DataType.e5m2, DataType.e4m3, DataType.f32,
OpcodeClass.TensorOp,
MathOperation.multiply_add)]
cluster_shapes_1sm = [[1,2,1], [2,1,1], [1,1,1], [1,4,1], [4,4,1]
, DynamicClusterShape
]
tile_schedulers = [
TileSchedulerType.Default,
]
# 1xSM MMA kernels
for math_inst in math_instructions_1sm:
tile_descriptions = []
for cluster_shape in cluster_shapes_1sm:
multiplier_1sm = (1, 1, 1) if cluster_shape == DynamicClusterShape else cluster_shape
tile_descriptions.append(
TileDescription([
math_inst.instruction_shape[0] * multiplier_1sm[0],
math_inst.instruction_shape[1] * multiplier_1sm[1],
math_inst.instruction_shape[2] * 4 * multiplier_1sm[2]],
0, [4, 1, 1], math_inst, min_cc, max_cc, cluster_shape,
[math_inst.instruction_shape[0], math_inst.instruction_shape[1],
math_inst.instruction_shape[2] * 4]))
tile_descriptions.append(
TileDescription([
math_inst.instruction_shape[0] * multiplier_1sm[0],
math_inst.instruction_shape[1] * multiplier_1sm[1],
math_inst.instruction_shape[2] * 4 * multiplier_1sm[2]],
0, [4, 1, 1], math_inst, min_cc, max_cc, cluster_shape,
[1, math_inst.instruction_shape[1],
math_inst.instruction_shape[2] * 4]))
tile_descriptions.append(
TileDescription([
math_inst.instruction_shape[0] * multiplier_1sm[0],
math_inst.instruction_shape[1] * multiplier_1sm[1],
math_inst.instruction_shape[2] * 4 * multiplier_1sm[2]],
0, [4, 1, 1], math_inst, min_cc, max_cc, cluster_shape,
[math_inst.instruction_shape[0], 1,
math_inst.instruction_shape[2] * 4]))
data_types = [
{
"a_type" : math_inst.element_a,
"b_type" : math_inst.element_b,
"c_type" : DataType.f16,
"d_type" : DataType.f16,
"acc_type" : math_inst.element_accumulator,
"epi_type" : epi_type,
},
{
"a_type" : math_inst.element_a,
"b_type" : math_inst.element_b,
"c_type" : DataType.bf16,
"d_type" : DataType.bf16,
"acc_type" : math_inst.element_accumulator,
"epi_type" : epi_type,
},
{
"a_type" : math_inst.element_a,
"b_type" : math_inst.element_b,
"c_type" : DataType.f32,
"d_type" : DataType.f32,
"acc_type" : math_inst.element_accumulator,
"epi_type" : epi_type,
},
{
"a_type" : math_inst.element_a,
"b_type" : math_inst.element_b,
"c_type" : DataType.void,
"d_type" : DataType.f16,
"acc_type" : math_inst.element_accumulator,
"epi_type" : epi_type,
},
{
"a_type" : math_inst.element_a,
"b_type" : math_inst.element_b,
"c_type" : DataType.void,
"d_type" : DataType.bf16,
"acc_type" : math_inst.element_accumulator,
"epi_type" : epi_type,
},
{
"a_type" : math_inst.element_a,
"b_type" : math_inst.element_b,
"c_type" : DataType.void,
"d_type" : DataType.f32,
"acc_type" : math_inst.element_accumulator,
"epi_type" : epi_type,
},
]
# Set alignment d based on Destination format.
for layout in layouts:
layout[2][1] = 128 // DataTypeSize[data_types[0]["d_type"]]
is_runtime_datatype = lambda runtime_datatype: runtime_datatype in (DataType.f4, DataType.f6, DataType.f8)
for data_type in data_types:
if ( data_type["a_type"] == DataType.e4m3 ) and ( data_type["b_type"] == DataType.e4m3 ) and\
( data_type["d_type"] == DataType.e5m2 ):
continue
is_runtime_datatype_a = is_runtime_datatype(data_type["a_type"])
is_runtime_datatype_b = is_runtime_datatype(data_type["d_type"])
# A/B datatypes should be both static or dynamic
if (is_runtime_datatype_a != is_runtime_datatype_b):
continue
# grouped GEMM does not support runtime data type yet
if grouped and (is_runtime_datatype_a or is_runtime_datatype_b):
continue
kernel_schedule = to_grouped_schedule(KernelScheduleType.BlockwiseTmaWarpSpecialized1SmSm100, grouped)
epi_schedule = to_grouped_schedule(EpilogueScheduleType.TmaWarpSpecialized1Sm, grouped)
CreateGemmUniversal3xOperator(manifest, layouts, tile_descriptions, data_type,
[[kernel_schedule, epi_schedule]],
tile_schedulers=tile_schedulers, gemm_kind=gemm_kind)
def GenerateSM100_TensorOp_mixed_8bits_UMMA_gemm(manifest, cuda_version):
# SM100 MMA with mixed F4/F6/F8 inputs + without block scale
if not CudaToolkitVersionSatisfies(cuda_version, 12, 0):
@@ -10318,6 +10647,11 @@ def GenerateSM100(manifest, cuda_version):
# StreamK is included in regular generation
GenerateSM100_TensorOp_mixed_8bits_UMMA_gemm(manifest, cuda_version)
# Blockwise kernels
GenerateSM100_TensorOp_fp8_UMMA_gemm_with_blockwise(manifest, cuda_version)
GenerateSM100_TensorOp_fp8_UMMA_gemm_with_blockwise(manifest, cuda_version, gemm_kind=GemmKind.GroupedBlockwiseUniversal3x)
#
# Sparse Gemm
#
@@ -10755,6 +11089,8 @@ def GenerateSM90(manifest, cuda_version):
GenerateSM90_SparseTensorOp_tf32_WGMMA_gemm(manifest, cuda_version)
GenerateSM90_SparseTensorOp_int8_WGMMA_gemm(manifest, cuda_version)
GenerateSM90_SparseTensorOp_fp8_WGMMA_gemm(manifest, cuda_version)
GenerateSM90_TensorOp_fp8_WGMMA_gemm_with_blockwise(manifest, cuda_version)
GenerateSM90_TensorOp_fp8_WGMMA_gemm_with_blockwise(manifest, cuda_version, gemm_kind=GemmKind.GroupedBlockwiseUniversal3x)
###################################################################################################
@@ -10819,6 +11155,8 @@ if __name__ == "__main__":
manifest = Manifest(args)
archs = args.architectures.split(';')
GenerateSM50(manifest, args.cuda_version)
GenerateSM60(manifest, args.cuda_version)
GenerateSM61(manifest, args.cuda_version)
@@ -10827,8 +11165,8 @@ if __name__ == "__main__":
GenerateSM80(manifest, args.cuda_version)
GenerateSM89(manifest, args.cuda_version)
GenerateSM90(manifest, args.cuda_version)
blackwell_enabled_arch = args.architectures in ["100a", "101a", "120a"]
blackwell_enabled_arch = any(arch in ["100a", "101a", "120a"] for arch in archs)
if blackwell_enabled_arch:
GenerateSM100(manifest, args.cuda_version)
GenerateSM120(manifest, args.cuda_version)