@@ -1036,84 +1036,16 @@ struct GroupedConvolutionBackwardDataKernel
10361036 }
10371037 else
10381038 {
1039- auto c_block_window = MakeCBlockWindow<memory_operation_enum::atomic_add>(
1040- c_ptr, kargs, group_id, block_idx_m, block_idx_n);
1041-
1042- EpiloguePipeline{}
1043- .template operator ()<decltype (c_block_window), decltype (c_block_tile)>(
1044- c_block_window, c_block_tile, d_block_window, smem_ptr_0);
1045- }
1046- }
1047-
1048- /* *
1049- * @brief Runs single GEMM problem cooperatively by whole workgroup.
1050- *
1051- * @note RunGemm2LDS in with two shared memory buffers using the ping pong buffer mechanism.
1052- *
1053- * @param a_ptr input A pointer
1054- * @param b_ptr input B pointer
1055- * @param c_ptr output C pointer
1056- * @param smem_ptr_0 The starting pointer of 1st shared memory block.
1057- * @param smem_ptr_1 The starting pointer of 2nd shared memory block.
1058- * @param kargs Grouped Convolution Backward Data kernel arguments
1059- * @param block_idx_m The GEMM's output M dimension tile index processed by this workgroup.
1060- * @param block_idx_n The GEMM's output N dimension tile index processed by this workgroup.
1061- *
1062- */
1063- CK_TILE_DEVICE static void RunGemm2LDS (const OutDataType* a_ptr,
1064- const InDataType* b_ptr,
1065- const std::array<const void *, NumDTensor>& ds_ptr,
1066- WeiDataType* c_ptr,
1067- void * __restrict__ smem_ptr_0,
1068- void * __restrict__ smem_ptr_1,
1069- const GroupedConvBwdDataKernelArgsSpecialized& kargs,
1070- const index_t splitted_k,
1071- const index_t block_idx_m,
1072- const index_t block_idx_n,
1073- const index_t block_idx_k,
1074- const index_t group_id)
1075- {
1076- // Create block windows using specialized methods
1077- const auto & a_block_window =
1078- MakeABlockWindow (a_ptr, kargs, group_id, block_idx_m, block_idx_k);
1079- const auto & b_block_window =
1080- MakeBBlockWindow (b_ptr, kargs, group_id, block_idx_n, block_idx_k);
1081- const auto & d_block_window =
1082- MakeDBlockWindows (ds_ptr, kargs, group_id, block_idx_m, block_idx_n);
1083-
1084- const index_t num_loop = amd_wave_read_first_lane (TilePartitioner::GetLoopNum (splitted_k));
1085- const bool has_hot_loop = GemmPipeline::BlockHasHotloop (num_loop);
1086- const TailNumber tail_num = GemmPipeline::GetBlockLoopTailNum (num_loop);
1087-
1088- // Run GEMM cooperatively by whole workgroup.
1089- const auto & c_block_tile = GemmPipeline{}.template operator ()(a_block_window,
1090- b_block_window,
1091- num_loop,
1092- has_hot_loop,
1093- tail_num,
1094- smem_ptr_0,
1095- smem_ptr_1);
1096-
1097- const index_t k_batch = amd_wave_read_first_lane (kargs.k_batch );
1098-
1099- // Run Epilogue Pipeline with k_batch dispatch
1100- if (k_batch == 1 )
1101- {
1102- auto c_block_window = MakeCBlockWindow<memory_operation_enum::set>(
1103- c_ptr, kargs, group_id, block_idx_m, block_idx_n);
1104-
1105- EpiloguePipeline{}
1106- .template operator ()<decltype (c_block_window), decltype (c_block_tile)>(
1107- c_block_window, c_block_tile, d_block_window, smem_ptr_0);
1108- }
1109- else
1110- {
1111- auto c_block_window = MakeCBlockWindow<memory_operation_enum::atomic_add>(
1112- c_ptr, kargs, group_id, block_idx_m, block_idx_n);
1039+ if constexpr (!(GroupedConvTraitsType_::VectorSizeC % 2 != 0 &&
1040+ is_any_of<OutDataType, fp16_t , bf16_t >::value))
1041+ {
1042+ auto c_block_window = MakeCBlockWindow<memory_operation_enum::atomic_add>(
1043+ c_ptr, kargs, group_id, block_idx_m, block_idx_n);
11131044
1114- EpiloguePipeline{}
1115- .template operator ()<decltype (c_block_window), decltype (c_block_tile)>(
1116- c_block_window, c_block_tile, d_block_window, smem_ptr_0);
1045+ EpiloguePipeline{}
1046+ .template operator ()<decltype (c_block_window), decltype (c_block_tile)>(
1047+ c_block_window, c_block_tile, d_block_window, smem_ptr_0);
1048+ }
11171049 }
11181050 }
11191051
@@ -1195,46 +1127,18 @@ struct GroupedConvolutionBackwardDataKernel
11951127 static_cast <InDataType*>(kargs.in_ptr ) + group_offset_c + input_batch_offset;
11961128
11971129 // allocate LDS
1198- __shared__ char smem_ptr_0[GetSmemSize ()];
1199-
1200- if constexpr (GemmPipeline::DoubleSmemBuffer == true )
1201- {
1202- __shared__ char smem_ptr_1[GemmPipeline::GetSmemSize ()];
1203- if constexpr (!(GroupedConvTraitsType_::VectorSizeC % 2 != 0 &&
1204- is_any_of<OutDataType, fp16_t , bf16_t >::value))
1205- {
1206- RunGemm2LDS (a_ptr,
1207- b_ptr,
1208- kargs.ds_ptr ,
1209- c_ptr,
1210- smem_ptr_0,
1211- smem_ptr_1,
1212- kargs,
1213- splitted_k,
1214- i_m,
1215- i_n,
1216- i_k,
1217- group_id);
1218- }
1219- }
1220- else
1221- {
1222- if constexpr (!(GroupedConvTraitsType_::VectorSizeC % 2 != 0 &&
1223- is_any_of<OutDataType, fp16_t , bf16_t >::value))
1224- {
1225- RunGemm (a_ptr,
1226- b_ptr,
1227- kargs.ds_ptr ,
1228- c_ptr,
1229- smem_ptr_0,
1230- kargs,
1231- splitted_k,
1232- i_m,
1233- i_n,
1234- i_k,
1235- group_id);
1236- }
1237- }
1130+ __shared__ char smem_ptr[GetSmemSize ()];
1131+ RunGemm (a_ptr,
1132+ b_ptr,
1133+ kargs.ds_ptr ,
1134+ c_ptr,
1135+ smem_ptr,
1136+ kargs,
1137+ splitted_k,
1138+ i_m,
1139+ i_n,
1140+ i_k,
1141+ group_id);
12381142 }
12391143};
12401144
0 commit comments