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