v3.8.0 update (#2082)
* 3.8 update * fix Markus' name --------- Co-authored-by: yuzhai <yuzhai@nvidia.com>
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user