v3.8.0 update (#2082)

* 3.8 update

* fix Markus' name

---------

Co-authored-by: yuzhai <yuzhai@nvidia.com>
This commit is contained in:
Yujia Zhai
2025-02-06 18:33:40 -08:00
committed by GitHub
parent affd1b693d
commit 833f6990e0
168 changed files with 24945 additions and 3436 deletions

View File

@@ -90,9 +90,11 @@ try:
raise ImportError("Disabling attempt to import cutlass_library")
from cutlass_library.library import *
from cutlass_library.manifest import *
from cutlass_library.emit_kernel_listing import emit_gemm_kernel_testlist
except ImportError:
from library import *
from manifest import *
from emit_kernel_listing import emit_gemm_kernel_testlist
###################################################################################################
#
@@ -177,7 +179,8 @@ def CreateGemmUniversal3xOperator(
complex_transforms=None,
epilogue_functor=EpilogueFunctor.LinearCombination,
swizzling_functor=SwizzlingFunctor.Identity1,
tile_schedulers=[TileSchedulerType.Default]):
tile_schedulers=[TileSchedulerType.Default],
gemm_kind=GemmKind.Universal3x):
if type(data_types) is dict:
data_types = [data_types]
@@ -206,7 +209,6 @@ def CreateGemmUniversal3xOperator(
D = TensorDescription(data_type["d_type"], layout[2][0], layout[2][1])
gemm_op_extra_args = {}
gemm_kind = GemmKind.Universal3x
element_compute = data_type.get("epi_type", data_type["acc_type"])
@@ -218,16 +220,43 @@ def CreateGemmUniversal3xOperator(
gemm_kind = GemmKind.BlockScaledUniversal3x
operation = GemmOperation(
gemm_kind, tile_description.minimum_compute_capability,
tile_description, A, B, C, element_compute, epilogue_functor, swizzling_functor, D,
kernel_schedule, epilogue_schedule, tile_scheduler, **gemm_op_extra_args)
A_dtype = data_type["a_type"]
B_dtype = data_type["b_type"]
A_dtype_bits = DataTypeSize[A_dtype]
B_dtype_bits = DataTypeSize[B_dtype]
is_A_dtype_narrow = A_dtype_bits < B_dtype_bits
if is_A_dtype_narrow:
narrow_dtype, wide_dtype = (A_dtype, B_dtype)
narrow_dtype_bits, wide_dtype_bits = (A_dtype_bits, B_dtype_bits)
else:
narrow_dtype, wide_dtype = (B_dtype, A_dtype)
narrow_dtype_bits, wide_dtype_bits = (B_dtype_bits, A_dtype_bits)
manifest.append(operation)
operations.append(operation)
mixed_input_modes = [None]
if narrow_dtype_bits != wide_dtype_bits:
if narrow_dtype == DataType.s4 and (wide_dtype == DataType.e4m3 or wide_dtype == DataType.e5m2):
mixed_input_modes = [MixedInputMode.ScaleOnly]
else:
mixed_input_modes = [MixedInputMode.ConvertOnly, MixedInputMode.ScaleOnly, MixedInputMode.ScaleWithZeroPoint]
mixed_input_shuffle_options = [False]
if (mixed_input_modes[0] is not None) and (wide_dtype_bits == 16) and (narrow_dtype_bits == 4 or narrow_dtype_bits == 8):
mixed_input_shuffle_options = [False, True]
for mixed_input_mode, mixed_input_shuffle in product(mixed_input_modes, mixed_input_shuffle_options):
operation = GemmOperation(
gemm_kind, tile_description.minimum_compute_capability,
tile_description, A, B, C, element_compute, epilogue_functor, swizzling_functor, D,
kernel_schedule, epilogue_schedule, tile_scheduler,
mixed_input_mode=mixed_input_mode, mixed_input_shuffle=mixed_input_shuffle, **gemm_op_extra_args)
manifest.append(operation)
operations.append(operation)
return operations
def is_grouped(gemm_kind):
return gemm_kind == GemmKind.GroupedGemmUniversal3x
# Generates 3.0 API based GemmUniversal API kernels. Alignment constraints are folded in with layouts
def CreateSparseGemmUniversal3xOperator(
manifest, layouts, tile_descriptions, data_types,
@@ -4934,12 +4963,7 @@ def GenerateSM80(manifest, cuda_version):
###################################################################################################
def GenerateSM89_TensorOp_16832_fp8(manifest, cuda_version):
if (
not CudaToolkitVersionSatisfies(cuda_version, 12, 4)
):
return
def GenerateSM89_TensorOp_16832_fp8(manifest, element_acc):
layouts = [
(LayoutType.RowMajor, LayoutType.ColumnMajor, LayoutType.ColumnMajor),
(LayoutType.RowMajor, LayoutType.ColumnMajor, LayoutType.RowMajor)
@@ -4948,49 +4972,48 @@ def GenerateSM89_TensorOp_16832_fp8(manifest, cuda_version):
math_instructions = [
MathInstruction(
[16, 8, 32],
DataType.e4m3, DataType.e4m3, DataType.f32,
DataType.e4m3, DataType.e4m3, element_acc,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[16, 8, 32],
DataType.e4m3, DataType.e5m2, DataType.f32,
DataType.e4m3, DataType.e5m2, element_acc,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[16, 8, 32],
DataType.e5m2, DataType.e4m3, DataType.f32,
DataType.e5m2, DataType.e4m3, element_acc,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[16, 8, 32],
DataType.e5m2, DataType.e5m2, DataType.f32,
DataType.e5m2, DataType.e5m2, element_acc,
OpcodeClass.TensorOp,
MathOperation.multiply_add),
MathInstruction(
[16, 8, 32],
DataType.e4m3, DataType.e4m3, DataType.f32,
DataType.e4m3, DataType.e4m3, element_acc,
OpcodeClass.TensorOp,
MathOperation.multiply_add_fast_accum),
MathInstruction(
[16, 8, 32],
DataType.e4m3, DataType.e5m2, DataType.f32,
DataType.e4m3, DataType.e5m2, element_acc,
OpcodeClass.TensorOp,
MathOperation.multiply_add_fast_accum),
MathInstruction(
[16, 8, 32],
DataType.e5m2, DataType.e4m3, DataType.f32,
DataType.e5m2, DataType.e4m3, element_acc,
OpcodeClass.TensorOp,
MathOperation.multiply_add_fast_accum),
MathInstruction(
[16, 8, 32],
DataType.e5m2, DataType.e5m2, DataType.f32,
DataType.e5m2, DataType.e5m2, element_acc,
OpcodeClass.TensorOp,
MathOperation.multiply_add_fast_accum),
]
min_cc = 89
max_cc = 89
max_cc = 100
alignment_constraints = [16,]
alignment_constraints_small_channels = [16, 8, 4]
@@ -5077,6 +5100,18 @@ def GenerateSM89_TensorOp_16832_fp8(manifest, cuda_version):
else:
op.C.alignment = 8
def GenerateSM89_TensorOp_16832_fp8_fp32acc(manifest, cuda_version):
if not CudaToolkitVersionSatisfies(cuda_version, 12, 4):
return
GenerateSM89_TensorOp_16832_fp8(manifest, DataType.f32)
def GenerateSM89_TensorOp_16832_fp8_fp16acc(manifest, cuda_version):
if not CudaToolkitVersionSatisfies(cuda_version, 12, 8):
return
GenerateSM89_TensorOp_16832_fp8(manifest, DataType.f16)
#
def GenerateSM89_SparseTensorOp_16864_fp8(manifest, cuda_version):
@@ -5177,7 +5212,8 @@ def GenerateSM89_SparseTensorOp_16864_fp8(manifest, cuda_version):
#
def GenerateSM89(manifest, cuda_version):
GenerateSM89_TensorOp_16832_fp8(manifest, cuda_version)
GenerateSM89_TensorOp_16832_fp8_fp32acc(manifest, cuda_version)
GenerateSM89_TensorOp_16832_fp8_fp16acc(manifest, cuda_version)
GenerateSM89_SparseTensorOp_16864_fp8(manifest, cuda_version)
###################################################################################################
@@ -5189,6 +5225,7 @@ try:
generate_tf32_math_instructions_sm90,
generate_int8_math_instructions_sm90,
generate_fp8_math_instructions_sm90,
generate_mixed_dtype_math_instructions_sm90,
make_sparse_math_instructions,
generate_tile_descriptions_sm90,
get_valid_schedules,
@@ -5201,6 +5238,7 @@ except ImportError:
generate_tf32_math_instructions_sm90,
generate_int8_math_instructions_sm90,
generate_fp8_math_instructions_sm90,
generate_mixed_dtype_math_instructions_sm90,
make_sparse_math_instructions,
generate_tile_descriptions_sm90,
get_valid_schedules,
@@ -5208,8 +5246,8 @@ except ImportError:
fix_alignments,
)
def GenerateSM90_TensorOp_16b_WGMMA_gemm(manifest, cuda_version):
if not CudaToolkitVersionSatisfies(cuda_version, 12, 0):
def GenerateSM90_TensorOp_16b_WGMMA_gemm(manifest, cuda_version, gemm_kind=GemmKind.Universal3x):
if not CudaToolkitVersionSatisfies(cuda_version, 12, 3 if is_grouped(gemm_kind) else 0):
return
instantiation_level = manifest.get_sm90_instantiation_level(pruned_level=100, default_level=131, exhaustive_level=9992)
@@ -5262,10 +5300,11 @@ def GenerateSM90_TensorOp_16b_WGMMA_gemm(manifest, cuda_version):
data_types=data_type,
instantiation_level=instantiation_level,
layout=layout,
gemm_kind=gemm_kind,
)
if len(schedules):
CreateGemmUniversal3xOperator(manifest, [layout], [tile_desc], data_type, 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,
@@ -5728,8 +5767,8 @@ def GenerateSM90_SparseTensorOp_int8_WGMMA_gemm(manifest, cuda_version):
tile_schedulers=[TileSchedulerType.StreamK])
def GenerateSM90_TensorOp_fp8_WGMMA_gemm(manifest, cuda_version):
if not CudaToolkitVersionSatisfies(cuda_version, 12, 0):
def GenerateSM90_TensorOp_fp8_WGMMA_gemm(manifest, cuda_version, gemm_kind=GemmKind.Universal3x):
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)
@@ -5783,10 +5822,11 @@ def GenerateSM90_TensorOp_fp8_WGMMA_gemm(manifest, cuda_version):
data_types=data_type,
instantiation_level=instantiation_level,
layout=layout,
gemm_kind=gemm_kind,
)
if len(schedules):
CreateGemmUniversal3xOperator(manifest, [layout], [tile_desc], data_type, 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,
@@ -5851,6 +5891,90 @@ def GenerateSM90_TensorOp_fp8_WGMMA_alignx_gemm(manifest, cuda_version):
stream_k_schedules,
tile_schedulers=[TileSchedulerType.StreamK])
def GenerateSM90_TensorOp_mixed_dtype_WGMMA_gemm(manifest, cuda_version):
if not CudaToolkitVersionSatisfies(cuda_version, 12, 1):
return
instantiation_level = manifest.get_sm90_instantiation_level(pruned_level=20, default_level=121, exhaustive_level=9999)
is_aligned = True
# layouts for ABC, their alignments will be fixed later based on the data type
layouts = [
[[LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 16]],
]
valid_types_for_a_b_acc = [
(DataType.e4m3, DataType.f16, DataType.f32),
(DataType.e4m3, DataType.bf16, DataType.f32),
(DataType.e5m2, DataType.f16, DataType.f32),
(DataType.e5m2, DataType.bf16, DataType.f32),
(DataType.s8, DataType.f16, DataType.f32),
(DataType.s8, DataType.bf16, DataType.f32),
(DataType.u8, DataType.f16, DataType.f32),
(DataType.u8, DataType.bf16, DataType.f32),
(DataType.s4, DataType.f16, DataType.f32),
(DataType.s4, DataType.bf16, DataType.f32),
(DataType.s4, DataType.e4m3, DataType.f32),
(DataType.s4, DataType.e5m2, DataType.f32),
(DataType.u4, DataType.f16, DataType.f32),
(DataType.u4, DataType.bf16, DataType.f32),
(DataType.u2, DataType.f16, DataType.f32),
(DataType.u2, DataType.bf16, DataType.f32),
(DataType.s2, DataType.f16, DataType.f32),
(DataType.s2, DataType.bf16, DataType.f32),
]
# Note: For sizeof(a_type) > sizeof(b_type), some generated kernels might crash due to a compiler bug. Disable it for now.
#swapped_valid_types_for_a_b_acc = [(b_type, a_type, acc_type) for a_type, b_type, acc_type in valid_types_for_a_b_acc]
#valid_types_for_a_b_acc = valid_types_for_a_b_acc + swapped_valid_types_for_a_b_acc
math_instructions = generate_mixed_dtype_math_instructions_sm90(instantiation_level, valid_types_for_a_b_acc)
valid_types_for_d = [DataType.f32]
valid_types_for_c = [DataType.f32]
tile_descriptions = generate_tile_descriptions_sm90(
math_instructions=math_instructions,
is_aligned=is_aligned,
level=instantiation_level)
for tile_desc in tile_descriptions:
math_inst = tile_desc.math_instruction
data_types = []
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,
)
)
for layout in layouts:
for data_type in data_types:
# Fix alignments, DataTypeSize are in the unit of bits
alignment_bits = 128
layout[0][1] = alignment_bits // DataTypeSize[data_type['a_type']]
layout[1][1] = alignment_bits // DataTypeSize[data_type['b_type']]
layout[2][1] = alignment_bits // DataTypeSize[data_type['c_type']]
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,
)
if len(schedules):
CreateGemmUniversal3xOperator(manifest, [layout], [tile_desc], data_type, schedules)
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])
def GenerateSM90_SparseTensorOp_fp8_WGMMA_gemm(manifest, cuda_version):
if not CudaToolkitVersionSatisfies(cuda_version, 12, 2):
@@ -6662,7 +6786,7 @@ def GenerateSM100_TensorOp_32b_UMMA_gemm(manifest, cuda_version):
CreateGemmUniversal3xOperator(manifest, layouts, tile_descriptions, data_types,
[[KernelScheduleType.TmaWarpSpecialized2SmSm100, epi_schedule]], tile_schedulers=tile_schedulers)
def GenerateSM100_TensorOp_16b_UMMA_gemm(manifest, cuda_version):
def GenerateSM100_TensorOp_16b_UMMA_gemm(manifest, cuda_version, gemm_kind=GemmKind.Universal3x):
if not CudaToolkitVersionSatisfies(cuda_version, 12, 8):
return
@@ -6680,6 +6804,8 @@ def GenerateSM100_TensorOp_16b_UMMA_gemm(manifest, cuda_version):
min_cc = 100
max_cc = 100
grouped = is_grouped(gemm_kind)
math_instructions_1sm = [
# f16 -> f16
#MathInstruction(
@@ -6736,6 +6862,7 @@ def GenerateSM100_TensorOp_16b_UMMA_gemm(manifest, cuda_version):
MathOperation.multiply_add)]
cluster_shapes_1sm = [[1,2,1], [1,1,1], [1,4,1],[4,4,1]
, DynamicClusterShape
]
tile_schedulers = [
@@ -6776,9 +6903,11 @@ def GenerateSM100_TensorOp_16b_UMMA_gemm(manifest, cuda_version):
for layout in layouts:
layout[2][1] = 128 // DataTypeSize[data_types[0]["d_type"]]
kernel_schedule = KernelScheduleType.TmaWarpSpecialized1SmSm100 if not grouped else KernelScheduleType.PtrArrayTmaWarpSpecialized1SmSm100
epi_schedule = EpilogueScheduleType.TmaWarpSpecialized1Sm if not grouped else EpilogueScheduleType.PtrArrayTmaWarpSpecialized1Sm
CreateGemmUniversal3xOperator(manifest, layouts, tile_descriptions, data_types,
[[KernelScheduleType.TmaWarpSpecialized1SmSm100, EpilogueScheduleType.TmaWarpSpecialized1Sm]],
tile_schedulers=tile_schedulers)
[[kernel_schedule, epi_schedule]],
tile_schedulers=tile_schedulers, gemm_kind=gemm_kind)
# for mixed precision kernels, also generate kernels that write output matrix in the A/B format
# Avoid emitting two kernels if the accumulator type does not differ from the input type (e.g. F16 accumulation)
@@ -6806,8 +6935,8 @@ def GenerateSM100_TensorOp_16b_UMMA_gemm(manifest, cuda_version):
layout[2][1] = 128 // DataTypeSize[data_types_mixed[0]["d_type"]]
CreateGemmUniversal3xOperator(manifest, layouts, tile_descriptions, data_types_mixed,
[[KernelScheduleType.TmaWarpSpecialized1SmSm100, EpilogueScheduleType.TmaWarpSpecialized1Sm]],
tile_schedulers=tile_schedulers)
[[kernel_schedule, epi_schedule]],
tile_schedulers=tile_schedulers, gemm_kind=gemm_kind)
# 2xSM MMA kernels
math_instructions_2sm = [
@@ -6886,6 +7015,7 @@ def GenerateSM100_TensorOp_16b_UMMA_gemm(manifest, cuda_version):
MathOperation.multiply_add)]
cluster_shapes_2sm = [[2,1,1], [2,2,1], [2,4,1], [4,1,1], [4,2,1], [4,4,1]
, DynamicClusterShape
]
for math_inst in math_instructions_2sm:
@@ -6921,13 +7051,16 @@ def GenerateSM100_TensorOp_16b_UMMA_gemm(manifest, cuda_version):
for layout in layouts:
layout[2][1] = 128 // DataTypeSize[data_types[0]["d_type"]]
if math_inst.instruction_shape[0] == 128:
if grouped:
epi_schedule = EpilogueScheduleType.PtrArrayTmaWarpSpecialized2Sm
elif math_inst.instruction_shape[0] == 128:
epi_schedule = EpilogueScheduleType.TmaWarpSpecialized2Sm
else:
epi_schedule = EpilogueScheduleType.ScheduleAuto
kernel_schedule = KernelScheduleType.TmaWarpSpecialized2SmSm100 if not grouped else KernelScheduleType.PtrArrayTmaWarpSpecialized2SmSm100
CreateGemmUniversal3xOperator(manifest, layouts, tile_descriptions, data_types,
[[KernelScheduleType.TmaWarpSpecialized2SmSm100, epi_schedule]], tile_schedulers=tile_schedulers)
[[kernel_schedule, epi_schedule]], tile_schedulers=tile_schedulers, gemm_kind=gemm_kind)
# for mixed precision kernels, also generate kernels that write output matrix in the A/B format
# Avoid emitting two kernels if the accumulator type does not differ from the input type (e.g. F16 accumulation)
@@ -6955,9 +7088,9 @@ def GenerateSM100_TensorOp_16b_UMMA_gemm(manifest, cuda_version):
layout[2][1] = 128 // DataTypeSize[data_types_mixed[0]["d_type"]]
CreateGemmUniversal3xOperator(manifest, layouts, tile_descriptions, data_types_mixed,
[[KernelScheduleType.TmaWarpSpecialized2SmSm100, epi_schedule]], tile_schedulers=tile_schedulers)
[[kernel_schedule, epi_schedule]], tile_schedulers=tile_schedulers, gemm_kind=gemm_kind)
def GenerateSM100_TensorOp_fp8_UMMA_gemm(manifest, cuda_version):
def GenerateSM100_TensorOp_fp8_UMMA_gemm(manifest, cuda_version, gemm_kind=GemmKind.Universal3x):
if not CudaToolkitVersionSatisfies(cuda_version, 12, 8):
return
@@ -6976,6 +7109,7 @@ def GenerateSM100_TensorOp_fp8_UMMA_gemm(manifest, cuda_version):
min_cc = 100
max_cc = 100
epi_type = DataType.f32
grouped = is_grouped(gemm_kind)
math_instructions_1sm = [
# inst 64x128
@@ -7038,6 +7172,7 @@ def GenerateSM100_TensorOp_fp8_UMMA_gemm(manifest, cuda_version):
MathOperation.multiply_add)]
cluster_shapes_1sm = [[1,2,1], [2,1,1], [1,1,1], [1,4,1], [4,4,1]
, DynamicClusterShape
]
tile_schedulers = [
@@ -7163,9 +7298,14 @@ def GenerateSM100_TensorOp_fp8_UMMA_gemm(manifest, cuda_version):
if ( data_type["a_type"] == DataType.e4m3 ) and ( data_type["b_type"] == DataType.e4m3 ) and\
( data_type["d_type"] == DataType.e5m2 ):
continue
# don't support runtime data type for grouped yet
if grouped and (data_type["a_type"] == DataType.f8 or data_type["b_type"] == DataType.f8):
continue
kernel_schedule = KernelScheduleType.TmaWarpSpecialized1SmSm100 if not grouped else KernelScheduleType.PtrArrayTmaWarpSpecialized1SmSm100
epi_schedule = EpilogueScheduleType.TmaWarpSpecialized1Sm if not grouped else EpilogueScheduleType.PtrArrayTmaWarpSpecialized1Sm
CreateGemmUniversal3xOperator(manifest, layouts, tile_descriptions, data_type,
[[KernelScheduleType.TmaWarpSpecialized1SmSm100, EpilogueScheduleType.TmaWarpSpecialized1Sm]],
tile_schedulers=tile_schedulers)
[[kernel_schedule, epi_schedule]],
tile_schedulers=tile_schedulers, gemm_kind=gemm_kind)
# 2xSM MMA kernels
math_instructions_2sm = [
@@ -7241,6 +7381,7 @@ def GenerateSM100_TensorOp_fp8_UMMA_gemm(manifest, cuda_version):
]
cluster_shapes_2sm = [[2,1,1], [2,2,1], [2,4,1], [4,1,1], [4,2,1], [4,4,1]
, DynamicClusterShape
]
for math_inst in math_instructions_2sm:
@@ -7361,15 +7502,20 @@ def GenerateSM100_TensorOp_fp8_UMMA_gemm(manifest, cuda_version):
if ( data_type["a_type"] == DataType.e4m3 ) and ( data_type["b_type"] == DataType.e4m3 ) and\
( data_type["d_type"] == DataType.e5m2 ):
continue
# don't support runtime data type for grouped yet
if grouped and (data_type["a_type"] == DataType.f8 or data_type["b_type"] == DataType.f8):
continue
if math_inst.instruction_shape[0] == 128:
if grouped:
epi_schedule = EpilogueScheduleType.PtrArrayTmaWarpSpecialized2Sm
elif math_inst.instruction_shape[0] == 128:
epi_schedule = EpilogueScheduleType.TmaWarpSpecialized2Sm
else:
epi_schedule = EpilogueScheduleType.ScheduleAuto
kernel_schedule = KernelScheduleType.TmaWarpSpecialized2SmSm100 if not grouped else KernelScheduleType.PtrArrayTmaWarpSpecialized2SmSm100
CreateGemmUniversal3xOperator(manifest, layouts, tile_descriptions, data_type,
[[KernelScheduleType.TmaWarpSpecialized2SmSm100, epi_schedule]], tile_schedulers=tile_schedulers)
[[kernel_schedule, epi_schedule]], tile_schedulers=tile_schedulers, gemm_kind=gemm_kind)
def GenerateSM100_TensorOp_mixed_8bits_UMMA_gemm_with_block_scaled(manifest, cuda_version):
@@ -7460,6 +7606,7 @@ def GenerateSM100_TensorOp_mixed_8bits_UMMA_gemm_with_block_scaled(manifest, cud
[2,1,1],
# [1,4,1],
[4,4,1]
, DynamicClusterShape
]
# 1xSM MMA kernels
@@ -7533,6 +7680,7 @@ def GenerateSM100_TensorOp_mixed_8bits_UMMA_gemm_with_block_scaled(manifest, cud
[4,1,1],
# [4,2,1],
[4,4,1]
, DynamicClusterShape
]
for math_inst in math_instructions_2sm:
@@ -7728,6 +7876,7 @@ def GenerateSM100_TensorOp_fp4_UMMA_gemm_with_block_scaled(manifest, cuda_versio
[2,1,1],
# [1,4,1],
[4,4,1]
, DynamicClusterShape
]
# 1xSM MMA kernels
@@ -7841,6 +7990,7 @@ def GenerateSM100_TensorOp_fp4_UMMA_gemm_with_block_scaled(manifest, cuda_versio
[4,1,1],
# [4,2,1],
[4,4,1]
, DynamicClusterShape
]
for math_inst in math_instructions_2sm:
@@ -8419,6 +8569,9 @@ def GenerateSM100(manifest, cuda_version):
GenerateSM100_TensorOp_int8_UMMA_gemm(manifest, cuda_version)
GenerateSM100_TensorOp_fp8_UMMA_gemm(manifest, cuda_version)
# grouped GEMM
GenerateSM100_TensorOp_fp8_UMMA_gemm(manifest, cuda_version, gemm_kind=GemmKind.GroupedGemmUniversal3x)
GenerateSM100_TensorOp_16b_UMMA_gemm(manifest, cuda_version, gemm_kind=GemmKind.GroupedGemmUniversal3x)
#
# Block Scaled Gemm
#
@@ -8800,7 +8953,10 @@ def GenerateSM90(manifest, cuda_version):
GenerateSM90_TensorOp_int8_WGMMA_alignx_gemm(manifest, cuda_version)
GenerateSM90_TensorOp_fp8_WGMMA_gemm(manifest, cuda_version)
GenerateSM90_TensorOp_fp8_WGMMA_alignx_gemm(manifest, cuda_version)
GenerateSM90_TensorOp_mixed_dtype_WGMMA_gemm(manifest, cuda_version)
GenerateSM90_TensorOp_1684(manifest, cuda_version)
GenerateSM90_TensorOp_16b_WGMMA_gemm(manifest, cuda_version, gemm_kind=GemmKind.GroupedGemmUniversal3x)
GenerateSM90_TensorOp_fp8_WGMMA_gemm(manifest, cuda_version, gemm_kind=GemmKind.GroupedGemmUniversal3x)
GenerateSM90_TensorOp_1684_complex(manifest, cuda_version)
GenerateSM90_TensorOp_1684_complex_gaussian(manifest, cuda_version)
GenerateSM90_TensorOp_1684_rank_k(manifest, cuda_version)
@@ -8899,6 +9055,12 @@ if __name__ == "__main__":
if 'library' in args.generator_target.split(','):
manifest.emit(GeneratorTarget.Library)
if 'kernel_testlist_l0' in args.generator_target.split(','):
emit_gemm_kernel_testlist(manifest, args.curr_build_dir, args.architectures, "functional_L0")
if 'kernel_testlist_l1' in args.generator_target.split(','):
emit_gemm_kernel_testlist(manifest, args.curr_build_dir, args.architectures, "functional_L1")
if args.selected_kernel_list is not None:
if len(manifest.selected_kernels) > 0:
with open(args.selected_kernel_list, 'w') as file_writer: