Skip to content
Merged
Show file tree
Hide file tree
Changes from 4 commits
Commits
Show all changes
56 commits
Select commit Hold shift + click to select a range
0c2d00d
Reference CGEMM + test stub
myamlak May 11, 2022
4d07aa1
Format.
myamlak May 12, 2022
b1c9458
Incomplete simple implementation
myamlak May 13, 2022
6a0883f
Library instances
myamlak May 13, 2022
14bd143
Sketch of tests
myamlak May 13, 2022
674f74a
Test fixes.
myamlak May 16, 2022
55927aa
Example added
myamlak May 16, 2022
ffe12e2
Cosmetics
myamlak May 16, 2022
a61f34f
Add elementwise operation kernel and example
May 16, 2022
c262612
Add comment
May 17, 2022
e00a943
Merge remote-tracking branch 'origin/develop' into myamlak/cgemm
myamlak May 17, 2022
b3767db
Merge remote-tracking branch 'origin/eltwise_op' into myamlak/cgemm
myamlak May 17, 2022
b456d5e
Add template argument of dim . Prepare to support multiple dimension
May 17, 2022
0d26477
Rename example
May 17, 2022
4af77e1
Support 1 dimension
May 17, 2022
492da45
Add static assert
May 17, 2022
ecdfe96
Add comment
May 17, 2022
5ae304d
Second auxiliary buffer added
myamlak May 17, 2022
0f84025
Extract pad
May 17, 2022
06e52d9
Remove redundant argument
May 17, 2022
7d44e78
Support any dimension for elementwise operation
May 17, 2022
b7a82d2
Remove line
May 17, 2022
83f7531
Let it be the multiple number of CU
May 18, 2022
c4d610b
Move thread per block to the parameter of constructor
May 18, 2022
5e10474
Merge remote-tracking branch 'origin/eltwise_op' into myamlak/cgemm
myamlak May 18, 2022
208ac1a
Consuming binary ops to do A+B / A-B
myamlak May 18, 2022
6ebcb66
Fix + cosmetics + bf16 test commented out temporarily
myamlak May 18, 2022
a7676df
Merge remote-tracking branch 'origin/develop' into myamlak/cgemm
myamlak May 19, 2022
f63ca8e
Format
myamlak May 19, 2022
f497e2b
Enabling bf16 test
myamlak May 19, 2022
18125c3
Revert "Enabling bf16 test"
myamlak May 20, 2022
5fd5daa
Fix + test reenabled
myamlak May 20, 2022
d731023
fix build
May 21, 2022
d00ecab
Revert "fix build"
rosenrodt May 23, 2022
3f8e846
post PR #235 merge fix
rosenrodt May 23, 2022
1cfff19
amend
rosenrodt May 23, 2022
326d331
Merge remote-tracking branch 'origin/develop' into myamlak/cgemm
myamlak May 23, 2022
4379d8d
Merge remote-tracking branch 'origin/fix_build_0521' into myamlak/cgemm
myamlak May 23, 2022
f73c3ea
Single workspace for cgemm + helper
myamlak May 23, 2022
d3ec209
Perf calc fix
myamlak May 23, 2022
b7d9cde
Merge remote-tracking branch 'origin/develop' into myamlak/cgemm
myamlak May 24, 2022
ac9ef30
Review remarks: static_cast
myamlak May 24, 2022
d478a38
Review remarks: binary ops templated
myamlak May 24, 2022
c82093c
Cleaning
myamlak May 24, 2022
e6914f2
Removal of instances and their tests
myamlak May 24, 2022
ee06099
Review remarks from aosew addressed
myamlak May 24, 2022
97ac500
Review remark: unnecessary attribute
myamlak May 25, 2022
bb1f808
Merge remote-tracking branch 'origin/develop' into myamlak/cgemm
myamlak May 26, 2022
80f038a
Post-merge fixes
myamlak May 26, 2022
bda2654
Merge remote-tracking branch 'origin/develop' into myamlak/cgemm
myamlak May 27, 2022
7d2fa99
Restrict 4gemm to PassThrough + bug fix
myamlak May 27, 2022
d35c6a8
Review remarks
myamlak May 30, 2022
5b250a9
Merge branch 'develop' into myamlak/cgemm
myamlak May 30, 2022
fe01b4d
Merge remote-tracking branch 'origin/develop' into myamlak/cgemm
May 30, 2022
52ccb9b
update licence
May 30, 2022
9fdbcf9
change cgemm example to fp16
May 30, 2022
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 5 additions & 4 deletions example/19_binary_elementwise/broadcast_add_3d_am_bmnk.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,8 @@ using ABDataType = F16;
using CDataType = F16;
using EltwiseComputeDataType = F32;

