Skip to content

Commit aad4cf0

Browse files
authored
Wmma support for gemm_bias_add_reduce (#3316)
* Add tests for gemm_bias_add_reduce * Initial working implementation * Generalize implementation of reduce epilogue * Add tests for all layouts * Add instances * Fix test archs * Fix xdl bug * Remove library/profiler duplications * Fix num_byted error profiler * Fix typos * Fix copyright
1 parent f9c6ba0 commit aad4cf0

15 files changed

Lines changed: 1426 additions & 143 deletions

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

Lines changed: 682 additions & 0 deletions
Large diffs are not rendered by default.

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

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -49,8 +49,11 @@ __launch_bounds__(CK_MAX_THREAD_PER_BLOCK, MinimumOccupancy)
4949

5050
auto splitk_batch_offset = typename GridwiseGemm::SplitKBatchOffset(karg, blockIdx.z);
5151

52-
auto epilogue_args =
53-
EpilogueType(p_reduces_grid, reduce_in_element_ops, reduce_out_element_ops, karg.M);
52+
auto epilogue_args = EpilogueType(p_reduces_grid,
53+
reduce_in_element_ops,
54+
reduce_out_element_ops,
55+
karg.M,
56+
tensor_operation::element_wise::PassThrough{});
5457

5558
GridwiseGemm::template Run<HasMainKBlockLoop, EGlobalMemoryDataOperation, TailNum>(
5659
p_shared, splitk_batch_offset, karg, epilogue_args);
@@ -188,6 +191,7 @@ struct DeviceGemmReduce_Wmma_CShuffleV3 : public DeviceGemmReduce<0, ReduceOpera
188191

189192
using ReduceTrait = ReduceTrait_<ReduceAccDataType,
190193
ReducePtrsGlobal,
194+
tensor_operation::element_wise::PassThrough,
191195
ReduceOperations,
192196
ReduceInElementwiseOperations,
193197
ReduceAccElementwiseOperations,

include/ck/tensor_operation/gpu/grid/epilogue_cshuffle_v3_reduce_wmma.hpp

Lines changed: 128 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@ namespace ck {
1010

1111
template <typename ReduceAccDataType,
1212
typename ReducePtrsGlobal,
13+
typename D0ElementwiseOperation,
1314
typename ReduceOperations,
1415
typename ReduceInElementwiseOperations,
1516
typename ReduceAccElementwiseOperations,
@@ -21,6 +22,7 @@ struct ReduceTrait_
2122
{
2223
using ReduceAccDataType_ = ReduceAccDataType;
2324
using ReducePtrsGlobal_ = ReducePtrsGlobal;
25+
using D0ElementwiseOperation_ = D0ElementwiseOperation;
2426
using ReduceOperations_ = ReduceOperations;
2527
using ReduceInElementwiseOperations_ = ReduceInElementwiseOperations;
2628
using ReduceAccElementwiseOperations_ = ReduceAccElementwiseOperations;
@@ -148,11 +150,13 @@ struct EpilogueReduceCShuffle
148150
typename ReduceTrait::ReducePtrsGlobal_ p_reduces_grid_,
149151
const typename ReduceTrait::ReduceInElementwiseOperations_ reduce_in_element_ops_,
150152
const typename ReduceTrait::ReduceAccElementwiseOperations_ reduce_out_element_ops_,
151-
const index_t MRaw_)
153+
const index_t MRaw_,
154+
const typename ReduceTrait::D0ElementwiseOperation_ d0_element_op_)
152155
: p_reduces_grid(p_reduces_grid_),
153156
reduce_in_element_ops(reduce_in_element_ops_),
154157
reduce_out_element_ops(reduce_out_element_ops_),
155158
MRaw(MRaw_),
159+
d0_element_op{d0_element_op_},
156160
reduce_grid_desc_m{MakeReduceGridDescriptor_M(MRaw)}
157161
{
158162
}
@@ -174,6 +178,13 @@ struct EpilogueReduceCShuffle
174178
const index_t& block_m_id,
175179
const index_t& block_n_id)
176180
{
181+
// HACK: this force m/n_block_data_idx_on_grid into SGPR
182+
const index_t m_block_data_idx_on_grid =
183+
__builtin_amdgcn_readfirstlane(block_m_id * MPerBlock);
184+
185+
const index_t n_block_data_idx_on_grid =
186+
__builtin_amdgcn_readfirstlane(block_n_id * NPerBlock);
187+
177188
auto reduce_grid_desc_mblock_mperblock =
178189
MakeReduceGridDescriptor_MBlock_MPerBlock(reduce_grid_desc_m);
179190

@@ -216,29 +227,6 @@ struct EpilogueReduceCShuffle
216227
c_block_desc_mrepeat_mwave_msubgroup_nrepeat_nwave_nthreadpersubgroup_maccvgprs =
217228
GetCShuffleLDSDescriptor();
218229

219-
// tuple of reference to C/Ds tensor descriptors
220-
const auto c_ds_desc_refs = concat_tuple_of_reference(
221-
tie(c_shuffle_block_desc_mshrepeat_mpershrepeat_nshrepeat_npershrepeat),
222-
generate_tie([&](auto i) -> const auto& // return type should be reference
223-
{ return ds_grid_desc_mblock_mperblock_nblock_nperblock[i]; },
224-
Number<NumDTensor>{}));
225-
226-
// Thread transfer LDS to Vmem
227-
auto cde_shuffle_block_copy_lds_and_global =
228-
Base::template GetLDSToVmemEpilogueDescriptor<EGlobalMemoryDataOperation, EDataType>(
229-
c_ds_desc_refs,
230-
e_grid_desc_mblock_mperblock_nblock_nperblock,
231-
cde_element_op,
232-
block_m_id,
233-
block_n_id);
234-
235-
// tuple of reference to C/Ds tensor buffers
236-
const auto c_ds_buf_refs = concat_tuple_of_reference(
237-
tie(c_shuffle_block_buf),
238-
generate_tie([&](auto i) -> const auto& // return type should be reference
239-
{ return ds_grid_buf[i]; },
240-
Number<NumDTensor>{}));
241-
242230
// LDS c_reduce_block_desc_mperblock_nperblock
243231
constexpr auto c_reduce_block_desc_mperblock_nperblock = transform_tensor_descriptor(
244232
c_shuffle_block_desc_mshrepeat_mpershrepeat_nshrepeat_npershrepeat,
@@ -346,6 +334,68 @@ struct EpilogueReduceCShuffle
346334
},
347335
Number<NumReduce>{});
348336

