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

@@ -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: