Skip to content

Commit de8ee37

Browse files
Fixing GEMM Multi D on Tile Engine (#3583)
1 parent 644cdbe commit de8ee37

1 file changed

Lines changed: 174 additions & 172 deletions

File tree

tile_engine/ops/gemm/gemm_instance_builder.py

Lines changed: 174 additions & 172 deletions
Original file line numberDiff line numberDiff line change
@@ -676,69 +676,71 @@ def populate_launch(
676676
if self.kernel_name_prefix == "gemm_multi_d":
677677
instance_code += """
678678
679-
// Kernel type
680-
using GemmKernelMultiD = ck_tile::GemmKernelMultiD<TilePartitioner, GemmPipeline, GemmEpilogue>;
681-
682-
// Kernel arguments
683-
auto kargs = GemmKernelMultiD::MakeKernelArgs(args);
684-
685-
if (!GemmKernelMultiD::IsSupportedArgument(kargs)) {
686-
throw std::runtime_error("Wrong! Arguments not supported! Skipping gemm!");
687-
}
679+
// Kernel type
680+
using GemmKernelMultiD = ck_tile::GemmKernelMultiD<TilePartitioner, GemmPipeline, GemmEpilogue>;
681+
682+
// Kernel arguments
683+
auto kargs = GemmKernelMultiD::MakeKernelArgs(args);
684+
685+
if (!GemmKernelMultiD::IsSupportedArgument(kargs)) {
686+
throw std::runtime_error("Wrong! Arguments not supported! Skipping gemm!");
687+
}
688688
689-
// Get grid and block sizes
690-
const dim3 grids = GemmKernelMultiD::GridSize(args.M, args.N, args.k_batch);
691-
const dim3 blocks = GemmKernelMultiD::BlockSize();
692-
693-
if(stream.log_level_ > 0) {
694-
std::cout << "Launching kernel with args: " << GemmKernelMultiD::GetName() << '\\n'
695-
<< "grid: {" << grids.x << ", " << grids.y << ", " << grids.z << "}"
696-
<< ", blocks: {" << blocks.x << ", " << blocks.y << ", " << blocks.z << "}"
697-
<< std::endl;
698-
}"""
689+
// Get grid and block sizes
690+
const dim3 grids = GemmKernelMultiD::GridSize(args.M, args.N, args.k_batch);
691+
const dim3 blocks = GemmKernelMultiD::BlockSize();
692+
693+
if(stream.log_level_ > 0) {
694+
std::cout << "Launching kernel with args: " << GemmKernelMultiD::GetName() << '\\n'
695+
<< "grid: {" << grids.x << ", " << grids.y << ", " << grids.z << "}"
696+
<< ", blocks: {" << blocks.x << ", " << blocks.y << ", " << blocks.z << "}"
697+
<< std::endl;
698+
}"""
699699

700700
instance_code += f"""
701-
// Launch kernel
702-
constexpr int kBlockPerCu = {k_block_per_cu};
703-
float ave_time = ck_tile::launch_kernel(
704-
stream,
705-
ck_tile::make_kernel<kBlockPerCu>(GemmKernelMultiD{{}}, grids, blocks, 0, kargs));
706-
707-
return ave_time;
708-
}};"""
701+
// Launch kernel
702+
constexpr int kBlockPerCu = {k_block_per_cu};
703+
float ave_time = ck_tile::launch_kernel(
704+
stream,
705+
ck_tile::make_kernel<kBlockPerCu>(GemmKernelMultiD{{}}, grids, blocks, 0, kargs));
706+
707+
return ave_time;
708+
}}
709+
}};
710+
"""
709711

710712
elif self.kernel_name_prefix in ["gemm_universal", "gemm_preshuffle"]:
711713
instance_code += f"""
712714
713715
// Kernel type
714716
using GemmKernel = ck_tile::GemmKernel<TilePartitioner, GemmPipeline, GemmEpilogue>;
715717
716-
// Kernel arguments
717-
auto kargs = GemmKernel::MakeKernelArgs(args);
718-
719-
if (!GemmKernel::IsSupportedArgument(kargs)) {{
720-
throw std::runtime_error("Wrong! Arguments not supported! Skipping gemm!");
721-
}}
722-
723-
// Get grid and block sizes
724-
const dim3 grids = {"GemmKernel::MaxOccupancyGridSize(stream)" if persistent in [True, "true"] else "GemmKernel::GridSize(args.M, args.N, args.k_batch)"};
725-
const dim3 blocks = GemmKernel::BlockSize();
726-
727-
if(stream.log_level_ > 0) {{
728-
std::cout << "Launching kernel with args: " << GemmKernel::GetName() << '\\n'
729-
<< "grid: {{" << grids.x << ", " << grids.y << ", " << grids.z << "}}"
730-
<< ", blocks: {{" << blocks.x << ", " << blocks.y << ", " << blocks.z << "}}"
731-
<< std::endl;
732-
}}"""
718+
// Kernel arguments
719+
auto kargs = GemmKernel::MakeKernelArgs(args);
720+
721+
if (!GemmKernel::IsSupportedArgument(kargs)) {{
722+
throw std::runtime_error("Wrong! Arguments not supported! Skipping gemm!");
723+
}}
724+
725+
// Get grid and block sizes
726+
const dim3 grids = {"GemmKernel::MaxOccupancyGridSize(stream)" if persistent in [True, "true"] else "GemmKernel::GridSize(args.M, args.N, args.k_batch)"};
727+
const dim3 blocks = GemmKernel::BlockSize();
728+
729+
if(stream.log_level_ > 0) {{
730+
std::cout << "Launching kernel with args: " << GemmKernel::GetName() << '\\n'
731+
<< "grid: {{" << grids.x << ", " << grids.y << ", " << grids.z << "}}"
732+
<< ", blocks: {{" << blocks.x << ", " << blocks.y << ", " << blocks.z << "}}"
733+
<< std::endl;
734+
}}"""
733735

734736
instance_code += f"""
735-
// Launch kernel
736-
constexpr int kBlockPerCu = {k_block_per_cu};
737-
float ave_time = ck_tile::launch_kernel(
738-
stream,
739-
ck_tile::make_kernel<kBlockPerCu>(GemmKernel{{}}, grids, blocks, 0, kargs));
740-
741-
return ave_time;
737+
// Launch kernel
738+
constexpr int kBlockPerCu = {k_block_per_cu};
739+
float ave_time = ck_tile::launch_kernel(
740+
stream,
741+
ck_tile::make_kernel<kBlockPerCu>(GemmKernel{{}}, grids, blocks, 0, kargs));
742+
743+
return ave_time;
742744
}}
743745
}};
744746
"""
@@ -747,8 +749,8 @@ def populate_launch(
747749
def populate_epilogue(self, epilogue):
748750
instance_code = """
749751
750-
// Epilogue
751-
"""
752+
// Epilogue
753+
"""
752754

753755
if epilogue == "cshuffle":
754756
if self.kernel_name_prefix == "gemm_universal":
@@ -769,145 +771,145 @@ def populate_epilogue(self, epilogue):
769771

770772
def populate_cshuffle_gemm_universal(self):
771773
instance_code = """
772-
using EpilogueProblem = ck_tile::CShuffleEpilogueProblem<
773-
ADataType,
774-
BDataType,
775-
ck_tile::tuple<>, // DsDataType
776-
AccDataType,
777-
CDataType,
778-
ck_tile::tuple<>, // DsLayout
779-
CLayout,
780-
ck_tile::element_wise::PassThrough,
781-
TileM, // kM_
782-
TileN, // kN_
783-
WarpPerBlock_M, // MWave_
784-
WarpPerBlock_N, // NWave_
785-
WarpTileM, // MPerXdl_
786-
WarpTileN, // NPerXdl_
787-
WarpTileK, // KPerXdl_
788-
TransposeC, // isCTransposed_
789-
NumWaveGroups>; // kNumWaveGroups_
790-
791-
using GemmEpilogue = ck_tile::CShuffleEpilogue<EpilogueProblem>;"""
774+
using EpilogueProblem = ck_tile::CShuffleEpilogueProblem<
775+
ADataType,
776+
BDataType,
777+
ck_tile::tuple<>, // DsDataType
778+
AccDataType,
779+
CDataType,
780+
ck_tile::tuple<>, // DsLayout
781+
CLayout,
782+
ck_tile::element_wise::PassThrough,
783+
TileM, // kM_
784+
TileN, // kN_
785+
WarpPerBlock_M, // MWave_
786+
WarpPerBlock_N, // NWave_
787+
WarpTileM, // MPerXdl_
788+
WarpTileN, // NPerXdl_
789+
WarpTileK, // KPerXdl_
790+
TransposeC, // isCTransposed_
791+
NumWaveGroups>; // kNumWaveGroups_
792+
793+
using GemmEpilogue = ck_tile::CShuffleEpilogue<EpilogueProblem>;"""
792794
return instance_code
793795

794796
def populate_cshuffle_gemm_multi_d(self):
795797
instance_code = """
796-
using EpilogueProblem = ck_tile::CShuffleEpilogueProblem<
797-
ADataType,
798-
BDataType,
799-
DsDataType,
800-
AccDataType,
801-
CDataType,
802-
DsLayout,
803-
CLayout,
804-
ElementWiseFn,
805-
TileM, // kM_
806-
TileN, // kN_
807-
WarpPerBlock_M, // MWave_
808-
WarpPerBlock_N, // NWave_
809-
WarpTileM, // MPerXdl_
810-
WarpTileN, // NPerXdl_
811-
WarpTileK, // KPerXdl_
812-
TransposeC>; // isCTransposed_
813-
814-
using GemmEpilogue = ck_tile::CShuffleEpilogue<EpilogueProblem>;"""
798+
using EpilogueProblem = ck_tile::CShuffleEpilogueProblem<
799+
ADataType,
800+
BDataType,
801+
DsDataType,
802+
AccDataType,
803+
CDataType,
804+
DsLayout,
805+
CLayout,
806+
ElementWiseFn,
807+
TileM, // kM_
808+
TileN, // kN_
809+
WarpPerBlock_M, // MWave_
810+
WarpPerBlock_N, // NWave_
811+
WarpTileM, // MPerXdl_
812+
WarpTileN, // NPerXdl_
813+
WarpTileK, // KPerXdl_
814+
TransposeC>; // isCTransposed_
815+
816+
using GemmEpilogue = ck_tile::CShuffleEpilogue<EpilogueProblem>;"""
815817
return instance_code
816818

817819
def populate_cshuffle_gemm_preshuffle(self):
818820
instance_code = """
819-
using EpilogueProblem = ck_tile::CShuffleEpilogueProblem<
820-
ADataType,
821-
BDataType,
822-
ck_tile::tuple<>, // DsDataType
823-
AccDataType,
824-
CDataType,
825-
ck_tile::tuple<>, // DsLayout
826-
CLayout,
827-
ck_tile::element_wise::PassThrough,
828-
TileM, // kM_
829-
TileN, // kN_
830-
WarpPerBlock_M, // MWave_
831-
WarpPerBlock_N, // NWave_
832-
WarpTileM, // MPerXdl_
833-
WarpTileN, // NPerXdl_
834-
WarpTileK, // KPerXdl_
835-
TransposeC, // isCTransposed_
836-
NumWaveGroups, // kNumWaveGroups_
837-
false, // FixedVectorSize_
838-
1, // VectorSizeC_
839-
PermuteN>; // isPermuteN_
840-
841-
using GemmEpilogue = ck_tile::CShuffleEpilogue<EpilogueProblem>;"""
821+
using EpilogueProblem = ck_tile::CShuffleEpilogueProblem<
822+
ADataType,
823+
BDataType,
824+
ck_tile::tuple<>, // DsDataType
825+
AccDataType,
826+
CDataType,
827+
ck_tile::tuple<>, // DsLayout
828+
CLayout,
829+
ck_tile::element_wise::PassThrough,
830+
TileM, // kM_
831+
TileN, // kN_
832+
WarpPerBlock_M, // MWave_
833+
WarpPerBlock_N, // NWave_
834+
WarpTileM, // MPerXdl_
835+
WarpTileN, // NPerXdl_
836+
WarpTileK, // KPerXdl_
837+
TransposeC, // isCTransposed_
838+
NumWaveGroups, // kNumWaveGroups_
839+
false, // FixedVectorSize_
840+
1, // VectorSizeC_
841+
PermuteN>; // isPermuteN_
842+
843+
using GemmEpilogue = ck_tile::CShuffleEpilogue<EpilogueProblem>;"""
842844
return instance_code
843845

844846
def populate_default_gemm_universal(self):
845847
instance_code = """
846-
using EpilogueProblem = ck_tile::DefaultGemm2DEpilogueProblem<
847-
ADataType,
848-
BDataType,
849-
ck_tile::tuple<>, // DsDataType
850-
AccDataType,
851-
CDataType,
852-
ck_tile::tuple<>, // DsLayout
853-
CLayout,
854-
ck_tile::element_wise::PassThrough,
855-
TileM, // kM_
856-
TileN, // kN_
857-
kPadM,
858-
kPadN,
859-
WarpTileM, // kMPerXdl_
860-
WarpTileN, // kNPerXdl_
861-
WarpTileK, // kKPerXdl_
862-
TransposeC>; // isCTransposed_
863-
864-
using GemmEpilogue = ck_tile::DefaultGemm2DEpilogue<EpilogueProblem>;"""
848+
using EpilogueProblem = ck_tile::DefaultGemm2DEpilogueProblem<
849+
ADataType,
850+
BDataType,
851+
ck_tile::tuple<>, // DsDataType
852+
AccDataType,
853+
CDataType,
854+
ck_tile::tuple<>, // DsLayout
855+
CLayout,
856+
ck_tile::element_wise::PassThrough,
857+
TileM, // kM_
858+
TileN, // kN_
859+
kPadM,
860+
kPadN,
861+
WarpTileM, // kMPerXdl_
862+
WarpTileN, // kNPerXdl_
863+
WarpTileK, // kKPerXdl_
864+
TransposeC>; // isCTransposed_
865+
866+
using GemmEpilogue = ck_tile::DefaultGemm2DEpilogue<EpilogueProblem>;"""
865867
return instance_code
866868