337+
// multiple Ds
338+
constexpr auto d_reduce_thread_desc_mblock_mperblock_nblock_nperblock =
339+
make_naive_tensor_descriptor_packed(
340+
make_tuple(I1, Number<mreduce_per_thread>{}, I1, Number<nreduce_per_thread>{}));
341+
342+
constexpr auto ds_reduce_thread_desc_mblock_mperblock_nblock_nperblock = generate_tuple(
343+
[&](auto) { return d_reduce_thread_desc_mblock_mperblock_nblock_nperblock; },
344+
Number<NumDTensor>{});
345+
346+
constexpr auto ds_thread_buf_size =
347+
d_reduce_thread_desc_mblock_mperblock_nblock_nperblock.GetElementSpaceSize();
348+
349+
auto c01_thread_buf =
350+
make_static_buffer<AddressSpaceEnum::Vgpr, typename ReduceTrait::ReduceAccDataType_>(
351+
Number<ds_thread_buf_size>{});
352+
353+
auto ds_thread_copy_global_to_vgpr = generate_tuple(
354+
[&](auto I) {
355+
return ThreadwiseTensorSliceTransfer_v2<
356+
remove_cvref_t<tuple_element_t<I.value, DsDataType>>,
357+
typename ReduceTrait::ReduceAccDataType_,
358+
decltype(ds_grid_desc_mblock_mperblock_nblock_nperblock[I]),
359+
remove_cvref_t<
360+
decltype(ds_reduce_thread_desc_mblock_mperblock_nblock_nperblock[I])>,
361+
Sequence<I1, mreduce_per_thread, I1, nreduce_per_thread>,
362+
Sequence<0, 1, 2, 3>,
363+
3,
364+
ReduceTrait::CReduceThreadLds2VGprCopySrcDstScalarPerVector_NPerBlock_,
365+
1,
366+
true>(ds_grid_desc_mblock_mperblock_nblock_nperblock[I],
367+
make_multi_index(
368+
I0,
369+
m_block_data_idx_on_grid + c_reduce_thread_data_idx_begin[I0],
370+
I0,
371+
n_block_data_idx_on_grid + c_reduce_thread_data_idx_begin[I1]));
372+
},
373+
Number<NumDTensor>{});
374+
375+
constexpr auto c_reduce_thread_desc_mblock_mperblock_nblock_nperblock =
376+
make_naive_tensor_descriptor_packed(
377+
make_tuple(I1, Number<mreduce_per_thread>{}, I1, Number<nreduce_per_thread>{}));
378+
379+
// Write E from Vgpr to Vmem
380+
auto c_reduce_thread_copy_vgpr_to_global = ThreadwiseTensorSliceTransfer_v1r3<
381+
typename ReduceTrait::ReduceAccDataType_,
382+
EDataType,
383+
decltype(c_reduce_thread_desc_mblock_mperblock_nblock_nperblock),
384+
decltype(e_grid_desc_mblock_mperblock_nblock_nperblock),
385+
tensor_operation::element_wise::PassThrough,
386+
Sequence<I1, mreduce_per_thread, I1, nreduce_per_thread>, // SliceLengths
387+
Sequence<0, 1, 2, 3>, // DimAccessOrder
388+
3, // DstVectorDim
389+
ReduceTrait::CReduceThreadLds2VGprCopySrcDstScalarPerVector_NPerBlock_,
390+
EGlobalMemoryDataOperation,
391+
1,
392+
true>{e_grid_desc_mblock_mperblock_nblock_nperblock,
393+
make_multi_index(I0,
394+
m_block_data_idx_on_grid + c_reduce_thread_data_idx_begin[I0],
395+
I0,
396+
n_block_data_idx_on_grid + c_reduce_thread_data_idx_begin[I1]),
397+
NumDTensor > 0 ? tensor_operation::element_wise::PassThrough{} : cde_element_op};
398+
349399
constexpr index_t num_access = sfc_c_vgpr.GetNumOfAccess();
350400

