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:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user