867869
def populate_default_gemm_multi_d(self):
868870
instance_code = """
869-
using EpilogueProblem = ck_tile::DefaultGemm2DEpilogueProblem<
870-
ADataType,
871-
BDataType,
872-
DsDataType,
873-
AccDataType,
874-
CDataType,
875-
DsLayout,
876-
CLayout,
877-
ElementWiseFn,
878-
TileM, // kM_
879-
TileN, // kN_
880-
kPadM,
881-
kPadN,
882-
WarpTileM, // kMPerXdl_
883-
WarpTileN, // kNPerXdl_
884-
WarpTileK, // kKPerXdl_
885-
TransposeC>; // isCTransposed_
886-
887-
using GemmEpilogue = ck_tile::DefaultGemm2DEpilogue<EpilogueProblem>;"""
871+
using EpilogueProblem = ck_tile::DefaultGemm2DEpilogueProblem<
872+
ADataType,
873+
BDataType,
874+
DsDataType,
875+
AccDataType,
876+
CDataType,
877+
DsLayout,
878+
CLayout,
879+
ElementWiseFn,
880+
TileM, // kM_
881+
TileN, // kN_
882+
kPadM,
883+
kPadN,
884+
WarpTileM, // kMPerXdl_
885+
WarpTileN, // kNPerXdl_
886+
WarpTileK, // kKPerXdl_
887+
TransposeC>; // isCTransposed_
888+
889+
using GemmEpilogue = ck_tile::DefaultGemm2DEpilogue<EpilogueProblem>;"""
888890
return instance_code
889891

890892
def populate_default_gemm_preshuffle(self):
891893
instance_code = """
892-
using EpilogueProblem = ck_tile::DefaultGemm2DEpilogueProblem<
893-
ADataType,
894-
BDataType,
895-
ck_tile::tuple<>, // DsDataType
896-
AccDataType,
897-
CDataType,
898-
ck_tile::tuple<>, // DsLayout
899-
CLayout,
900-
ck_tile::element_wise::PassThrough,
901-
TileM, // kM_
902-
TileN, // kN_
903-
kPadM,
904-
kPadN,
905-
WarpTileM, // kMPerXdl_
906-
WarpTileN, // kNPerXdl_
907-
WarpTileK, // kKPerXdl_
908-
TransposeC>; // isCTransposed_
909-
910-
using GemmEpilogue = ck_tile::DefaultGemm2DEpilogue<EpilogueProblem>;"""
894+
using EpilogueProblem = ck_tile::DefaultGemm2DEpilogueProblem<
895+
ADataType,
896+
BDataType,
897+
ck_tile::tuple<>, // DsDataType
898+
AccDataType,
899+
CDataType,
900+
ck_tile::tuple<>, // DsLayout
901+
CLayout,
902+
ck_tile::element_wise::PassThrough,
903+
TileM, // kM_
904+
TileN, // kN_
905+
kPadM,
906+
kPadN,
907+
WarpTileM, // kMPerXdl_
908+
WarpTileN, // kNPerXdl_
909+
WarpTileK, // kKPerXdl_
910+
TransposeC>; // isCTransposed_
911+
912+
using GemmEpilogue = ck_tile::DefaultGemm2DEpilogue<EpilogueProblem>;"""
911913
return instance_code
912914

913915
def _generate_cmake_individual_targets(self, kernel_list):

0 commit comments

Comments
 (0)