351401
static_assert(num_access == sfc_cde_global.GetNumOfAccess(), "wrong!");
@@ -365,22 +415,60 @@ struct EpilogueReduceCShuffle
365415

366416
// make sure it's safe to read from LDS
367417
block_sync_lds();
368-
369-
// each block loads its C data from LDS, D from global, applies elementwise
370-
// operation and stores result E to global
371-
cde_shuffle_block_copy_lds_and_global.Run(
372-
c_ds_desc_refs,
373-
c_ds_buf_refs,
374-
tie(e_grid_desc_mblock_mperblock_nblock_nperblock),
375-
tie(e_grid_buf));
376-
377418
{
378419
c_reduce_thread_copy_lds_to_vgpr.Run(c_reduce_block_desc_mperblock_nperblock,
379420
c_shuffle_block_buf,
380421
c_reduce_thread_desc_mperblock_nperblock,
381422
make_tuple(I0, I0),
382423
c_reduce_thread_buf);
383424

425+
// Note: currently multiple Ds supports only Bias + Add.
426+
// It needs to be generalized for other operations (currently not needed)
427+
if constexpr(NumDTensor > 0)
428+
{
429+
auto& d0_thread_copy_global_to_vgpr = ds_thread_copy_global_to_vgpr(I0);
430+
// d0 / d1 operations
431+
d0_thread_copy_global_to_vgpr.Run(
432+
ds_grid_desc_mblock_mperblock_nblock_nperblock[I0],
433+
ds_grid_buf[I0],
434+
ds_reduce_thread_desc_mblock_mperblock_nblock_nperblock[I0],
435+
make_tuple(I0, I0, I0, I0),
436+
c01_thread_buf);
437+
438+
// c = activation(c + bias)
439+
static_for<0, c_reduce_thread_desc_mperblock_nperblock.GetElementSize(), 1>{}(
440+
[&](auto i) {
441+
typename ReduceTrait::ReduceAccDataType_ out;
442+
cde_element_op(out, c_reduce_thread_buf(i) + c01_thread_buf(i));
443+
c_reduce_thread_buf(i) = out;
444+
});
445+
446+
auto& d1_thread_copy_global_to_vgpr = ds_thread_copy_global_to_vgpr(I1);
447+
448+
d1_thread_copy_global_to_vgpr.Run(
449+
ds_grid_desc_mblock_mperblock_nblock_nperblock[I1],
450+
ds_grid_buf[I1],
451+
ds_reduce_thread_desc_mblock_mperblock_nblock_nperblock[I1],
452+
make_tuple(I0, I0, I0, I0),
453+
c01_thread_buf);
454+
455+
// c = c + c1_function(c1)
456+
static_for<0, c_reduce_thread_desc_mperblock_nperblock.GetElementSize(), 1>{}(
457+
[&](auto i) {
458+
d0_element_op(c01_thread_buf(i), c01_thread_buf(i));
459+
c_reduce_thread_buf(i) += c01_thread_buf(i);
460+
});
461+
}
462+
463+
// Write E
464+
c_reduce_thread_copy_vgpr_to_global.Run(
465+
c_reduce_thread_desc_mblock_mperblock_nblock_nperblock,
466+
make_tuple(I0, I0, I0, I0),
467+
c_reduce_thread_buf,
468+
e_grid_desc_mblock_mperblock_nblock_nperblock,
469+
e_grid_buf);
470+
471+
// Reduction
384472
static_for<0, NumReduce, 1>{}([&](auto In) {
385473
auto& p_reduce_grid = p_reduces_grid[In];
386474

@@ -448,14 +536,15 @@ struct EpilogueReduceCShuffle
448536
{
449537
constexpr auto cde_global_step = sfc_cde_global.GetForwardStep(access_id);
450538
// move on Ds
451-
static_for<0, NumDTensor, 1>{}([&](auto i) {
452-
cde_shuffle_block_copy_lds_and_global.MoveSrcSliceWindow(
453-
c_ds_desc_refs, i + I1, cde_global_step);
539+
static_for<0, NumDTensor, 1>{}([&](auto I) {
540+
auto& d_thread_copy_global_to_vgpr = ds_thread_copy_global_to_vgpr(I);
541+
d_thread_copy_global_to_vgpr.MoveSrcSliceWindow(
542+
ds_grid_desc_mblock_mperblock_nblock_nperblock[I], cde_global_step);
454543
});
455544

456545
// move on E
457-
cde_shuffle_block_copy_lds_and_global.MoveDstSliceWindow(
458-
tie(e_grid_desc_mblock_mperblock_nblock_nperblock), cde_global_step);
546+
c_reduce_thread_copy_vgpr_to_global.MoveDstSliceWindow(
547+
e_grid_desc_mblock_mperblock_nblock_nperblock, cde_global_step);
459548
}
460549
});
461550
}
@@ -464,6 +553,7 @@ struct EpilogueReduceCShuffle
464553
typename ReduceTrait::ReduceInElementwiseOperations_ reduce_in_element_ops;
465554
typename ReduceTrait::ReduceAccElementwiseOperations_ reduce_out_element_ops;
466555
index_t MRaw;
556+
typename ReduceTrait::D0ElementwiseOperation_ d0_element_op;
467557
ReduceGridDesc_M reduce_grid_desc_m;
468558
};
469559

include/ck/tensor_operation/gpu/grid/gridwise_gemm_bias_add_reduce_xdl_cshuffle_v1.hpp

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -897,6 +897,8 @@ struct GridwiseGemmBiasAddReduce_k0mk1_k0nk1_mn_xdl_cshuffle_v1
897897
static_assert(num_access == sfc_c_global.GetNumOfAccess(), "wrong!");
898898

899899
static_for<0, num_access, 1>{}([&](auto access_id) {
900+
block_sync_lds();
901+
900902
// each thread write its data from VGPR to LDS
901903
c_thread_copy_vgpr_to_lds.Run(c_thread_desc_m0_n0_m1_n1_m2_m3_m4_n2,
902904
sfc_c_vgpr.GetIndexTupleOfNumber(access_id),

library/include/ck/library/tensor_operation_instance/gpu/device_gemm_mean_squaremean_instance.hpp

Lines changed: 41 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@ namespace instance {
1919

2020
using DeviceGemmAddAddMeanSquareMeanPtr = ck::tensor_operation::device::DeviceGemmReducePtr<1, 2>;
2121

22+
#if defined(CK_USE_XDL)
2223
void add_device_gemm_bias_add_mean_squaremean_xdl_cshuffle_f16_f16_f16_f16_f16_f32_f32_mk_kn_mn_instances(
2324
std::vector<DeviceGemmAddAddMeanSquareMeanPtr>&);
2425
void add_device_gemm_bias_add_mean_squaremean_xdl_cshuffle_f16_f16_f16_f16_f16_f32_f32_mk_nk_mn_instances(
@@ -27,6 +28,18 @@ void add_device_gemm_bias_add_mean_squaremean_xdl_cshuffle_f16_f16_f16_f16_f16_f
2728
std::vector<DeviceGemmAddAddMeanSquareMeanPtr>&);
2829
void add_device_gemm_bias_add_mean_squaremean_xdl_cshuffle_f16_f16_f16_f16_f16_f32_f32_km_nk_mn_instances(
2930
std::vector<DeviceGemmAddAddMeanSquareMeanPtr>&);
31+
#endif // CK_USE_XDL
32+
33+
#if defined(CK_USE_WMMA)
34+
void add_device_gemm_bias_add_mean_squaremean_wmma_cshuffle_f16_f16_f16_f16_f16_f32_f32_mk_kn_mn_instances(
35+
std::vector<DeviceGemmAddAddMeanSquareMeanPtr>&);
36+
void add_device_gemm_bias_add_mean_squaremean_wmma_cshuffle_f16_f16_f16_f16_f16_f32_f32_mk_nk_mn_instances(
37+
std::vector<DeviceGemmAddAddMeanSquareMeanPtr>&);
38+
void add_device_gemm_bias_add_mean_squaremean_wmma_cshuffle_f16_f16_f16_f16_f16_f32_f32_km_kn_mn_instances(
39+
std::vector<DeviceGemmAddAddMeanSquareMeanPtr>&);
40+
void add_device_gemm_bias_add_mean_squaremean_wmma_cshuffle_f16_f16_f16_f16_f16_f32_f32_km_nk_mn_instances(
41+
std::vector<DeviceGemmAddAddMeanSquareMeanPtr>&);
42+
#endif // CK_USE_WMMA
3043

3144
template <typename ADataType,
3245
typename BDataType,
@@ -45,33 +58,61 @@ auto get_device_gemm_add_add_mean_squaremean_instances()
4558
is_same<BLayout, tensor_layout::gemm::RowMajor>::value &&
4659
is_same<CLayout, tensor_layout::gemm::RowMajor>::value)
4760
{
61+
#if defined(CK_USE_XDL)
4862
ck::tensor_operation::device::instance::
4963
add_device_gemm_bias_add_mean_squaremean_xdl_cshuffle_f16_f16_f16_f16_f16_f32_f32_mk_kn_mn_instances(
5064
op_ptrs);
65+
#endif
66+
#if defined(CK_USE_WMMA)
67+
ck::tensor_operation::device::instance::
68+
add_device_gemm_bias_add_mean_squaremean_wmma_cshuffle_f16_f16_f16_f16_f16_f32_f32_mk_kn_mn_instances(
69+
op_ptrs);
70+
#endif
5171
}
5272
else if constexpr(is_same<ALayout, tensor_layout::gemm::RowMajor>::value &&
5373
is_same<BLayout, tensor_layout::gemm::ColumnMajor>::value &&
5474
is_same<CLayout, tensor_layout::gemm::RowMajor>::value)
5575
{
76+
#if defined(CK_USE_XDL)
5677
ck::tensor_operation::device::instance::
5778
add_device_gemm_bias_add_mean_squaremean_xdl_cshuffle_f16_f16_f16_f16_f16_f32_f32_mk_nk_mn_instances(
5879
op_ptrs);
80+
#endif
81+
#if defined(CK_USE_WMMA)
82+
ck::tensor_operation::device::instance::
83+
add_device_gemm_bias_add_mean_squaremean_wmma_cshuffle_f16_f16_f16_f16_f16_f32_f32_mk_nk_mn_instances(
84+
op_ptrs);
85+
#endif
5986
}
6087
else if constexpr(is_same<ALayout, tensor_layout::gemm::ColumnMajor>::value &&
6188
is_same<BLayout, tensor_layout::gemm::RowMajor>::value &&
6289
is_same<CLayout, tensor_layout::gemm::RowMajor>::value)
6390
{
91+
#if defined(CK_USE_XDL)
6492
ck::tensor_operation::device::instance::
6593
add_device_gemm_bias_add_mean_squaremean_xdl_cshuffle_f16_f16_f16_f16_f16_f32_f32_km_kn_mn_instances(
6694
op_ptrs);
95+
#endif
96+
#if defined(CK_USE_WMMA)
97+
ck::tensor_operation::device::instance::
98+
add_device_gemm_bias_add_mean_squaremean_wmma_cshuffle_f16_f16_f16_f16_f16_f32_f32_km_kn_mn_instances(
99+
op_ptrs);
100+
#endif
67101
}
68102
else if constexpr(is_same<ALayout, tensor_layout::gemm::ColumnMajor>::value &&
69103
is_same<BLayout, tensor_layout::gemm::ColumnMajor>::value &&
70104
is_same<CLayout, tensor_layout::gemm::RowMajor>::value)
71105
{
106+
#if defined(CK_USE_XDL)
72107
ck::tensor_operation::device::instance::
73108
add_device_gemm_bias_add_mean_squaremean_xdl_cshuffle_f16_f16_f16_f16_f16_f32_f32_km_nk_mn_instances(
74109
op_ptrs);
110+
#endif
111+
#if defined(CK_USE_WMMA)
112+
ck::tensor_operation::device::instance::
113+
add_device_gemm_bias_add_mean_squaremean_wmma_cshuffle_f16_f16_f16_f16_f16_f32_f32_km_nk_mn_instances(
114+
op_ptrs);
115+
#endif
75116
}
76117
}
77118

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,15 @@
11
# Copyright (c) Advanced Micro Devices, Inc., or its affiliates.
22
# SPDX-License-Identifier: MIT
33

4-
# ONLY XDL_KERNELS
4+
# ONLY XDL_AND_WMMA_KERNELS
55
add_instance_library(device_gemm_bias_add_reduce_instance
66
device_gemm_bias_add_mean_squaremean_xdl_cshuffle_f16_f16_f16_f32_f32_mk_kn_mn_instance.cpp
77
device_gemm_bias_add_mean_squaremean_xdl_cshuffle_f16_f16_f16_f32_f32_mk_nk_mn_instance.cpp
88
device_gemm_bias_add_mean_squaremean_xdl_cshuffle_f16_f16_f16_f32_f32_km_kn_mn_instance.cpp
99
device_gemm_bias_add_mean_squaremean_xdl_cshuffle_f16_f16_f16_f32_f32_km_nk_mn_instance.cpp
10+
11+
device_gemm_bias_add_mean_squaremean_wmma_cshuffle_f16_f16_f16_f32_f32_mk_kn_mn_instance.cpp
12+
device_gemm_bias_add_mean_squaremean_wmma_cshuffle_f16_f16_f16_f32_f32_mk_nk_mn_instance.cpp
13+
device_gemm_bias_add_mean_squaremean_wmma_cshuffle_f16_f16_f16_f32_f32_km_kn_mn_instance.cpp
14+
device_gemm_bias_add_mean_squaremean_wmma_cshuffle_f16_f16_f16_f32_f32_km_nk_mn_instance.cpp
1015
)

0 commit comments

Comments
 (0)