Skip to content

Commit bc497be

Browse files
authored
[CK TILE] Fix grouped conv kernels splitk and double lds (#3527)
1 parent f449a5f commit bc497be

3 files changed

Lines changed: 54 additions & 330 deletions

File tree

include/ck_tile/ops/grouped_convolution/kernel/grouped_convolution_backward_data_kernel.hpp

Lines changed: 21 additions & 117 deletions
Original file line numberDiff line numberDiff line change
@@ -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

include/ck_tile/ops/grouped_convolution/kernel/grouped_convolution_backward_weight_kernel.hpp

Lines changed: 9 additions & 96 deletions
Original file line numberDiff line numberDiff line change
@@ -829,66 +829,14 @@ struct GroupedConvolutionBackwardWeightKernel
829829
}
830830
else
831831
{
832-
auto c_block_window = MakeCBlockWindow<memory_operation_enum::atomic_add>(
833-
c_ptr, kargs, block_idx_m, block_idx_n);
834-
835-
EpiloguePipeline{}(c_block_window, c_block_tile, d_block_window, smem_ptr_0);
836-
}
837-
}
838-
839-
/**
840-
* @brief Runs single GEMM problem cooperatively by whole workgroup.
841-
*
842-
* @note RunGEMM2LDS in with two shared memory buffers using the ping pong buffer mechanism.
843-
*
844-
* @param a_ptr input A pointer
845-
* @param b_ptr input B pointer
846-
* @param c_ptr output C pointer
847-
* @param smem_ptr_0 The starting pointer of 1st shared memory block.
848-
* @param smem_ptr_1 The starting pointer of 2nd shared memory block.
849-
* @param kargs Grouped Convolution Backward Weight kernel arguments
850-
* @param block_idx_m The GEMM's output M dimension tile index processed by this workgroup.
851-
* @param block_idx_n The GEMM's output N dimension tile index processed by this workgroup.
852-
*
853-
*/
854-
CK_TILE_DEVICE static void RunGemm2LDS(const OutDataType* a_ptr,
855-
const InDataType* b_ptr,
856-
const std::array<const void*, NumDTensor>& ds_ptr,
857-
WeiDataType* c_ptr,
858-
void* __restrict__ smem_ptr_0,
859-
void* __restrict__ smem_ptr_1,
860-
const GroupedConvBwdWeightKernelArgsSpecialized& kargs,
861-
const index_t num_loop,
862-
const index_t block_idx_m,
863-
const index_t block_idx_n,
864-
const index_t block_idx_k)
865-
{
866-
// Create block windows using helper methods
867-
const auto& a_block_window = MakeABlockWindow(a_ptr, kargs, block_idx_m, block_idx_k);
868-
const auto& b_block_window = MakeBBlockWindow(b_ptr, kargs, block_idx_n, block_idx_k);
869-
const auto& d_block_window = MakeDBlockWindows(ds_ptr, kargs, block_idx_m, block_idx_n);
870-
871-
// Run GEMM cooperatively by whole workgroup.
872-
const auto& c_block_tile = GemmPipeline{}.template operator()(
873-
a_block_window, b_block_window, num_loop, smem_ptr_0, smem_ptr_1);
874-
875-
// Run Epilogue Pipeline with k_batch dispatching
876-
if(kargs.k_batch == 1)
877-
{
878-
auto c_block_window = MakeCBlockWindow<memory_operation_enum::set>(
879-
c_ptr, kargs, block_idx_m, block_idx_n);
880-
881-
EpiloguePipeline{}(c_block_window, c_block_tile, d_block_window, smem_ptr_0);
882-
}
883-
else
884-
{
885-
#if defined(__gfx11__)
886-
return;
887-
#endif
888-
auto c_block_window = MakeCBlockWindow<memory_operation_enum::atomic_add>(
889-
c_ptr, kargs, block_idx_m, block_idx_n);
832+
if constexpr(!(GroupedConvTraitsType_::VectorSizeC % 2 != 0 &&
833+
is_any_of<WeiDataType, fp16_t, bf16_t>::value))
834+
{
835+
auto c_block_window = MakeCBlockWindow<memory_operation_enum::atomic_add>(
836+
c_ptr, kargs, block_idx_m, block_idx_n);
890837

891-
EpiloguePipeline{}(c_block_window, c_block_tile, d_block_window, smem_ptr_0);
838+
EpiloguePipeline{}(c_block_window, c_block_tile, d_block_window, smem_ptr_0);
839+
}
892840
}
893841
}
894842

@@ -949,44 +897,9 @@ struct GroupedConvolutionBackwardWeightKernel
949897
const InDataType* b_ptr = static_cast<const InDataType*>(kargs.in_ptr) + group_offset_b;
950898
WeiDataType* c_ptr = static_cast<WeiDataType*>(kargs.wei_ptr) + group_offset_c;
951899

952-
__shared__ char smem_ptr_0[GetSmemSize()];
900+
__shared__ char smem_ptr[GetSmemSize()];
953901

954-
if constexpr(GemmPipeline::DoubleSmemBuffer == true)
955-
{
956-
__shared__ char smem_ptr_1[GemmPipeline::GetSmemSize()];
957-
if constexpr(!(GroupedConvTraitsType_::VectorSizeC % 2 != 0 &&
958-
is_any_of<WeiDataType, fp16_t, bf16_t>::value))
959-
{
960-
RunGemm2LDS(a_ptr,
961-
b_ptr,
962-
kargs.ds_ptr,
963-
c_ptr,
964-
smem_ptr_0,
965-
smem_ptr_1,
966-
kargs,
967-
num_loop,
968-
i_m,
969-
i_n,
970-
i_k);
971-
}
972-
}
973-
else
974-
{
975-
if constexpr(!(GroupedConvTraitsType_::VectorSizeC % 2 != 0 &&
976-
is_any_of<WeiDataType, fp16_t, bf16_t>::value))
977-
{
978-
RunGemm(a_ptr,
979-
b_ptr,
980-
kargs.ds_ptr,
981-
c_ptr,
982-
smem_ptr_0,
983-
kargs,
984-
num_loop,
985-
i_m,
986-
i_n,
987-
i_k);
988-
}
989-
}
902+
RunGemm(a_ptr, b_ptr, kargs.ds_ptr, c_ptr, smem_ptr, kargs, num_loop, i_m, i_n, i_k);
990903
}
991904
}
992905
};

0 commit comments

Comments
 (0)