using Add = ck::tensor_operation::binary_element_wise::Add;
using Add = ck::tensor_operation::binary_element_wise::
Add<EltwiseComputeDataType, EltwiseComputeDataType, EltwiseComputeDataType>;

using DeviceElementwiseAddInstance =
ck::tensor_operation::device::DeviceBinaryElementwise<ABDataType,
Expand Down Expand Up @@ -48,11 +49,11 @@ void host_broadcast3D_am_bmnk(HostTensorC& C,
for(std::size_t n = 0; n < shape[1]; ++n)
for(std::size_t k = 0; k < shape[2]; ++k)
{
ComputeDataType a_val = static_cast<ComputeDataType>(A(m));
ComputeDataType b_val = static_cast<ComputeDataType>(B(m, n, k));
ComputeDataType a_val = ck::type_convert<ComputeDataType>(A(m));
ComputeDataType b_val = ck::type_convert<ComputeDataType>(B(m, n, k));
ComputeDataType c_val = 0;
functor(c_val, a_val, b_val);
C(m, n, k) = static_cast<ctype>(c_val);
C(m, n, k) = ck::type_convert<ctype>(c_val);
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -42,48 +42,54 @@ namespace ck {
namespace tensor_operation {
namespace device {

template <typename ALayout,
typename BLayout,
typename CLayout,
typename ADataType,
typename BDataType,
typename CDataType,
typename GemmAccDataType,
typename CShuffleDataType,
typename AElementwiseOperation,
typename BElementwiseOperation,
typename CElementwiseOperation,
GemmSpecialization GemmSpec,
index_t NumGemmKPrefetchStage,
index_t BlockSize,
index_t MPerBlock,
index_t NPerBlock,
index_t KPerBlock,
index_t AK1,
index_t BK1,
index_t MPerXDL,
index_t NPerXDL,
index_t MXdlPerWave,
index_t NXdlPerWave,
typename ABlockTransferThreadClusterLengths_AK0_M_AK1,
typename ABlockTransferThreadClusterArrangeOrder,
typename ABlockTransferSrcAccessOrder,
index_t ABlockTransferSrcVectorDim,
index_t ABlockTransferSrcScalarPerVector,
index_t ABlockTransferDstScalarPerVector_AK1,
bool ABlockLdsExtraM,
typename BBlockTransferThreadClusterLengths_BK0_N_BK1,
typename BBlockTransferThreadClusterArrangeOrder,
typename BBlockTransferSrcAccessOrder,
index_t BBlockTransferSrcVectorDim,
index_t BBlockTransferSrcScalarPerVector,
index_t BBlockTransferDstScalarPerVector_BK1,
bool BBlockLdsExtraN,
index_t CShuffleMXdlPerWavePerShuffle,
index_t CShuffleNXdlPerWavePerShuffle,
typename CShuffleBlockTransferClusterLengths_MBlock_MPerBlock_NBlock_NPerBlock,
index_t CShuffleBlockTransferScalarPerVector_NPerBlock,
LoopScheduler LoopSched = make_default_loop_scheduler()>
template <
typename ALayout,
typename BLayout,
typename CLayout,
typename ADataType,
typename BDataType,
typename CDataType,
typename GemmAccDataType,
typename CShuffleDataType,
typename AElementwiseOperation,
typename BElementwiseOperation,
typename CElementwiseOperation,
GemmSpecialization GemmSpec,
index_t NumGemmKPrefetchStage,
index_t BlockSize,
index_t MPerBlock,
index_t NPerBlock,
index_t KPerBlock,
index_t AK1,
index_t BK1,
index_t MPerXDL,
index_t NPerXDL,
index_t MXdlPerWave,
index_t NXdlPerWave,
typename ABlockTransferThreadClusterLengths_AK0_M_AK1,
typename ABlockTransferThreadClusterArrangeOrder,
typename ABlockTransferSrcAccessOrder,
index_t ABlockTransferSrcVectorDim,
index_t ABlockTransferSrcScalarPerVector,
index_t ABlockTransferDstScalarPerVector_AK1,
bool ABlockLdsExtraM,
typename BBlockTransferThreadClusterLengths_BK0_N_BK1,
typename BBlockTransferThreadClusterArrangeOrder,
typename BBlockTransferSrcAccessOrder,
index_t BBlockTransferSrcVectorDim,
index_t BBlockTransferSrcScalarPerVector,
index_t BBlockTransferDstScalarPerVector_BK1,
bool BBlockLdsExtraN,
index_t CShuffleMXdlPerWavePerShuffle,
index_t CShuffleNXdlPerWavePerShuffle,
typename CShuffleBlockTransferClusterLengths_MBlock_MPerBlock_NBlock_NPerBlock,
index_t CShuffleBlockTransferScalarPerVector_NPerBlock,
LoopScheduler LoopSched = make_default_loop_scheduler(),
enable_if_t<
is_same_v<AElementwiseOperation, ck::tensor_operation::element_wise::PassThrough> &&
is_same_v<BElementwiseOperation, ck::tensor_operation::element_wise::PassThrough> &&
is_same_v<CElementwiseOperation, ck::tensor_operation::element_wise::PassThrough>,
bool> = false>
struct DeviceCGemm_4Gemm_Xdl_CShuffle
: public DeviceCGemm<AElementwiseOperation, BElementwiseOperation, CElementwiseOperation>
{
Expand All @@ -93,39 +99,37 @@ struct DeviceCGemm_4Gemm_Xdl_CShuffle
static constexpr auto I1 = Number<1>{};
static constexpr auto I2 = Number<2>{};

static constexpr auto ScalarPerVector = Number<4>{};
static constexpr auto MPerThread = Number<4>{};
static constexpr auto AScalarPerVector = Number<4>{};
static constexpr auto BScalarPerVector = Number<4>{};
static constexpr auto CScalarPerVector = Number<4>{};

template <typename Desc_M0>
static auto PadDescriptor_M0_1d(Desc_M0 desc_m0, index_t gridSize, index_t blockSize)
template <typename Desc_M>
static auto PadDescriptor_M_1d(Desc_M desc_m, index_t gridSize, index_t blockSize)
{
const auto m0 = desc_m0.GetLength(I0);
const index_t loop_step = gridSize * blockSize * ScalarPerVector;
const auto pad = math::integer_least_multiple(m0, loop_step) - m0;
const auto desc_m0_pad =
transform_tensor_descriptor(desc_m0,
make_tuple(make_right_pad_transform(m0, pad)),
const auto M = desc_m.GetLength(I0);
const index_t loop_step = gridSize * blockSize * MPerThread;
const auto pad = math::integer_least_multiple(M, loop_step) - M;
const auto desc_m_pad =
transform_tensor_descriptor(desc_m,
make_tuple(make_right_pad_transform(M, pad)),
make_tuple(Sequence<0>{}),
make_tuple(Sequence<0>{}));
return desc_m0_pad;
return desc_m_pad;
}

static auto MakeDescriptor_M0(const std::vector<int>& shape,
const std::vector<int>& stride,
index_t gridSize,
index_t blockSize)
static auto MakeDescriptor_M(const std::vector<index_t>& lengths,
const std::vector<index_t>& strides,
index_t gridSize,
index_t blockSize)
{
auto tupleOfShape = generate_tuple([&](auto I) { return shape[I]; }, Number<2>{});
auto tupleOfStride = generate_tuple([&](auto I) { return stride[I]; }, Number<2>{});
auto tupleOfShape = generate_tuple([&](auto I) { return lengths[I]; }, Number<1>{});
auto tupleOfStride = generate_tuple([&](auto I) { return strides[I]; }, Number<1>{});
Comment thread
aosewski marked this conversation as resolved.
Outdated

// nd desc - [s0, s1, s2, ...]
const auto desc = make_naive_tensor_descriptor(tupleOfShape, tupleOfStride);

const auto desc_m0 = transform_tensor_descriptor(
desc,
make_tuple(make_merge_transform(tupleOfShape)),
make_tuple(generate_sequence_v2([&](auto I) { return I; }, Number<2>{})),
make_tuple(Sequence<0>{}));

return PadDescriptor_M0_1d(desc_m0, gridSize, blockSize);
return PadDescriptor_M_1d(desc, gridSize, blockSize);
}

static auto MakeAGridDescriptor_AK0_M_AK1(index_t MRaw, index_t KRaw, index_t StrideA)
Expand Down Expand Up @@ -395,7 +399,7 @@ struct DeviceCGemm_4Gemm_Xdl_CShuffle
using AGridDesc_AK0_M_AK1 = decltype(MakeAGridDescriptor_AK0_M_AK1(1, 1, 1));
using BGridDesc_BK0_N_BK1 = decltype(MakeBGridDescriptor_BK0_N_BK1(1, 1, 1));
using CGridDesc_M_N = decltype(MakeCGridDescriptor_M_N(1, 1, 1));
using GridDesc_M0 = decltype(MakeDescriptor_M0({1, 1}, {1, 1}, 1, 1));
using CGridDesc_M = decltype(MakeDescriptor_M({1, 1}, {1, 1}, 1, 1));

// GridwiseGemm
using GridwiseGemm = GridwiseGemm_k0mk1_k0nk1_mn_xdl_cshuffle_v1<
Expand Down Expand Up @@ -492,13 +496,13 @@ struct DeviceCGemm_4Gemm_Xdl_CShuffle

if constexpr(is_same<tensor_layout::gemm::RowMajor, CLayout>::value)
{
c_grid_desc_m0_ =
DeviceOp::MakeDescriptor_M0({MRaw, NRaw}, {StrideC, I1}, grid_size, BlockSize);
c_grid_desc_m_ =
DeviceOp::MakeDescriptor_M({MRaw, NRaw}, {StrideC, I1}, grid_size, BlockSize);
}
else if constexpr(is_same<tensor_layout::gemm::ColumnMajor, CLayout>::value)
{
c_grid_desc_m0_ =
DeviceOp::MakeDescriptor_M0({MRaw, NRaw}, {I1, StrideC}, grid_size, BlockSize);
c_grid_desc_m_ =
DeviceOp::MakeDescriptor_M({MRaw, NRaw}, {I1, StrideC}, grid_size, BlockSize);
}

p_aux_2_grid_ = p_workspace + c_grid_desc_m_n_.GetElementSpaceSize();
Expand All @@ -516,7 +520,7 @@ struct DeviceCGemm_4Gemm_Xdl_CShuffle
AGridDesc_AK0_M_AK1 a_grid_desc_ak0_m_ak1_;
BGridDesc_BK0_N_BK1 b_grid_desc_bk0_n_bk1_;
CGridDesc_M_N c_grid_desc_m_n_;
GridDesc_M0 c_grid_desc_m0_;
CGridDesc_M c_grid_desc_m_;
typename GridwiseGemm::CGridDescriptor_MBlock_MPerBlock_NBlock_NPerBlock
c_grid_desc_mblock_mperblock_nblock_nperblock_;
typename GridwiseGemm::DefaultBlock2CTileMap block_2_ctile_map_;
Expand Down Expand Up @@ -556,27 +560,41 @@ struct DeviceCGemm_4Gemm_Xdl_CShuffle
CDataType,
CDataType,
CDataType,
GridDesc_M0,
CGridDesc_M,
CGridDesc_M,
CGridDesc_M,
Add,
ScalarPerVector>;
MPerThread,
AScalarPerVector,
BScalarPerVector,
CScalarPerVector>;
using GridwiseBinSubstract = GridwiseBinaryElementwise_1D<CDataType,
CDataType,
CDataType,
CDataType,
GridDesc_M0,
CGridDesc_M,
CGridDesc_M,
CGridDesc_M,
Substract,
ScalarPerVector>;
MPerThread,
AScalarPerVector,
BScalarPerVector,
CScalarPerVector>;
const auto add_kernel = kernel_binary_elementwise_1d<GridwiseBinAdd,
CDataType,
CDataType,
CDataType,
GridDesc_M0,
CGridDesc_M,
CGridDesc_M,
CGridDesc_M,
Add>;
const auto substract_kernel = kernel_binary_elementwise_1d<GridwiseBinSubstract,
CDataType,
CDataType,
CDataType,
GridDesc_M0,
CGridDesc_M,
CGridDesc_M,
CGridDesc_M,
Substract>;

if(GridwiseGemm::CalculateHasMainKBlockLoop(K))
Expand Down Expand Up @@ -637,9 +655,9 @@ struct DeviceCGemm_4Gemm_Xdl_CShuffle
arg.p_aux_grid_,
arg.p_aux_2_grid_,
arg.p_c_grid_real_,
arg.c_grid_desc_m0_,
arg.c_grid_desc_m0_,
arg.c_grid_desc_m0_,
arg.c_grid_desc_m_,
arg.c_grid_desc_m_,
arg.c_grid_desc_m_,
Substract{});

ave_time +=
Expand Down Expand Up @@ -685,9 +703,9 @@ struct DeviceCGemm_4Gemm_Xdl_CShuffle
arg.p_aux_grid_,
arg.p_aux_2_grid_,
arg.p_c_grid_imag_,
arg.c_grid_desc_m0_,
arg.c_grid_desc_m0_,
arg.c_grid_desc_m0_,
arg.c_grid_desc_m_,
arg.c_grid_desc_m_,
arg.c_grid_desc_m_,
Add{});
}
else
Expand Down Expand Up @@ -748,9 +766,9 @@ struct DeviceCGemm_4Gemm_Xdl_CShuffle
arg.p_aux_grid_,
arg.p_aux_2_grid_,
arg.p_c_grid_real_,
arg.c_grid_desc_m0_,
arg.c_grid_desc_m0_,
arg.c_grid_desc_m0_,
arg.c_grid_desc_m_,
arg.c_grid_desc_m_,
arg.c_grid_desc_m_,
Substract{});

ave_time +=
Expand Down Expand Up @@ -796,9 +814,9 @@ struct DeviceCGemm_4Gemm_Xdl_CShuffle
arg.p_aux_grid_,
arg.p_aux_2_grid_,
arg.p_c_grid_imag_,
arg.c_grid_desc_m0_,
arg.c_grid_desc_m0_,
arg.c_grid_desc_m0_,
arg.c_grid_desc_m_,
arg.c_grid_desc_m_,
arg.c_grid_desc_m_,
Add{});
}

Expand Down
4 changes: 2 additions & 2 deletions include/ck/tensor_operation/gpu/device/device_gemm_dl.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -60,8 +60,8 @@ template <
index_t CThreadTransferDstScalarPerVector,
enable_if_t<
is_same_v<AElementwiseOperation, ck::tensor_operation::element_wise::PassThrough> &&
is_same_v<AElementwiseOperation, ck::tensor_operation::element_wise::PassThrough> &&
is_same_v<AElementwiseOperation, ck::tensor_operation::element_wise::PassThrough>,
is_same_v<BElementwiseOperation, ck::tensor_operation::element_wise::PassThrough> &&
is_same_v<CElementwiseOperation, ck::tensor_operation::element_wise::PassThrough>,
bool> = false>
struct DeviceGemmDl
: public DeviceGemm<AElementwiseOperation, BElementwiseOperation, CElementwiseOperation>
Expand Down