@@ -10,6 +10,7 @@ namespace ck {
1010
1111template <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
0 commit comments