Add couple configs into generator.py for mixed input MM (#1350)

* Add couple configs into generator.py for mixed input MM

* change one unit test name; reenable 128x32 in the profiler

* Added U8/BF16 tests.

---------

Co-authored-by: Haicheng Wu <haichengw@nvidia.com>
Co-authored-by: Haicheng Wu <57973641+hwu36@users.noreply.github.com>
This commit is contained in:
Aleksandar Samardžić
2024-08-16 06:59:29 +02:00
committed by GitHub
parent 865be73a97
commit 3f084f7f3c
21 changed files with 1506 additions and 80 deletions

View File

@@ -178,30 +178,16 @@ class GemmOperation:
if self.is_complex():
extended_name = "${core_name}"
else:
# e.g. f16_f16_f32_void_f32 kernel
if self.C.element != self.tile_description.math_instruction.element_accumulator and \
self.A.element != self.tile_description.math_instruction.element_accumulator:
extended_name = "${element_c}_${core_name}_${element_a}"
if self.is_mixed_input():
extended_name += "_${element_b}"
# e.g. f32_f32_f32_void_f32 kernel
elif self.C.element != self.tile_description.math_instruction.element_accumulator and \
self.A.element == self.tile_description.math_instruction.element_accumulator:
extended_name = "${element_c}_${core_name}"
if self.is_mixed_input():
extended_name += "_${element_b}"
# e.g. f16_f16_f32_f32_f32 kernel
elif self.C.element == self.tile_description.math_instruction.element_accumulator and \
self.A.element != self.tile_description.math_instruction.element_accumulator:
extended_name = "${core_name}_${element_a}"
if self.is_mixed_input():
extended_name += "_${element_b}"
# e.g. f32_f32_f32_f32_f32 kernel
if self.is_mixed_input():
extended_name = "${core_name}_${element_a}_${element_b}"
if self.C.element != self.tile_description.math_instruction.element_accumulator:
extended_name = "${element_c}_" + extended_name
else:
extended_name = "${core_name}"
if self.C.element != self.tile_description.math_instruction.element_accumulator:
extended_name = "${element_c}_" + extended_name
if self.A.element != self.tile_description.math_instruction.element_accumulator:
extended_name += "_${element_a}"
extended_name = SubstituteTemplate(extended_name, {
'element_a': DataTypeNames[self.A.element],