v3.8.0 update (#2082)
* 3.8 update * fix Markus' name --------- Co-authored-by: yuzhai <yuzhai@nvidia.com>
This commit is contained in:
@@ -43,7 +43,7 @@ import os.path
|
||||
import shutil
|
||||
import sys
|
||||
import copy
|
||||
from typing import Any, Optional, Sequence, Tuple
|
||||
from typing import Any, Optional, Sequence, Tuple, List
|
||||
|
||||
try:
|
||||
import builtins
|
||||
@@ -153,6 +153,17 @@ def generate_int8_math_instruction_shapes_sm90(level: int):
|
||||
]
|
||||
return filtered_list_of_wgmma_shapes
|
||||
|
||||
def generate_mixed_dtype_math_instructions_shapes_sm90(wgmma_level: int, a_type: DataType, b_type: DataType):
|
||||
# DataTypeSize are in the unit of bits
|
||||
a_bytes = DataTypeSize[a_type] // 8
|
||||
b_bytes = DataTypeSize[b_type] // 8
|
||||
if a_bytes == 4 or b_bytes == 4:
|
||||
return generate_tf32_math_instruction_shapes_sm90(wgmma_level)
|
||||
elif a_bytes == 2 or b_bytes == 2:
|
||||
return generate_fp16_bf16_math_instruction_shapes_sm90(wgmma_level)
|
||||
else:
|
||||
return generate_fp8_math_instruction_shapes_sm90(wgmma_level)
|
||||
|
||||
###########
|
||||
|
||||
def generate_tf32_math_instructions_sm90(level: int):
|
||||
@@ -219,6 +230,22 @@ def generate_fp8_math_instructions_sm90(level: int):
|
||||
]
|
||||
return math_instructions
|
||||
|
||||
def generate_mixed_dtype_math_instructions_sm90(level: int, types_of_a_b_acc: List[Tuple[DataType, DataType, DataType]]):
|
||||
wgmma_level = get_wgmma_level_from_global_level(level)
|
||||
math_instructions = []
|
||||
for a_type, b_type, acc_type in types_of_a_b_acc:
|
||||
math_instruction_shapes = generate_mixed_dtype_math_instructions_shapes_sm90(wgmma_level, a_type, b_type)
|
||||
for math_instruction_shape in math_instruction_shapes:
|
||||
math_instructions += [
|
||||
MathInstruction(
|
||||
math_instruction_shape,
|
||||
a_type, b_type, acc_type,
|
||||
OpcodeClass.TensorOp,
|
||||
MathOperation.multiply_add
|
||||
),
|
||||
]
|
||||
return math_instructions
|
||||
|
||||
def generate_int8_math_instructions_sm90(level: int):
|
||||
wgmma_level = get_wgmma_level_from_global_level(level)
|
||||
math_instructions = []
|
||||
@@ -407,7 +434,7 @@ def can_tile_desc_use_shmem_in_epilogue(tile_description, data_types):
|
||||
|
||||
|
||||
def get_valid_schedules(tile_description, cuda_version, is_aligned, data_types, layout,
|
||||
instantiation_level, enable_fp8_fast_acc=True):
|
||||
instantiation_level, enable_fp8_fast_acc=True, gemm_kind=GemmKind.Universal3x):
|
||||
# Level 0: prune according to existing generator.py behavior
|
||||
# Level >= 1: no pruning
|
||||
level = get_pruning_level_from_global_level(instantiation_level)
|
||||
@@ -428,8 +455,6 @@ def get_valid_schedules(tile_description, cuda_version, is_aligned, data_types,
|
||||
is_fp32 = data_types["a_type"] in FP32_TYPES and data_types["b_type"] in FP32_TYPES
|
||||
requires_transposed_epilogue = is_fp32 and layout[0][0] == LayoutType.RowMajor and layout[1][0] == LayoutType.RowMajor
|
||||
|
||||
is_sparse = tile_description.math_instruction.opcode_class == OpcodeClass.SparseTensorOp
|
||||
|
||||
can_do_cooperative = is_tile_desc_compatible_with_cooperative(tile_description)
|
||||
can_do_tma_epilogue = is_aligned and not requires_transposed_epilogue and can_tile_desc_use_shmem_in_epilogue(tile_description, data_types)
|
||||
|
||||
@@ -464,6 +489,16 @@ def get_valid_schedules(tile_description, cuda_version, is_aligned, data_types,
|
||||
if is_fp32 and (is_tn or is_nn) and (cta_n % cta_k != 0):
|
||||
return [], []
|
||||
|
||||
grouped = gemm_kind == GemmKind.GroupedGemmUniversal3x
|
||||
if grouped:
|
||||
# the following cases are unsupported by grouped GEMM
|
||||
if not is_aligned:
|
||||
return [], []
|
||||
if not can_do_tma_epilogue:
|
||||
return [], []
|
||||
if requires_transposed_epilogue:
|
||||
return [], []
|
||||
|
||||
# Early pruning
|
||||
if level < 1:
|
||||
# Don't stamp out FP16/BF16 kernels smaller than or equal to 64x128x64
|
||||
@@ -477,20 +512,23 @@ def get_valid_schedules(tile_description, cuda_version, is_aligned, data_types,
|
||||
if not is_void_c or d_type not in FP8_TYPES:
|
||||
return [], []
|
||||
if CudaToolkitVersionSatisfies(cuda_version, 12, 1) and can_do_cooperative and can_do_tma_epilogue:
|
||||
return [
|
||||
schedules = []
|
||||
if not grouped:
|
||||
schedules.append(
|
||||
[
|
||||
KernelScheduleType.TmaWarpSpecializedCooperative,
|
||||
EpilogueScheduleType.TmaWarpSpecializedCooperative
|
||||
])
|
||||
schedules.append(
|
||||
[
|
||||
KernelScheduleType.TmaWarpSpecializedCooperative,
|
||||
EpilogueScheduleType.TmaWarpSpecializedCooperative
|
||||
],
|
||||
[
|
||||
KernelScheduleType.TmaWarpSpecializedCooperativeFP8FastAccum,
|
||||
EpilogueScheduleType.TmaWarpSpecializedCooperative
|
||||
],
|
||||
] , []
|
||||
KernelScheduleType.TmaWarpSpecializedCooperativeFP8FastAccum if not grouped else KernelScheduleType.KernelPtrArrayTmaWarpSpecializedCooperativeFP8FastAccum,
|
||||
EpilogueScheduleType.TmaWarpSpecializedCooperative if not grouped else EpilogueScheduleType.PtrArrayTmaWarpSpecializedCooperative,
|
||||
])
|
||||
return schedules, []
|
||||
return [], []
|
||||
|
||||
if is_fp8 and not is_large_fp8_tile:
|
||||
valid_dtypes_for_c = [DataType.f32, DataType.bf16, DataType.f16]
|
||||
valid_dtypes_for_c = [DataType.f32, DataType.bf16, DataType.f16, DataType.void]
|
||||
# Prune all configs with fp8 source, and all configs with non-fp8 output
|
||||
# that have different dtypes for source and output.
|
||||
if c_type not in valid_dtypes_for_c or (d_type not in FP8_TYPES and c_type != d_type):
|
||||
@@ -504,6 +542,33 @@ def get_valid_schedules(tile_description, cuda_version, is_aligned, data_types,
|
||||
if is_void_c and not can_do_tma_epilogue:
|
||||
return [], []
|
||||
|
||||
# For mixed input data types
|
||||
a_type_size = DataTypeSize[data_types["a_type"]]
|
||||
b_type_size = DataTypeSize[data_types["b_type"]]
|
||||
if a_type_size != b_type_size and CudaToolkitVersionSatisfies(cuda_version, 12, 1):
|
||||
schedules = []
|
||||
epilogue_schedule = EpilogueScheduleType.TmaWarpSpecialized
|
||||
if a_type_size > b_type_size:
|
||||
epilogue_schedule = EpilogueScheduleType.EpilogueTransposed
|
||||
schedules.append([
|
||||
KernelScheduleType.TmaWarpSpecialized,
|
||||
epilogue_schedule
|
||||
])
|
||||
schedules.append([
|
||||
KernelScheduleType.TmaWarpSpecializedPingpong,
|
||||
epilogue_schedule
|
||||
])
|
||||
if cta_m >= 128:
|
||||
if a_type_size > b_type_size:
|
||||
epilogue_schedule = EpilogueScheduleType.EpilogueTransposed
|
||||
else:
|
||||
epilogue_schedule = EpilogueScheduleType.TmaWarpSpecializedCooperative
|
||||
schedules.append([
|
||||
KernelScheduleType.TmaWarpSpecializedCooperative,
|
||||
epilogue_schedule
|
||||
])
|
||||
return schedules, []
|
||||
|
||||
if not is_aligned:
|
||||
schedules = [[KernelScheduleType.CpAsyncWarpSpecialized,
|
||||
default_epilogue]]
|
||||
@@ -521,6 +586,15 @@ def get_valid_schedules(tile_description, cuda_version, is_aligned, data_types,
|
||||
|
||||
return schedules, stream_k_schedules
|
||||
|
||||
if grouped:
|
||||
pingpong = KernelScheduleType.KernelPtrArrayTmaWarpSpecializedPingpong if not is_fp8 else KernelScheduleType.KernelPtrArrayTmaWarpSpecializedPingpongFP8FastAccum
|
||||
cooperative = KernelScheduleType.KernelPtrArrayTmaWarpSpecializedCooperative if not is_fp8 else KernelScheduleType.KernelPtrArrayTmaWarpSpecializedCooperativeFP8FastAccum
|
||||
if can_do_tma_epilogue:
|
||||
schedules.append([pingpong, EpilogueScheduleType.PtrArrayTmaWarpSpecializedPingpong])
|
||||
if can_do_cooperative:
|
||||
schedules.append([cooperative, EpilogueScheduleType.PtrArrayTmaWarpSpecializedCooperative])
|
||||
return schedules, []
|
||||
|
||||
schedules = []
|
||||
# Pruning: emit Void-C kernels with persistent kernels only
|
||||
if level >= 1 or not is_void_c:
|
||||
|
||||
Reference in New Issue
Block a user