Skip to content

Commit f9c6ba0

Browse files
Implement grouped gemm fastgelu for RDNA4 (#3303)
* Implement grouped gemm fastgelu for RDNA4 * chore: some cleanup and minor inconsistencies in grouped gemm profiler * chore: clarified logic and reporting of supported instance warnings
1 parent a7d6b1e commit f9c6ba0

24 files changed

Lines changed: 665 additions & 399 deletions

File tree

include/ck/tensor_operation/gpu/device/impl/device_grouped_gemm_wmma_splitk_cshuffle_v3.hpp

Lines changed: 72 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
#include <sstream>
88

99
#include "ck/ck.hpp"
10+
#include "ck/tensor_operation/gpu/element/unary_element_wise_operation.hpp"
1011
#include "ck/utility/env.hpp"
1112
#include "ck/host_utility/device_prop.hpp"
1213
#include "ck/host_utility/kernel_launch.hpp"
@@ -242,7 +243,6 @@ struct DeviceGroupedGemm_Wmma_CShuffleV3 : public DeviceGroupedGemmSplitK<ALayou
242243
static constexpr index_t B2E_M01 = 8;
243244
using GroupedGemmBlock2ETileMap = OffsettedBlockToCTileMap<Block2ETileMapKSplit>;
244245
using KernelArgument = typename GridwiseGemm::Argument;
245-
using PassThrough = ck::tensor_operation::element_wise::PassThrough;
246246
template <typename KernelArgument_>
247247
struct GemmTransKernelArgBase
248248
{
@@ -274,23 +274,38 @@ struct DeviceGroupedGemm_Wmma_CShuffleV3 : public DeviceGroupedGemmSplitK<ALayou
274274
}
275275

276276
// Argument
277-
// TODO: Add A/B/CDE element op?
278277
struct Argument : public BaseArgument
279278
{
280279

281280
Argument(std::vector<const void*>& p_As,
282281
std::vector<const void*>& p_Bs,
282+
std::vector<std::array<const void*, NumDTensor>>& p_Ds,
283283
std::vector<void*>& p_Es,
284-
std::vector<GemmDesc>& gemm_descs)
285-
: Argument(p_As, p_Bs, p_Es, gemm_descs, DefaultKBatch)
284+
std::vector<GemmDesc>& gemm_descs,
285+
AElementwiseOperation a_element_op,
286+
BElementwiseOperation b_element_op,
287+
CDEElementwiseOperation c_element_op)
288+
: Argument(p_As,
289+
p_Bs,
290+
p_Ds,
291+
p_Es,
292+
gemm_descs,
293+
a_element_op,
294+
b_element_op,
295+
c_element_op,
296+
DefaultKBatch)
286297
{
287298
// TODO: use occupancy api to calculate appropriate batch size.
288299
}
289300

290301
Argument(std::vector<const void*>& p_As,
291302
std::vector<const void*>& p_Bs,
303+
std::vector<std::array<const void*, NumDTensor>>& p_Ds,
292304
std::vector<void*>& p_Es,
293305
std::vector<GemmDesc>& gemm_descs,
306+
AElementwiseOperation a_element_op,
307+
BElementwiseOperation b_element_op,
308+
CDEElementwiseOperation c_element_op,
294309
index_t kbatch)
295310
: K_BATCH{kbatch}, gemm_kernel_host_args_{nullptr}
296311
{
@@ -299,9 +314,11 @@ struct DeviceGroupedGemm_Wmma_CShuffleV3 : public DeviceGroupedGemmSplitK<ALayou
299314

300315
if(!(group_count_ == ck::type_convert<ck::index_t>(p_As.size()) &&
301316
group_count_ == ck::type_convert<ck::index_t>(p_Bs.size()) &&
317+
((NumDTensor == 0 && p_Ds.size() == 0) ||
318+
group_count_ == ck::type_convert<ck::index_t>(p_Ds.size())) &&
302319
group_count_ == ck::type_convert<ck::index_t>(p_Es.size())))
303320
{
304-
throw std::runtime_error("wrong! group_count_ != p_As/b/c.size");
321+
throw std::runtime_error("wrong! group_count_ != p_As/b/d/e.size");
305322
}
306323

307324
gemm_kernel_args_.reserve(group_count_);
@@ -320,9 +337,22 @@ struct DeviceGroupedGemm_Wmma_CShuffleV3 : public DeviceGroupedGemmSplitK<ALayou
320337
continue;
321338
}
322339

323-
const index_t stride_a = gemm_descs[i].stride_A_;
324-
const index_t stride_b = gemm_descs[i].stride_B_;
325-
const index_t stride_c = gemm_descs[i].stride_C_;
340+
const index_t stride_a = gemm_descs[i].stride_A_;
341+
const index_t stride_b = gemm_descs[i].stride_B_;
342+
const index_t stride_c = gemm_descs[i].stride_C_;
343+
const auto& stride_d_vec = gemm_descs[i].stride_Ds_;
344+
345+
if(!(NumDTensor == ck::type_convert<ck::index_t>(stride_d_vec.size())))
346+
{
347+
throw std::runtime_error("wrong! stride D mismatch");
348+
}
349+
350+
// Copy D stride vector to fixed-size array
351+
std::array<index_t, NumDTensor> stride_ds;
352+
if constexpr(NumDTensor > 0)
353+
{
354+
std::copy(stride_d_vec.begin(), stride_d_vec.end(), stride_ds);
355+
}
326356

327357
const index_t m_padded = GridwiseGemm::CalculateMPadded(M);
328358
const index_t n_padded = GridwiseGemm::CalculateNPadded(N);
@@ -346,19 +376,19 @@ struct DeviceGroupedGemm_Wmma_CShuffleV3 : public DeviceGroupedGemmSplitK<ALayou
346376

347377
auto karg = KernelArgument(std::array<const void*, 1>{p_As[i]},
348378
std::array<const void*, 1>{p_Bs[i]},
349-
std::array<const void*, 0>{}, // p_ds_grid_
379+
p_Ds[i],
350380
type_convert<EDataType*>(p_Es[i]),
351381
M,
352382
N,
353383
K,
354384
std::array<index_t, 1>{stride_a},
355385
std::array<index_t, 1>{stride_b},
356-
std::array<index_t, 0>{}, // StrideDs_
386+
stride_ds,
357387
stride_c,
358388
K_BATCH,
359-
PassThrough{},
360-
PassThrough{},
361-
PassThrough{},
389+
a_element_op,
390+
b_element_op,
391+
c_element_op,
362392
false);
363393

364394
gemm_kernel_args_.emplace_back(
@@ -632,6 +662,23 @@ struct DeviceGroupedGemm_Wmma_CShuffleV3 : public DeviceGroupedGemmSplitK<ALayou
632662
}
633663
}
634664

