@@ -31,17 +31,14 @@ using S = ck::Sequence<Is...>;
3131
3232using Empty_Tuple = ck::Tuple<>;
3333using PassThrough = ck::tensor_operation::element_wise::PassThrough;
34+ using FastGelu = ck::tensor_operation::element_wise::FastGelu;
3435
3536using AccDataType = F32 ;
3637using DsDataType = Empty_Tuple;
3738
3839using DsLayout = Empty_Tuple;
3940using ELayout = Row;
4041
41- using AElementOp = PassThrough;
42- using BElementOp = PassThrough;
43- using CDEElementOp = PassThrough;
44-
4542static constexpr auto PipelineV1 = BlockGemmPipelineVersion::v1;
4643static constexpr auto PipelineV3 = BlockGemmPipelineVersion::v3;
4744static constexpr auto IntrawaveScheduler = BlockGemmPipelineScheduler::Intrawave;
@@ -54,6 +51,9 @@ template <typename T,
5451 device::GemmSpecialization GemmSpec,
5552 BlockGemmPipelineScheduler BlkGemmPipeSched,
5653 BlockGemmPipelineVersion BlkGemmPipelineVer,
54+ typename AElementOp,
55+ typename BElementOp,
56+ typename CDEElementOp,
5757 enable_if_t <sizeof (T) == 2 , bool > = false >
5858using device_grouped_gemm_wmma_universal_km_kn_mn_instances =
5959 std::tuple<
@@ -73,6 +73,9 @@ template <typename T,
7373 device::GemmSpecialization GemmSpec,
7474 BlockGemmPipelineScheduler BlkGemmPipeSched,
7575 BlockGemmPipelineVersion BlkGemmPipelineVer,
76+ typename AElementOp,
77+ typename BElementOp,
78+ typename CDEElementOp,
7679 enable_if_t <sizeof (T) == 2 , bool > = false >
7780using device_grouped_gemm_wmma_universal_km_nk_mn_instances = std::tuple<
7881 // clang-format off
@@ -91,6 +94,9 @@ template <typename T,
9194 device::GemmSpecialization GemmSpec,
9295 BlockGemmPipelineScheduler BlkGemmPipeSched,
9396 BlockGemmPipelineVersion BlkGemmPipelineVer,
97+ typename AElementOp,
98+ typename BElementOp,
99+ typename CDEElementOp,
94100 enable_if_t <sizeof (T) == 2 , bool > = false >
95101using device_grouped_gemm_wmma_universal_mk_kn_mn_instances =
96102 std::tuple<
@@ -110,6 +116,9 @@ template <typename T,
110116 device::GemmSpecialization GemmSpec,
111117 BlockGemmPipelineScheduler BlkGemmPipeSched,
112118 BlockGemmPipelineVersion BlkGemmPipelineVer,
119+ typename AElementOp,
120+ typename BElementOp,
121+ typename CDEElementOp,
113122 enable_if_t <sizeof (T) == 2 , bool > = false >
114123using device_grouped_gemm_wmma_universal_mk_nk_mn_instances =
115124 std::tuple<
@@ -124,17 +133,38 @@ using device_grouped_gemm_wmma_universal_mk_nk_mn_instances =
124133 // clang-format on
125134 >;
126135
136+ // List of instance variants to add (pipeline/scheduler/padding combinations)
137+ // Some are disabled now, can be re-enabled if needed
138+ using InstanceVariant =
139+ ck::Tuple<device::GemmSpecialization, BlockGemmPipelineScheduler, BlockGemmPipelineVersion>;
140+ static constexpr InstanceVariant InstanceVariants[] = {
141+
142+ make_tuple (GemmDefault, IntrawaveScheduler, PipelineV1),
143+ // make_tuple(GemmDefault, InterwaveScheduler, PipelineV1),
144+ make_tuple (GemmDefault, IntrawaveScheduler, PipelineV3),
145+
146+ make_tuple (GemmMNKPadding, IntrawaveScheduler, PipelineV1),
147+ // make_tuple(GemmMNKPadding, InterwaveScheduler, PipelineV1),
148+ // make_tuple(GemmMNKPadding, IntrawaveScheduler, PipelineV3),
149+ };
150+
127151// Helper function to add a list of layout instances with specific A/B/E datatypes for all supported
128152// padding/scheduler/pipeline version combinations
129153template <typename ALayout,
130154 typename BLayout,
131155 template <device::GemmSpecialization GemmSpec,
132156 BlockGemmPipelineScheduler BlkGemmPipeSched,
133- BlockGemmPipelineVersion BlkGemmPipelineVer>
157+ BlockGemmPipelineVersion BlkGemmPipelineVer,
158+ typename AElementOp,
159+ typename BElementOp,
160+ typename CDEElementOp>
134161 typename LayoutInstances,
135162 typename ADataType, // NOTE: type parameters as last so that they can be inferred from the
136163 typename BDataType, // vector argument
137- typename EDataType>
164+ typename EDataType,
165+ typename AElementOp,
166+ typename BElementOp,
167+ typename CDEElementOp>
138168void add_device_grouped_gemm_wmma_universal_instances (
139169 std::vector<std::unique_ptr<DeviceGroupedGemm<ALayout,
140170 BLayout,
@@ -148,18 +178,17 @@ void add_device_grouped_gemm_wmma_universal_instances(
148178 BElementOp,
149179 CDEElementOp>>>& instances)
150180{
151- add_device_operation_instances (instances,
152- LayoutInstances<GemmDefault, IntrawaveScheduler, PipelineV1>{});
153- add_device_operation_instances (instances,
154- LayoutInstances<GemmDefault, InterwaveScheduler, PipelineV1>{});
155- add_device_operation_instances (instances,
156- LayoutInstances<GemmDefault, IntrawaveScheduler, PipelineV3>{});
157- add_device_operation_instances (
158- instances, LayoutInstances<GemmMNKPadding, IntrawaveScheduler, PipelineV1>{});
159- add_device_operation_instances (
160- instances, LayoutInstances<GemmMNKPadding, InterwaveScheduler, PipelineV1>{});
161- add_device_operation_instances (
162- instances, LayoutInstances<GemmMNKPadding, IntrawaveScheduler, PipelineV3>{});
181+ // Add all instances from our instance list
182+ static_for<0 , std::size (InstanceVariants), 1 >{}([&](auto i) {
183+ constexpr auto instance = InstanceVariants[i];
184+ add_device_operation_instances (instances,
185+ LayoutInstances<instance.At (Number<0 >{}),
186+ instance.At (Number<1 >{}),
187+ instance.At (Number<2 >{}),
188+ AElementOp,
189+ BElementOp,
190+ CDEElementOp>{});
191+ });
163192}
164193
165194// Helper function to add a list of layout instances for instances with matching A/B/E data types
@@ -170,8 +199,14 @@ template <typename T,
170199 template <typename T2 ,
171200 device::GemmSpecialization GemmSpec,
172201 BlockGemmPipelineScheduler BlkGemmPipeSched,
173- BlockGemmPipelineVersion BlkGemmPipelineVer>
174- typename LayoutInstances>
202+ BlockGemmPipelineVersion BlkGemmPipelineVer,
203+ typename AElementOp,
204+ typename BElementOp,
205+ typename CDEElementOp>
206+ typename LayoutInstances,
207+ typename AElementOp, // NOTE: element-wise op parameters as last so that they can be
208+ typename BElementOp, // inferred from the vector argument
209+ typename CDEElementOp>
175210void add_device_grouped_gemm_wmma_universal_instances (
176211 std::vector<std::unique_ptr<DeviceGroupedGemm<ALayout,
177212 BLayout,
@@ -185,18 +220,18 @@ void add_device_grouped_gemm_wmma_universal_instances(
185220 BElementOp,
186221 CDEElementOp>>>& instances)
187222{
188- add_device_operation_instances (
189- instances, LayoutInstances<T, GemmDefault, IntrawaveScheduler, PipelineV1 >{});
190- add_device_operation_instances (
191- instances, LayoutInstances<T, GemmDefault, InterwaveScheduler, PipelineV1>{});
192- add_device_operation_instances (
193- instances, LayoutInstances<T, GemmDefault, IntrawaveScheduler, PipelineV3 >{});
194- add_device_operation_instances (
195- instances, LayoutInstances<T, GemmMNKPadding, IntrawaveScheduler, PipelineV1 >{});
196- add_device_operation_instances (
197- instances, LayoutInstances<T, GemmMNKPadding, InterwaveScheduler, PipelineV1>{});
198- add_device_operation_instances (
199- instances, LayoutInstances<T, GemmMNKPadding, IntrawaveScheduler, PipelineV3>{ });
223+ // Add all instances from our instance list
224+ static_for< 0 , std::size (InstanceVariants), 1 >{}([&]( auto i) {
225+ constexpr auto instance = InstanceVariants[i];
226+ add_device_operation_instances ( instances,
227+ LayoutInstances<T,
228+ instance. At (Number< 0 >{}),
229+ instance. At (Number< 1 >{}),
230+ instance. At (Number< 2 >{}),
231+ AElementOp,
232+ BElementOp,
233+ CDEElementOp>{});
234+ });
200235}
201236
202237} // namespace instance
0 commit comments