665+
if constexpr(!std::is_same_v<CDEElementwiseOperation,
666+
ck::tensor_operation::element_wise::PassThrough>)
667+
{
668+
if(arg.K_BATCH > 1)
669+
{
670+
// Using SplitK and a C element op would require a two stage kernel where the second
671+
// stage applies the op on the accumulated results
672+
if(ck::EnvIsEnabled(CK_ENV(CK_LOGGING)))
673+
{
674+
std::cout << "C element operators are not supported when using SplitK. Set "
675+
"K_BATCH to 1 or remove the operator."
676+
<< std::endl;
677+
}
678+
return false;
679+
}
680+
}
681+
635682
if constexpr(std::is_same_v<ComputeTypeA, f8_t> || std::is_same_v<ComputeTypeA, bf8_t> ||
636683
std::is_same_v<ComputeTypeB, f8_t> || std::is_same_v<ComputeTypeB, bf8_t>)
637684
{
@@ -681,14 +728,15 @@ struct DeviceGroupedGemm_Wmma_CShuffleV3 : public DeviceGroupedGemmSplitK<ALayou
681728

682729
static auto MakeArgument(std::vector<const void*>& p_As,
683730
std::vector<const void*>& p_Bs,
684-
std::vector<std::array<const void*, NumDTensor>>&,
731+
std::vector<std::array<const void*, NumDTensor>>& p_Ds,
685732
std::vector<void*>& p_Es,
686733
std::vector<GemmDesc> gemm_descs,
687-
AElementwiseOperation,
688-
BElementwiseOperation,
689-
CDEElementwiseOperation)
734+
AElementwiseOperation a_element_op,
735+
BElementwiseOperation b_element_op,
736+
CDEElementwiseOperation c_element_op)
690737
{
691-
return Argument{p_As, p_Bs, p_Es, gemm_descs};
738+
return Argument{
739+
p_As, p_Bs, p_Ds, p_Es, gemm_descs, a_element_op, b_element_op, c_element_op};
692740
}
693741

694742
static auto MakeInvoker() { return Invoker{}; }
@@ -697,14 +745,15 @@ struct DeviceGroupedGemm_Wmma_CShuffleV3 : public DeviceGroupedGemmSplitK<ALayou
697745
std::unique_ptr<BaseArgument>
698746
MakeArgumentPointer(std::vector<const void*>& p_As,
699747
std::vector<const void*>& p_Bs,
700-
std::vector<std::array<const void*, NumDTensor>>&,
748+
std::vector<std::array<const void*, NumDTensor>>& p_Ds,
701749
std::vector<void*>& p_Es,
702750
std::vector<GemmDesc>& gemm_descs,
703-
AElementwiseOperation,
704-
BElementwiseOperation,
705-
CDEElementwiseOperation) override
751+
AElementwiseOperation a_element_op,
752+
BElementwiseOperation b_element_op,
753+
CDEElementwiseOperation c_element_op) override
706754
{
707-
return std::make_unique<Argument>(p_As, p_Bs, p_Es, gemm_descs);
755+
return std::make_unique<Argument>(
756+
p_As, p_Bs, p_Ds, p_Es, gemm_descs, a_element_op, b_element_op, c_element_op);
708757
}
709758

710759
// polymorphic

library/include/ck/library/tensor_operation_instance/gpu/grouped_gemm/device_grouped_gemm_wmma_splitk_instance.hpp

Lines changed: 67 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -31,17 +31,14 @@ using S = ck::Sequence<Is...>;
3131

3232
using Empty_Tuple = ck::Tuple<>;
3333
using PassThrough = ck::tensor_operation::element_wise::PassThrough;
34+
using FastGelu = ck::tensor_operation::element_wise::FastGelu;
3435

3536
using AccDataType = F32;
3637
using DsDataType = Empty_Tuple;
3738

3839
using DsLayout = Empty_Tuple;
3940
using ELayout = Row;
4041

41-
using AElementOp = PassThrough;
42-
using BElementOp = PassThrough;
43-
using CDEElementOp = PassThrough;
44-
4542
static constexpr auto PipelineV1 = BlockGemmPipelineVersion::v1;
4643
static constexpr auto PipelineV3 = BlockGemmPipelineVersion::v3;
4744
static 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>
5858
using 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>
7780
using 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>
95101
using 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>
114123
using 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
129153
template <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>
138168
void 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>
175210
void 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

Comments
 (0)