21#ifndef EIGEN_TENSOR_TENSOR_CONTRACTION_SYCL_H
22#define EIGEN_TENSOR_TENSOR_CONTRACTION_SYCL_H
25#include "./InternalHeaderCheck.h"
32#ifndef EIGEN_SYCL_DISABLE_GEMV
47template <
typename Scalar,
typename StorageIndex, StorageIndex NCWindow, StorageIndex CFactor, StorageIndex NCFactor>
50 static constexpr StorageIndex LocalThreadSizeC = EIGEN_SYCL_LOCAL_THREAD_DIM0;
52 static constxpr StorageIndex LocalThreadSizeNC = EIGEN_SYCL_LOCAL_THREAD_DIM1;
54 static constexpr StorageIndex TileSizeDimNC = NCWindow / NCFactor;
56 static constexpr StorageIndex TileSizeDimC = CFactor * LocalThreadSizeNC * LocalThreadSizeC;
58 static constexpr StorageIndex WorkLoadPerThreadNC = TileSizeDimNC / LocalThreadSizeNC;
60 static constexpr StorageIndex WorkLoadPerThreadC = TileSizeDimC / LocalThreadSizeC;
62 static constexpr bool BC =
false;
83template <
typename Scalar,
typename StorageIndex, StorageIndex REG_SIZE_M, StorageIndex REG_SIZE_N, StorageIndex TSDK>
86 static constexpr StorageIndex TileSizeDimK = TSDK;
89#ifndef EIGEN_SYCL_REG_M
90 static constexpr StorageIndex WorkLoadPerThreadM = REG_SIZE_M;
92 static constexpr StorageIndex WorkLoadPerThreadM = EIGEN_SYCL_REG_M;
96#ifndef EIGEN_SYCL_REG_N
97 static constexpr StorageIndex WorkLoadPerThreadN = REG_SIZE_N;
99 static constexpr StorageIndex WorkLoadPerThreadN = EIGEN_SYCL_REG_N;
102 static constexpr StorageIndex LocalThreadSizeM = EIGEN_SYCL_LOCAL_THREAD_DIM0;
104 static constexpr StorageIndex LocalThreadSizeN = EIGEN_SYCL_LOCAL_THREAD_DIM1;
106 static constexpr StorageIndex TileSizeDimM = LocalThreadSizeM * WorkLoadPerThreadM;
108 static constexpr StorageIndex TileSizeDimN = LocalThreadSizeN * WorkLoadPerThreadN;
110 static constexpr StorageIndex LoadPerThreadLhs =
111 ((TileSizeDimK * WorkLoadPerThreadM * WorkLoadPerThreadN) / (TileSizeDimN));
113 static constexpr StorageIndex LoadPerThreadRhs =
114 ((TileSizeDimK * WorkLoadPerThreadM * WorkLoadPerThreadN) / (TileSizeDimM));
116 static constexpr bool BC =
true;
119 static constexpr bool DoubleBuffer =
120#ifdef EIGEN_SYCL_DISABLE_DOUBLE_BUFFER
131enum class contraction_type { local, no_local };
135enum class data_source { global_mem, local_mem, private_mem };
162template <
bool PacketLoad,
bool is_coalesced_layout, bool,
typename PacketType,
typename TensorMapper,
163 typename StorageIndex>
164static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<PacketLoad, PacketType> read(
165 const TensorMapper &tensorMapper,
const StorageIndex &NCIndex,
const StorageIndex &CIndex,
const StorageIndex &ld) {
166 const StorageIndex row = (is_coalesced_layout) ? NCIndex : CIndex;
167 const StorageIndex col = (is_coalesced_layout) ? CIndex : NCIndex;
168 return tensorMapper.get_tensor().template packet<Unaligned>(row + (col * ld));
193template <
bool PacketLoad,
bool,
bool IsRhs,
typename PacketType,
typename TensorMapper,
typename StorageIndex>
194static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<!PacketLoad, PacketType> read(
195 const TensorMapper &tensorMapper,
const StorageIndex &NCIndex,
const StorageIndex &CIndex,
const StorageIndex &) {
196 const StorageIndex row = (IsRhs) ? CIndex : NCIndex;
197 const StorageIndex col = (IsRhs) ? NCIndex : CIndex;
198 return tensorMapper(row, col);
222template <
typename StorageIndex, StorageIndex ld, data_source dt,
typename PacketType,
typename DataScalar>
223static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<dt != data_source::global_mem, void> write(
224 PacketType &packet_data, DataScalar ptr) {
225 constexpr int PacketSize = Eigen::internal::unpacket_traits<PacketType>::size;
227 for (
int i = 0; i < PacketSize; i++) {
228 *ptr = PacketWrapper<PacketType, PacketSize>::scalarize(i, packet_data);
248template <data_source dt,
typename PacketType,
typename DataScalar>
249static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
250 std::enable_if_t<Eigen::internal::unpacket_traits<PacketType>::size != 1 && dt == data_source::global_mem,
void>
251 write(PacketType &packet_data, DataScalar *ptr) {
252 ::Eigen::internal::pstoreu<DataScalar, PacketType>(ptr, packet_data);
268template <data_source dt,
typename PacketType,
typename DataScalar>
269static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
270 std::enable_if_t<Eigen::internal::unpacket_traits<PacketType>::size == 1 && dt == data_source::global_mem,
void>
271 write(PacketType &packet_data, DataScalar *ptr) {
280template <
bool is_
internal>
281EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
bool check_boundary(
bool) {
291EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
bool check_boundary<false>(
bool cond) {
321template <
bool is_transposed,
bool is_rhs_,
bool packet_load_,
typename PacketType>
323 static constexpr bool packet_load = packet_load_;
324 typedef typename Eigen::internal::unpacket_traits<PacketType>::type OutScalar;
325 static constexpr bool is_rhs = is_rhs_;
326 typedef std::conditional_t<packet_load, PacketType, OutScalar> OutType;
327 static constexpr int elements_per_access = Eigen::internal::unpacket_traits<OutType>::size;
328 static constexpr bool is_coalesced_layout = !(is_transposed ^ is_rhs);
329 static constexpr int nc_stride = (is_coalesced_layout ? elements_per_access : 1);
330 static constexpr int c_stride = (is_coalesced_layout ? 1 : elements_per_access);
372template <
typename StorageIndex>
373struct ThreadProperties {
374 const StorageIndex linearLocalThreadId;
375 const StorageIndex kGroupId;
376 const StorageIndex mGroupOffset;
377 const StorageIndex nGroupOffset;
378 const StorageIndex kGroupOffset;
379 const StorageIndex mLocalOffset;
380 const StorageIndex nLocalOffset;
381 const StorageIndex mGlobalOffset;
382 const StorageIndex nGlobalOffset;
384 const bool is_internal;
386 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE ThreadProperties(
387 const StorageIndex linearLocalThreadId_,
const StorageIndex kGroupId_,
const StorageIndex mGroupOffset_,
388 const StorageIndex nGroupOffset_,
const StorageIndex kGroupOffset_,
const StorageIndex mLocalOffset_,
389 const StorageIndex nLocalOffset_,
const StorageIndex mGlobalOffset_,
const StorageIndex nGlobalOffset_,
390 StorageIndex kSize_,
const bool is_internal_)
391 : linearLocalThreadId(linearLocalThreadId_),
393 mGroupOffset(mGroupOffset_),
394 nGroupOffset(nGroupOffset_),
395 kGroupOffset(kGroupOffset_),
396 mLocalOffset(mLocalOffset_),
397 nLocalOffset(nLocalOffset_),
398 mGlobalOffset(mGlobalOffset_),
399 nGlobalOffset(nGlobalOffset_),
401 is_internal(is_internal_) {}
454template <
typename OutScalar,
typename LhsScalar,
typename RhsScalar,
typename OutAccessor,
typename LhsMapper,
455 typename RhsMapper,
typename StorageIndex,
typename Properties,
typename TripleDim,
bool Vectorizable,
456 typename input_mapper_properties,
bool IsFinal, contraction_type contraction_tp>
457class TensorContractionKernel {
459 typedef typename Eigen::TensorSycl::internal::Vectorise<OutScalar, Eigen::SyclDevice, Vectorizable>::PacketReturnType
461 static constexpr int PacketSize =
462 Eigen::TensorSycl::internal::Vectorise<OutScalar, Eigen::SyclDevice, Vectorizable>::PacketSize;
463 static constexpr bool is_lhs_transposed =
464 !::Eigen::internal::TensorContractionInputMapperTrait<LhsMapper>::inner_dim_contiguous;
465 static constexpr bool is_rhs_transposed =
466 !::Eigen::internal::TensorContractionInputMapperTrait<RhsMapper>::inner_dim_contiguous;
468 typedef BlockProperties<is_lhs_transposed,
false, input_mapper_properties::is_lhs_matrix && Vectorizable,
472 typedef BlockProperties<is_rhs_transposed,
true, input_mapper_properties::is_rhs_matrix && Vectorizable,
476 static constexpr StorageIndex NStride =
477 contraction_tp == contraction_type::local ? Properties::WorkLoadPerThreadN : RHSBlockProperties::nc_stride;
479 typedef cl::sycl::accessor<OutScalar, 1, cl::sycl::access::mode::read_write, cl::sycl::access::target::local> Scratch;
480 typedef cl::sycl::multi_ptr<OutScalar, cl::sycl::access::address_space::local_space> local_ptr;
481 typedef OutScalar * private_ptr;
482 typedef std::conditional_t<contraction_tp == contraction_type::local, local_ptr, private_ptr> tile_ptr;
483 static constexpr StorageIndex LSDL = contraction_tp == contraction_type::local
484 ? Properties::TileSizeDimM + Properties::BC
485 : Properties::WorkLoadPerThreadM;
486 static constexpr StorageIndex LSDR = contraction_tp == contraction_type::local
487 ? Properties::TileSizeDimN + Properties::BC
488 : Properties::WorkLoadPerThreadN;
489 static constexpr StorageIndex LocalOffset = Properties::LocalThreadSizeM * Properties::LocalThreadSizeN;
503 template <contraction_type, StorageIndex>
506 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE MemHolder(local_ptr block_start_ptr) : ptr(block_start_ptr) {}
511 template <StorageIndex MemSize>
512 struct MemHolder<contraction_type::no_local, MemSize> {
513 OutScalar ptr[MemSize] = {OutScalar{0}};
540 tile_ptr lhs_scratch_ptr_compute;
541 tile_ptr rhs_scratch_ptr_compute;
542 const std::pair<StorageIndex, StorageIndex> lhs_extract_index;
543 const std::pair<StorageIndex, StorageIndex> rhs_extract_index;
544 template <contraction_type tp = contraction_tp, std::enable_if_t<tp == contraction_type::no_local,
int> = 0>
546 : lhs_scratch_extract{},
547 rhs_scratch_extract{},
548 lhs_scratch_ptr_compute(lhs_scratch_extract.ptr),
549 rhs_scratch_ptr_compute(rhs_scratch_extract.ptr),
550 lhs_extract_index(std::pair<StorageIndex, StorageIndex>(StorageIndex{0}, StorageIndex{0})),
551 rhs_extract_index(std::pair<StorageIndex, StorageIndex>(StorageIndex{0}, StorageIndex{0})) {}
553 template <contraction_type tp = contraction_tp, std::enable_if_t<tp == contraction_type::local,
int> = 0>
555 local_ptr block_start_ptr)
556 : lhs_scratch_extract{block_start_ptr},
557 rhs_scratch_extract{lhs_scratch_extract.ptr +
558 ((Properties::DoubleBuffer + 1) * LSDL * Properties::TileSizeDimK)},
559 lhs_scratch_ptr_compute(lhs_scratch_extract.ptr + thread_properties.mLocalOffset),
560 rhs_scratch_ptr_compute(rhs_scratch_extract.ptr + thread_properties.nLocalOffset),
562 local_id_extract<LHSBlockProperties, Properties::TileSizeDimM>(thread_properties.linearLocalThreadId)),
564 local_id_extract<RHSBlockProperties, Properties::TileSizeDimN>(thread_properties.linearLocalThreadId)) {}
571 const StorageIndex groupSizeM;
572 const StorageIndex groupSizeN;
573 const StorageIndex numTiles;
574 const TripleDim triple_dim;
576 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorContractionKernel(Scratch scratch_,
const LhsMapper lhs_,
577 const RhsMapper rhs_, OutAccessor out_res_,
578 const StorageIndex groupSizeM_,
579 const StorageIndex groupSizeN_,
580 const StorageIndex numTiles_,
581 const TripleDim triple_dim_)
586 groupSizeM(groupSizeM_),
587 groupSizeN(groupSizeN_),
589 triple_dim(triple_dim_) {}
592 const RhsMapper rhs_, OutAccessor out_res_,
593 const StorageIndex groupSizeM_,
594 const StorageIndex numTiles_,
595 const TripleDim triple_dim_)
598 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void operator()(cl::sycl::nd_item<1> itemID)
const {
599 const StorageIndex linearLocalThreadId = itemID.get_local_id(0);
600 const StorageIndex nLocalThreadId = linearLocalThreadId / Properties::LocalThreadSizeM;
601 const StorageIndex mLocalThreadId = linearLocalThreadId % Properties::LocalThreadSizeM;
602 const StorageIndex mGroupId = itemID.get_group(0) % groupSizeM;
603 const StorageIndex tmp = itemID.get_group(0) / groupSizeM;
604 const StorageIndex nGroupId = IsFinal ? tmp : tmp % groupSizeN;
605 const StorageIndex kGroupId = IsFinal ? 0 : tmp / groupSizeN;
606 const StorageIndex mGroupOffset = mGroupId * Properties::TileSizeDimM;
607 const StorageIndex nGroupOffset = nGroupId * Properties::TileSizeDimN;
608 const StorageIndex mLocalOffset = PacketSize * mLocalThreadId;
609 const StorageIndex nLocalOffset = NStride * nLocalThreadId;
610 const StorageIndex mGlobalOffset = mGroupOffset + mLocalOffset;
611 const StorageIndex nGlobalOffset = nGroupOffset + nLocalOffset;
613 const StorageIndex kSizePerWG = IsFinal ? triple_dim.K : numTiles * Properties::TileSizeDimK;
614 StorageIndex kGroupOffset = kGroupId * kSizePerWG;
615 const bool is_internal = triple_dim.M - mGroupOffset >= Properties::TileSizeDimM &&
616 triple_dim.N - nGroupOffset >= Properties::TileSizeDimN &&
617 triple_dim.K - kGroupOffset >= kSizePerWG;
619 StorageIndex kSize = IsFinal ? triple_dim.K : std::min(kSizePerWG, triple_dim.K - kGroupOffset);
622 kGroupOffset += kSize;
624 auto thread_properties =
625 ThreadProperties<StorageIndex>(linearLocalThreadId, kGroupId, mGroupOffset, nGroupOffset, kGroupOffset,
626 mLocalOffset, nLocalOffset, mGlobalOffset, nGlobalOffset, kSize, is_internal);
628 auto out_ptr = out_res + (IsFinal ? 0 : thread_properties.kGroupId * triple_dim.M * triple_dim.N);
630 (thread_properties.is_internal) ? compute_panel<true>(itemID, thread_properties, out_ptr)
631 : compute_panel<false>(itemID, thread_properties, out_ptr);
636 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void compute_block_per_tile(OutScalar *lhs_block_ptr, OutScalar *rhs_block_ptr,
637 PacketReturnType *privateRes)
const {
638 StorageIndex idx = 0;
639 constexpr StorageIndex lhs_stride =
640 contraction_tp == contraction_type::local ? (PacketSize * Properties::LocalThreadSizeM) : 1;
642 for (StorageIndex wLPTN = 0; wLPTN < Properties::WorkLoadPerThreadN; wLPTN++) {
643 auto rhsPacket = PacketReturnType{*(rhs_block_ptr + wLPTN)};
644 StorageIndex lhs_index = 0;
646 for (StorageIndex wLPTM = 0; wLPTM < Properties::WorkLoadPerThreadM / PacketSize; wLPTM++) {
647 PacketReturnType lhsPack{};
648 Eigen::TensorSycl::internal::PacketWrapper<PacketReturnType, PacketSize>::set_packet(lhsPack,
649 lhs_block_ptr + lhs_index);
650 privateRes[idx] = ::Eigen::internal::pmadd(lhsPack, rhsPacket, privateRes[idx]);
652 lhs_index += lhs_stride;
660 template <
bool is_
internal_block, StorageIndex PrivateNStr
ide,
typename OutPtr>
661 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void store(OutPtr *out_ptr, PacketReturnType *privateRes,
662 StorageIndex mGlobalOffset, StorageIndex nGlobalOffset)
const {
663 auto chk_bound = [&](
const StorageIndex &mIndex,
const StorageIndex &nIndex) EIGEN_DEVICE_FUNC {
664 return (mIndex + PacketSize - 1 < triple_dim.M && nGlobalOffset + nIndex < triple_dim.N);
669 constexpr StorageIndex GlobalNStride = contraction_tp == contraction_type::local ? 1 : Properties::LocalThreadSizeN;
671 for (StorageIndex wLPTN = 0; wLPTN < Properties::WorkLoadPerThreadN / PrivateNStride; wLPTN++) {
673 StorageIndex outputLD = 0;
678 for (StorageIndex nId = 0; nId < PrivateNStride; nId++) {
679 StorageIndex globalRow = mGlobalOffset;
681 for (StorageIndex wLPTM = 0; wLPTM < Properties::WorkLoadPerThreadM / PacketSize; wLPTM++) {
682 PacketReturnType privetOut = privateRes[wLPTM];
683 if (check_boundary<is_internal_block>(chk_bound(globalRow, nId))) {
686 write<data_source::global_mem>(privetOut, out_ptr + outputLD + globalRow);
689 for (StorageIndex mId = 0; mId < PacketSize; mId++) {
690 StorageIndex mOffset = globalRow + mId;
691 if (mOffset < triple_dim.M && (nGlobalOffset + nId < triple_dim.N)) {
692 out_ptr[mOffset + outputLD] =
693 Eigen::TensorSycl::internal::PacketWrapper<PacketReturnType, PacketSize>::scalarize(mId, privetOut);
697 globalRow += (PacketSize * Properties::LocalThreadSizeM);
699 outputLD += triple_dim.M;
700 privateRes += Properties::WorkLoadPerThreadM / PacketSize;
702 out_ptr += (GlobalNStride * outputLD);
704 nGlobalOffset += (PrivateNStride * GlobalNStride);
708 template <
typename InputBlockProperties,
bool is_internal_block,
typename Input,
typename PrivateReg,
709 contraction_type contract_tp = contraction_tp>
710 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<contract_tp == contraction_type::no_local> extract_block(
711 const Input &inpt, PrivateReg private_ptr,
const std::pair<StorageIndex, StorageIndex> &,
712 const StorageIndex &ncOffset,
const StorageIndex cOffset)
const {
713 constexpr StorageIndex LocalThreadSizeNC =
714 InputBlockProperties::is_rhs ? Properties::LocalThreadSizeN : Properties::LocalThreadSizeM;
715 constexpr StorageIndex WorkLoadPerThreadNC =
716 InputBlockProperties::is_rhs ? Properties::WorkLoadPerThreadN : Properties::WorkLoadPerThreadM;
717 const StorageIndex &NC = InputBlockProperties::is_rhs ? triple_dim.N : triple_dim.M;
719 auto chk_bound = [&](
const StorageIndex &CIndex,
const StorageIndex &NCIndex) EIGEN_DEVICE_FUNC {
720 return ((CIndex + InputBlockProperties::c_stride - 1 < triple_dim.K) &&
721 (NCIndex + InputBlockProperties::nc_stride - 1 < NC));
723 const StorageIndex ld = InputBlockProperties::is_coalesced_layout ? NC : triple_dim.K;
724 StorageIndex cIndex = cOffset;
727 for (StorageIndex cId = 0; cId < Properties::TileSizeDimK / InputBlockProperties::c_stride; cId++) {
728 StorageIndex ncIndex = ncOffset;
730 for (StorageIndex ncId = 0; ncId < WorkLoadPerThreadNC / InputBlockProperties::nc_stride; ncId++) {
731 if (check_boundary<is_internal_block>(chk_bound(cIndex, ncIndex))) {
733 read<InputBlockProperties::packet_load, InputBlockProperties::is_coalesced_layout,
734 InputBlockProperties::is_rhs,
typename InputBlockProperties::OutType>(inpt, ncIndex, cIndex, ld);
736 write<StorageIndex, (InputBlockProperties::is_coalesced_layout ? 1 : WorkLoadPerThreadNC),
737 data_source::private_mem>(val, private_ptr);
740 for (StorageIndex i = 0; i < InputBlockProperties::elements_per_access; i++) {
741 const StorageIndex ncInd = ncIndex + (InputBlockProperties::is_coalesced_layout ? i : 0);
742 const StorageIndex cInd = cIndex + (InputBlockProperties::is_coalesced_layout ? 0 : i);
744 (ncInd < NC && cInd < triple_dim.K)
745 ? read<false, InputBlockProperties::is_coalesced_layout, InputBlockProperties::is_rhs, OutScalar>(
746 inpt, ncInd, cInd, ld)
748 write<StorageIndex, (InputBlockProperties::is_coalesced_layout ? 1 : WorkLoadPerThreadNC),
749 data_source::private_mem>(
750 val, private_ptr + (InputBlockProperties::is_coalesced_layout ? i : 0) +
751 ((InputBlockProperties::is_coalesced_layout ? 0 : i) * WorkLoadPerThreadNC));
757 ncIndex = (!InputBlockProperties::is_rhs && InputBlockProperties::nc_stride == 1 && PacketSize != 1)
758 ? ncOffset + (ncId + 1) % PacketSize + ((ncId + 1) / PacketSize) * LocalThreadSizeNC
759 : (ncIndex + InputBlockProperties::nc_stride * LocalThreadSizeNC);
760 private_ptr += InputBlockProperties::nc_stride;
763 private_ptr += (InputBlockProperties::c_stride - 1) * WorkLoadPerThreadNC;
764 cIndex += InputBlockProperties::c_stride;
767 template <
typename InputBlockProperties, StorageIndex TileSizeDimNC>
768 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::pair<StorageIndex, StorageIndex> local_id_extract(
769 const StorageIndex &linearLocalThreadId) {
770 const StorageIndex localThreadNC =
771 (InputBlockProperties::is_coalesced_layout)
772 ? linearLocalThreadId % (TileSizeDimNC / InputBlockProperties::nc_stride)
773 : linearLocalThreadId / (Properties::TileSizeDimK / InputBlockProperties::c_stride);
774 const StorageIndex localThreadC =
775 (InputBlockProperties::is_coalesced_layout)
776 ? linearLocalThreadId / (TileSizeDimNC / InputBlockProperties::nc_stride)
777 : linearLocalThreadId % (Properties::TileSizeDimK / InputBlockProperties::c_stride);
778 return std::pair<StorageIndex, StorageIndex>(localThreadNC, localThreadC);
781 template <
bool db = Properties::DoubleBuffer, contraction_type ctp = contraction_tp>
782 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<db && ctp == contraction_type::local> sync_mem(
783 const cl::sycl::nd_item<1> &,
bool &db_offset)
noexcept {
784 db_offset = !db_offset;
787 template <
bool db = Properties::DoubleBuffer, contraction_type ctp = contraction_tp>
788 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<!db && ctp == contraction_type::local> sync_mem(
789 const cl::sycl::nd_item<1> &itemID,
bool &)
noexcept {
790 itemID.barrier(cl::sycl::access::fence_space::local_space);
793 template <contraction_type ctp = contraction_tp>
794 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<ctp == contraction_type::no_local> sync_mem(
795 const cl::sycl::nd_item<1> &,
bool &)
noexcept {
799 template <
bool need_sync, contraction_type ctp = contraction_tp>
800 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<need_sync && ctp == contraction_type::no_local>
801 sync_thread(
const cl::sycl::nd_item<1> &
802#ifdef EIGEN_SYCL_ARM_GPU_CACHE_OPTIMISATION
806#ifdef EIGEN_SYCL_ARM_GPU_CACHE_OPTIMISATION
807 itemID.barrier(cl::sycl::access::fence_spacce::local_space);
812 template <
bool need_sync, contraction_type ctp = contraction_tp>
813 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<need_sync && ctp == contraction_type::local>
814 sync_thread(
const cl::sycl::nd_item<1> &itemID) {
815 itemID.barrier(cl::sycl::access::fence_space::local_space);
817 template <
bool need_sync>
818 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<!need_sync> sync_thread(
const cl::sycl::nd_item<1> &) {
822 template <
bool is_
internal_block>
823 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void compute_tile_per_panel(
const cl::sycl::nd_item<1> &itemID,
824 ThreadProperties<StorageIndex> &thread_properties,
826 PacketReturnType *privateRes,
827 bool &db_offset)
const {
829 extract_block<RHSBlockProperties, is_internal_block>(
830 rhs, tiled_input_block.rhs_scratch_extract.ptr + (db_offset * Properties::TileSizeDimK * LSDR),
831 tiled_input_block.rhs_extract_index,
832 contraction_tp == contraction_type::local ? thread_properties.nGroupOffset : thread_properties.nGlobalOffset,
833 thread_properties.kGroupOffset - thread_properties.kSize);
835 sync_thread<contraction_tp == contraction_type::no_local>(itemID);
838 extract_block<LHSBlockProperties, is_internal_block>(
839 lhs, tiled_input_block.lhs_scratch_extract.ptr + (db_offset * LSDL * Properties::TileSizeDimK),
840 tiled_input_block.lhs_extract_index,
841 contraction_tp == contraction_type::local ? thread_properties.mGroupOffset : thread_properties.mGlobalOffset,
842 thread_properties.kGroupOffset - thread_properties.kSize);
844 sync_thread<contraction_tp == contraction_type::local>(itemID);
846 StorageIndex lhs_offset = (db_offset * LSDL * Properties::TileSizeDimK);
847 StorageIndex rhs_offset = (db_offset * Properties::TileSizeDimK * LSDR);
849 for (StorageIndex k = 0; k < Properties::TileSizeDimK; k++) {
850 compute_block_per_tile(tiled_input_block.lhs_scratch_ptr_compute + lhs_offset,
851 tiled_input_block.rhs_scratch_ptr_compute + rhs_offset, privateRes);
856 thread_properties.kSize -= Properties::TileSizeDimK;
857 sync_mem(itemID, db_offset);
861 template <
bool is_
internal_block,
typename OutPtr>
862 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void compute_panel(
const cl::sycl::nd_item<1> &itemID,
863 ThreadProperties<StorageIndex> &thread_properties,
864 OutPtr out_ptr)
const {
865 auto tiled_input_block =
TiledMemory{thread_properties, scratch.get_pointer()};
867 PacketReturnType privateRes[Properties::WorkLoadPerThreadM * Properties::WorkLoadPerThreadN / PacketSize] = {
868 PacketReturnType{0}};
871 while (thread_properties.kSize >= Properties::TileSizeDimK) {
872 compute_tile_per_panel<is_internal_block>(itemID, thread_properties, tiled_input_block, privateRes, db_offset);
874 if (thread_properties.kSize > 0) {
875 compute_tile_per_panel<false>(itemID, thread_properties, tiled_input_block, privateRes, db_offset);
879 store<is_internal_block,
880 contraction_tp == contraction_type::local ?
static_cast<StorageIndex
>(1) : RHSBlockProperties::nc_stride>(
881 out_ptr + thread_properties.nGlobalOffset * triple_dim.M, privateRes, thread_properties.mGlobalOffset,
882 thread_properties.nGlobalOffset);
885 template <
typename InputBlockProperties,
bool is_internal_block,
typename Input,
typename Local,
886 contraction_type contract_tp = contraction_tp>
887 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<contract_tp == contraction_type::local> extract_block(
888 const Input &inpt, Local local_ptr,
const std::pair<StorageIndex, StorageIndex> &local_index,
889 const StorageIndex &ncOffset,
const StorageIndex cOffset)
const {
890 constexpr StorageIndex TileSizeDimNC =
891 InputBlockProperties::is_rhs ? Properties::TileSizeDimN : Properties::TileSizeDimM;
892 constexpr StorageIndex LoadPerThread =
893 InputBlockProperties::is_rhs ? Properties::LoadPerThreadRhs : Properties::LoadPerThreadLhs;
894 constexpr StorageIndex LSD = InputBlockProperties::is_rhs ? LSDR : LSDL;
895 static_assert(((LocalOffset % (TileSizeDimNC / InputBlockProperties::nc_stride) == 0) &&
896 (LocalOffset % (Properties::TileSizeDimK / InputBlockProperties::c_stride) == 0)),
897 " LocalOffset must be divisible by stride");
898 const StorageIndex &NC = InputBlockProperties::is_rhs ? triple_dim.N : triple_dim.M;
899 StorageIndex localThreadNC = local_index.first;
900 StorageIndex localThreadC = local_index.second;
901 auto chk_bound = [&](
const StorageIndex &CIndex,
const StorageIndex &NCIndex) EIGEN_DEVICE_FUNC {
902 return ((CIndex + InputBlockProperties::c_stride - 1 < triple_dim.K) &&
903 (NCIndex + InputBlockProperties::nc_stride - 1 < NC));
906 for (StorageIndex lPT = 0; lPT < LoadPerThread / InputBlockProperties::elements_per_access; lPT++) {
907 const StorageIndex CIndex = cOffset + (InputBlockProperties::c_stride * localThreadC);
908 const StorageIndex NCIndex = ncOffset + (InputBlockProperties::nc_stride * localThreadNC);
909 const StorageIndex ld = InputBlockProperties::is_coalesced_layout ? NC : triple_dim.K;
910 if (check_boundary<is_internal_block>(chk_bound(CIndex, NCIndex))) {
912 read<InputBlockProperties::packet_load, InputBlockProperties::is_coalesced_layout,
913 InputBlockProperties::is_rhs,
typename InputBlockProperties::OutType>(inpt, NCIndex, CIndex, ld);
914 write<StorageIndex, (InputBlockProperties::is_coalesced_layout ? 1 : LSD), data_source::local_mem>(
915 val, local_ptr + (InputBlockProperties::nc_stride * localThreadNC) +
916 (InputBlockProperties::c_stride * localThreadC * LSD));
919 for (StorageIndex i = 0; i < InputBlockProperties::elements_per_access; i++) {
920 const StorageIndex nCInd = NCIndex + (InputBlockProperties::is_coalesced_layout ? i : 0);
921 const StorageIndex cInd = CIndex + (InputBlockProperties::is_coalesced_layout ? 0 : i);
923 (nCInd < NC && cInd < triple_dim.K)
924 ? read<false, InputBlockProperties::is_coalesced_layout, InputBlockProperties::is_rhs, OutScalar>(
925 inpt, nCInd, cInd, ld)
928 write<StorageIndex, (InputBlockProperties::is_coalesced_layout ? 1 : LSD), data_source::local_mem>(
929 val, local_ptr + (InputBlockProperties::nc_stride * localThreadNC) +
930 (InputBlockProperties::is_coalesced_layout ? i : 0) +
931 ((InputBlockProperties::c_stride * localThreadC +
932 (InputBlockProperties::is_coalesced_layout ? 0 : i)) *
936 localThreadNC += (InputBlockProperties::is_coalesced_layout)
937 ? LocalOffset % (TileSizeDimNC / InputBlockProperties::nc_stride)
938 : LocalOffset / (Properties::TileSizeDimK / InputBlockProperties::c_stride);
939 localThreadC += (InputBlockProperties::is_coalesced_layout)
940 ? LocalOffset / (TileSizeDimNC / InputBlockProperties::nc_stride)
941 : LocalOffset % (Properties::TileSizeDimK / InputBlockProperties::c_stride);
946#ifndef EIGEN_SYCL_DISABLE_GEMV
989template <
typename OutScalar,
typename OutAccessor,
typename VectorMapper,
typename TensorMapper,
typename StorageIndex,
990 typename Properties, StorageIndex KFactor,
bool Vectorizable,
bool is_lhs_vec,
bool IsFinal>
991struct GeneralVectorTensor {
992 typedef typename Eigen::TensorSycl::internal::Vectorise<OutScalar, Eigen::SyclDevice, Vectorizable>::PacketReturnType
994 static constexpr int PacketSize =
995 Eigen::TensorSycl::internal::Vectorise<OutScalar, Eigen::SyclDevice, Vectorizable>::PacketSize;
996 typedef cl::sycl::accessor<OutScalar, 1, cl::sycl::access::mode::read_write, cl::sycl::access::target::local> Scratch;
998 static constexpr StorageIndex OutScratchOffset =
999 KFactor * Properties::LocalThreadSizeC * Properties::LocalThreadSizeNC;
1006 const VectorMapper vec;
1007 const TensorMapper mat;
1008 OutAccessor out_res;
1009 const StorageIndex nonContractGroupSize;
1010 const StorageIndex nonContractDim;
1011 const StorageIndex contractDim;
1013 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE GeneralVectorTensor(Scratch scratch_,
const VectorMapper vec_,
1014 const TensorMapper mat_, OutAccessor out_res_,
1015 const StorageIndex nonContractGroupSize_,
1016 const StorageIndex nonContractDim_,
1017 const StorageIndex contractDim_)
1018 : scratch(scratch_),
1022 nonContractGroupSize(nonContractGroupSize_),
1023 nonContractDim(nonContractDim_),
1024 contractDim(contractDim_) {}
1026 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void operator()(cl::sycl::nd_item<1> itemID)
const {
1027 auto scratch_ptr = scratch.get_pointer();
1028 const StorageIndex linearLocalThreadId = itemID.get_local_id(0);
1029 StorageIndex nonContractId = is_lhs_vec ? linearLocalThreadId / Properties::LocalThreadSizeC
1030 : linearLocalThreadId % Properties::LocalThreadSizeNC;
1031 StorageIndex contractId = is_lhs_vec ? linearLocalThreadId % Properties::LocalThreadSizeC
1032 : linearLocalThreadId / Properties::LocalThreadSizeNC;
1033 const StorageIndex cGroupSize = itemID.get_group_range(0) / nonContractGroupSize;
1034 const StorageIndex nonContractGroupId =
1035 is_lhs_vec ? itemID.get_group(0) / cGroupSize : itemID.get_group(0) % nonContractGroupSize;
1036 const StorageIndex contractGroupId =
1037 is_lhs_vec ? itemID.get_group(0) % cGroupSize : itemID.get_group(0) / nonContractGroupSize;
1038 auto out_ptr = out_res + (IsFinal ? 0 : contractGroupId * nonContractDim);
1040 const StorageIndex nonContractGroupOffset = nonContractGroupId * Properties::TileSizeDimNC;
1041 const StorageIndex contractGroupOffset = contractGroupId * Properties::TileSizeDimC;
1042 auto outScratchIndex = nonContractId + contractId * Properties::LocalThreadSizeNC;
1043 const StorageIndex globalNonContractDimOffset = nonContractGroupOffset + nonContractId;
1044 const StorageIndex globalContractDimOffset = contractGroupOffset + contractId;
1045 auto local_output = scratch_ptr + OutScratchOffset;
1046 const bool is_internal = nonContractDim - nonContractGroupOffset >= Properties::TileSizeDimNC &&
1047 contractDim - contractGroupOffset >= Properties::TileSizeDimC;
1049 ? compute_panel<true>(itemID, vec, mat, local_output, out_ptr,
1050#ifdef EIGEN_SYCL_LOCAL_MEM_UNSET_OR_ON
1051 scratch_ptr, contractGroupOffset,
1053 nonContractGroupOffset, linearLocalThreadId, contractDim, nonContractDim, contractId,
1054 nonContractId, globalContractDimOffset, globalNonContractDimOffset, outScratchIndex)
1055 : compute_panel<false>(itemID, vec, mat, local_output, out_ptr,
1056#ifdef EIGEN_SYCL_LOCAL_MEM_UNSET_OR_ON
1057 scratch_ptr, contractGroupOffset,
1059 nonContractGroupOffset, linearLocalThreadId, contractDim, nonContractDim, contractId,
1060 nonContractId, globalContractDimOffset, globalNonContractDimOffset, outScratchIndex);
1062 template <
bool is_
internal_block,
typename OutPtr>
1063 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void compute_panel(
1064 const cl::sycl::nd_item<1> &itemID,
const VectorMapper &vec,
const TensorMapper &mat, OutScalar *local_output,
1066#ifdef EIGEN_SYCL_LOCAL_MEM_UNSET_OR_ON
1067 OutScalar *scratch_ptr,
const StorageIndex contractGroupOffset,
1069 const StorageIndex nonContractGroupOffset,
const StorageIndex linearLocalThreadId, StorageIndex contractDim,
1070 StorageIndex nonContractDim, StorageIndex contractId, StorageIndex nonContractId,
1071 StorageIndex globalContractDimOffset, StorageIndex globalNonContractDimOffset, StorageIndex outScratchIndex) {
1072 OutScalar outScalar[Properties::WorkLoadPerThreadNC] = {OutScalar(0)};
1074#ifdef EIGEN_SYCL_LOCAL_MEM_UNSET_OR_ON
1075 const StorageIndex vectorOffset = contractGroupOffset + linearLocalThreadId;
1076 extract_block<VecBlockProperties, is_internal_block, KFactor,
1077 Properties::LocalThreadSizeNC * Properties::LocalThreadSizeC>(vec, scratch_ptr, linearLocalThreadId,
1078 vectorOffset, contractDim);
1080 itemID.barrier(cl::sycl::access::fence_space::local_space);
1081 auto in_scratch_ptr = scratch_ptr + contractId;
1084 StorageIndex privateOffsetC = 0;
1086 for (StorageIndex i = 0; i < Properties::WorkLoadPerThreadC; i++) {
1087 StorageIndex privateOffsetNC = 0;
1088 bool contract_conds = ((globalContractDimOffset + privateOffsetC) < contractDim);
1089#ifdef EIGEN_SYCL_LOCAL_MEM_UNSET_OR_ON
1090 auto vecScalar = *in_scratch_ptr;
1092 auto vecScalar = (check_boundary<is_internal_block>(contract_conds))
1093 ? vec(is_lhs_vec ? StorageIndex(0) : globalContractDimOffset + privateOffsetC,
1094 is_lhs_vec ? globalContractDimOffset + privateOffsetC : StorageIndex(0))
1098 for (StorageIndex j = 0; j < Properties::WorkLoadPerThreadNC; j++) {
1099 auto matScalar = (check_boundary<is_internal_block>(
1100 contract_conds && ((globalNonContractDimOffset + privateOffsetNC) < nonContractDim)))
1101 ? mat(is_lhs_vec ? globalContractDimOffset + privateOffsetC
1102 : globalNonContractDimOffset + privateOffsetNC,
1103 is_lhs_vec ? globalNonContractDimOffset + privateOffsetNC
1104 : globalContractDimOffset + privateOffsetC)
1107 outScalar[j] = ::Eigen::internal::pmadd(matScalar, vecScalar, outScalar[j]);
1108 privateOffsetNC += Properties::LocalThreadSizeNC;
1110 privateOffsetC += Properties::LocalThreadSizeC;
1111#ifdef EIGEN_SYCL_LOCAL_MEM_UNSET_OR_ON
1112 in_scratch_ptr += Properties::LocalThreadSizeC;
1116 auto out_scratch_ptr = local_output + outScratchIndex;
1119 for (StorageIndex j = 0; j < Properties::WorkLoadPerThreadNC; j++) {
1120 *out_scratch_ptr = outScalar[j];
1122 out_scratch_ptr += (Properties::LocalThreadSizeNC * Properties::LocalThreadSizeC);
1124 EIGEN_IF_CONSTEXPR (is_lhs_vec) {
1125 nonContractId = linearLocalThreadId % Properties::LocalThreadSizeNC;
1126 contractId = linearLocalThreadId / Properties::LocalThreadSizeNC;
1127 outScratchIndex = nonContractId + contractId * Properties::LocalThreadSizeNC;
1130 out_scratch_ptr = local_output + outScratchIndex;
1132 for (StorageIndex j = 0; j < Properties::WorkLoadPerThreadNC; j++) {
1134 for (StorageIndex offset = Properties::LocalThreadSizeC >> 1; offset > 0; offset >>= 1) {
1135 itemID.barrier(cl::sycl::access::fence_space::local_space);
1136 if (contractId < offset) {
1137 StorageIndex myNeigbourId = (Properties::LocalThreadSizeNC * offset);
1138 *out_scratch_ptr += out_scratch_ptr[myNeigbourId];
1142 out_scratch_ptr += (Properties::LocalThreadSizeNC * Properties::LocalThreadSizeC);
1145 if (contractId == 0) {
1146 out_scratch_ptr = local_output + nonContractId;
1147 StorageIndex global_final_offset = nonContractGroupOffset + nonContractId;
1148 out_ptr += global_final_offset;
1150 for (StorageIndex j = 0; j < Properties::WorkLoadPerThreadNC; j++) {
1151 if (check_boundary<is_internal_block>(global_final_offset < nonContractDim)) {
1152 auto res = *out_scratch_ptr;
1155 out_ptr += Properties::LocalThreadSizeNC;
1158 out_scratch_ptr += (Properties::LocalThreadSizeNC * Properties::LocalThreadSizeC);
1159 if (!(is_internal_block)) global_final_offset += Properties::LocalThreadSizeNC;
1164 template <
typename InputBlockProperties,
bool is_internal_block,
int CFactor,
int GroupSize,
typename Input,
1166 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void extract_block(
const Input &inpt, Local *local_ptr,
1167 const StorageIndex &linearLocalThreadId,
1168 const StorageIndex &cOffset,
const StorageIndex &C) {
1169 local_ptr += InputBlockProperties::c_stride * linearLocalThreadId;
1170 StorageIndex cIndex = cOffset;
1171 for (StorageIndex cId = 0; cId < CFactor / InputBlockProperties::c_stride; cId++) {
1172 if (check_boundary<is_internal_block>(cIndex + InputBlockProperties::c_stride - 1 < C)) {
1173 auto val = read<InputBlockProperties::packet_load, InputBlockProperties::is_coalesced_layout,
1174 InputBlockProperties::is_rhs,
typename InputBlockProperties::OutType>(inpt, StorageIndex(0),
1175 cIndex, StorageIndex(1));
1176 write<StorageIndex, 1, data_source::local_mem>(val, local_ptr);
1179 for (StorageIndex i = 0; i < InputBlockProperties::elements_per_access; i++) {
1182 ? read<false, InputBlockProperties::is_coalesced_layout, InputBlockProperties::is_rhs, OutScalar>(
1183 inpt, StorageIndex(0), cIndex + i, StorageIndex(1))
1185 write<StorageIndex, 1, data_source::local_mem>(val, local_ptr + i);
1188 local_ptr += InputBlockProperties::c_stride * GroupSize;
1189 cIndex += InputBlockProperties::c_stride * GroupSize;
1195#ifndef EIGEN_SYCL_DISABLE_SCALAR
1228template <
typename OutScalar,
typename LhsScalar,
typename RhsScalar,
typename OutAccessor,
typename LhsMapper,
1229 typename RhsMapper,
typename StorageIndex,
bool Vectorizable>
1230struct GeneralScalarContraction {
1231 typedef cl::sycl::accessor<OutScalar, 1, cl::sycl::access::mode::read_write, cl::sycl::access::target::local> Scratch;
1233 const LhsMapper lhs;
1234 const RhsMapper rhs;
1235 OutAccessor out_res;
1236 const StorageIndex rng;
1238 EIGEN_DEVICE_FUNC GeneralScalarContraction(Scratch scratch_,
const LhsMapper lhs_,
const RhsMapper rhs_,
1239 OutAccessor out_res_,
const StorageIndex rng_)
1240 : scratch(scratch_), lhs(lhs_), rhs(rhs_), out_res(out_res_), rng(rng_) {}
1242 EIGEN_DEVICE_FUNC
void operator()(cl::sycl::nd_item<1> itemID)
const {
1243 auto out_ptr = out_res;
1244 OutScalar *scratch_ptr = scratch.get_pointer();
1246 StorageIndex globalid = itemID.get_global_id(0);
1247 StorageIndex localid = itemID.get_local_id(0);
1248 OutScalar accumulator = OutScalar(0);
1249 for (StorageIndex i = globalid; i < rng; i += itemID.get_global_range(0)) {
1250 accumulator = Eigen::internal::pmadd(lhs(0, i), rhs(i, 0), accumulator);
1252 auto out_scratch_ptr = scratch_ptr + localid;
1253 *out_scratch_ptr = accumulator;
1254 for (StorageIndex offset = itemID.get_local_range(0) >> 1; offset > 0; offset >>= 1) {
1255 itemID.barrier(cl::sycl::access::fence_space::local_space);
1256 if (localid < offset) {
1257 *out_scratch_ptr = (accumulator += out_scratch_ptr[offset]);
1261 out_ptr[itemID.get_group(0)] = accumulator;
1270template <
typename Indices,
typename LeftArgType,
typename RightArgType,
typename OutputKernelType>
1273 :
public TensorContractionEvaluatorBase<TensorEvaluator<
1274 const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, Eigen::SyclDevice>> {
1275 static_assert(std::is_same<OutputKernelType, const NoOpOutputKernel>::value,
1276 "SYCL tensor contraction does not support output kernels.");
1278 typedef Eigen::SyclDevice Device;
1281 typedef TensorContractionEvaluatorBase<Self> Base;
1283 typedef std::remove_const_t<typename XprType::Scalar> Scalar;
1284 typedef typename XprType::Index StorageIndex;
1286 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
1287 typedef typename Base::Storage Storage;
1288 typedef typename Base::EvaluatorPointerType EvaluatorPointerType;
1290 const StorageIndex M;
1291 const StorageIndex N;
1292 const StorageIndex K;
1293 TripleDim(
const StorageIndex M_,
const StorageIndex N_,
const StorageIndex K_) : M(M_), N(N_), K(K_) {}
1296 PacketAccess = (PacketType<CoeffReturnType, Device>::size > 1),
1297 BlockAccess =
false,
1300 static constexpr int Layout = TensorEvaluator<LeftArgType, Device>::Layout;
1301 static constexpr int LDims = Base::LDims;
1302 static constexpr int RDims = Base::RDims;
1303 static constexpr int ContractDims = Base::ContractDims;
1305 typedef array<StorageIndex, LDims> left_dim_mapper_t;
1306 typedef array<StorageIndex, RDims> right_dim_mapper_t;
1308 typedef array<StorageIndex, ContractDims> contract_t;
1309 typedef array<StorageIndex, LDims - ContractDims> left_nocontract_t;
1310 typedef array<StorageIndex, RDims - ContractDims> right_nocontract_t;
1312 static constexpr int NumDims = LDims + RDims - 2 * ContractDims;
1314 typedef DSizes<StorageIndex, NumDims> Dimensions;
1316 typedef TensorEvaluator<typename Base::EvalLeftArgType, Device> LeftEvaluator;
1317 typedef TensorEvaluator<typename Base::EvalRightArgType, Device> RightEvaluator;
1318 typedef std::remove_const_t<typename LeftEvaluator::CoeffReturnType> LhsScalar;
1319 typedef std::remove_const_t<typename RightEvaluator::CoeffReturnType> RhsScalar;
1321 typedef typename LeftEvaluator::Dimensions LeftDimensions;
1322 typedef typename RightEvaluator::Dimensions RightDimensions;
1324 template <
bool lhs_inner_dim_contiguous,
bool rhs_inner_dim_contiguous,
bool rhs_inner_dim_reordered>
1325 struct input_mapper_propertis {
1326 static constexpr bool is_lhs_matrix = (LDims == 2 && ContractDims == 1) || lhs_inner_dim_contiguous;
1327 static constexpr bool is_rhs_matrix =
1328 (RDims == 2 && ContractDims == 1) || (rhs_inner_dim_contiguous && !rhs_inner_dim_reordered);
1331 TensorEvaluator(
const XprType &op,
const Device &device) : Base(op, device) {}
1334 EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(
typename Base::EvaluatorPointerType data) {
1335 this->m_leftImpl.evalSubExprsIfNeeded(
nullptr);
1336 this->m_rightImpl.evalSubExprsIfNeeded(
nullptr);
1338 this->m_result = this->m_device.get(
1339 static_cast<Scalar *
>(this->m_device.allocate_temp(this->dimensions().TotalSize() *
sizeof(Scalar))));
1340 data = this->m_result;
1343 return (this->m_result !=
nullptr);
1345 const Eigen::SyclDevice &device()
const {
return this->m_device; }
1346 void evalToSycl(
typename Base::EvaluatorPointerType buffer)
const {
1347 if (this->m_lhs_inner_dim_contiguous) {
1348 if (this->m_rhs_inner_dim_contiguous) {
1349 if (this->m_rhs_inner_dim_reordered) {
1350 evalTyped<true, true, true, Unaligned>(buffer);
1352 evalTyped<true, true, false, Unaligned>(buffer);
1355 if (this->m_rhs_inner_dim_reordered) {
1356 evalTyped<true, false, true, Unaligned>(buffer);
1358 evalTyped<true, false, false, Unaligned>(buffer);
1362 if (this->m_rhs_inner_dim_contiguous) {
1363 if (this->m_rhs_inner_dim_reordered) {
1364 evalTyped<false, true, true, Unaligned>(buffer);
1366 evalTyped<false, true, false, Unaligned>(buffer);
1369 if (this->m_rhs_inner_dim_reordered) {
1370 evalTyped<false, false, true, Unaligned>(buffer);
1372 evalTyped<false, false, false, Unaligned>(buffer);
1378 template <
bool lhs_inner_dim_contiguous,
bool rhs_inner_dim_contiguous,
bool rhs_inner_dim_reordered,
int Alignment>
1379 void evalTyped(
typename Base::EvaluatorPointerType buffer)
const {
1380 const auto triple_dim = TripleDim{this->m_i_size, this->m_j_size, this->m_k_size};
1381 typedef internal::TensorContractionInputMapper<
1382 LhsScalar, StorageIndex, internal::Lhs, LeftEvaluator, left_nocontract_t, contract_t,
1383 PacketType<CoeffReturnType, Device>::size, lhs_inner_dim_contiguous,
false,
Unaligned, MakePointer>
1386 typedef internal::TensorContractionInputMapper<RhsScalar, StorageIndex, internal::Rhs, RightEvaluator,
1387 right_nocontract_t, contract_t,
1388 PacketType<CoeffReturnType, Device>::size, rhs_inner_dim_contiguous,
1389 rhs_inner_dim_reordered,
Unaligned, MakePointer>
1393 LhsMapper lhs(this->m_leftImpl, this->m_left_nocontract_strides, this->m_i_strides,
1394 this->m_left_contracting_strides, this->m_k_strides);
1396 RhsMapper rhs(this->m_rightImpl, this->m_right_nocontract_strides, this->m_j_strides,
1397 this->m_right_contracting_strides, this->m_k_strides);
1399#ifndef EIGEN_SYCL_DISABLE_SCALAR
1400 if (triple_dim.M == 1 && triple_dim.N == 1) {
1401 launchSC(buffer, lhs, rhs, triple_dim.K);
1404#ifndef EIGEN_SYCL_DISABLE_GEMV
1405 if (triple_dim.M != 1 && triple_dim.N == 1) {
1406 LaunchVT<false>(buffer, rhs, lhs, triple_dim.M, triple_dim.K);
1407 }
else if (triple_dim.M == 1 && triple_dim.N != 1) {
1408 LaunchVT<true>(buffer, lhs, rhs, triple_dim.N, triple_dim.K);
1412 typedef input_mapper_propertis<lhs_inner_dim_contiguous, rhs_inner_dim_contiguous, rhs_inner_dim_reordered>
1413 inpt_mapper_properties;
1414#ifndef EIGEN_SYCL_DISABLE_SKINNY
1415 bool skinny =
false;
1416 auto platform_name = this->device().getPlatformName();
1418 if (platform_name.find(
"AMD") == 0) {
1419 skinny = (triple_dim.M < triple_dim.K || triple_dim.N < triple_dim.K) &&
1420 ((triple_dim.M < 1024 && triple_dim.N < 1024) ||
1421 (uint64_t(triple_dim.M * triple_dim.N) < uint64_t(triple_dim.K)));
1423 skinny = (((std::max(triple_dim.K, triple_dim.N) / std::min(triple_dim.K, triple_dim.N)) > 100) ||
1424 ((std::max(triple_dim.K, triple_dim.M) / std::min(triple_dim.K, triple_dim.M)) > 100) ||
1425 ((std::max(triple_dim.N, triple_dim.M) / std::min(triple_dim.N, triple_dim.M)) > 100));
1428 adjustTT<true, inpt_mapper_properties>(buffer, lhs, rhs, triple_dim);
1431 adjustTT<false, inpt_mapper_properties>(buffer, lhs, rhs, triple_dim);
1435 template <
bool skinny,
typename input_mapper_properties,
typename LhsMapper,
typename RhsMapper>
1436 void EIGEN_ALWAYS_INLINE adjustTT(EvaluatorPointerType buffer,
const LhsMapper &lhs,
const RhsMapper &rhs,
1437 const TripleDim &triple_dim)
const {
1438#ifdef EIGEN_SYCL_LOCAL_MEM_UNSET_OR_ON
1439 if (device().has_local_memory()) {
1440 typedef TensorSycl::internal::TTPanelSize<CoeffReturnType, StorageIndex, 4, 4, 16> PanelParameters;
1441 launchTT<TensorSycl::internal::contraction_type::local, skinny, input_mapper_properties, PanelParameters>(
1442 buffer, lhs, rhs, triple_dim);
1445#ifdef EIGEN_SYCL_LOCAL_MEM_UNSET_OR_OFF
1446 if (!(device().has_local_memory())) {
1447 typedef TensorSycl::internal::TTPanelSize<CoeffReturnType, StorageIndex, 4, 4, 4> PanelParameters;
1448 launchTT<TensorSycl::internal::contraction_type::no_local, skinny, input_mapper_properties, PanelParameters>(
1449 buffer, lhs, rhs, triple_dim);
1454 template <TensorSycl::internal::contraction_type ct,
bool skinny,
typename input_mapper_properties,
1455 typename Properties,
typename LhsMapper,
typename RhsMapper>
1456 void launchTT(EvaluatorPointerType buffer,
const LhsMapper &lhs,
const RhsMapper &rhs,
1457 const TripleDim &triple_dim)
const {
1458 const StorageIndex roundUpM = Eigen::TensorSycl::internal::roundUp(triple_dim.M, Properties::TileSizeDimM);
1459 const StorageIndex roundUpN = Eigen::TensorSycl::internal::roundUp(triple_dim.N, Properties::TileSizeDimN);
1460 const StorageIndex groupSizeM = roundUpM / Properties::TileSizeDimM;
1461 const StorageIndex groupSizeN = roundUpN / Properties::TileSizeDimN;
1463 const StorageIndex roundUpK = Eigen::TensorSycl::internal::roundUp(triple_dim.K, Properties::TileSizeDimK);
1464 StorageIndex totalTilesK = roundUpK / Properties::TileSizeDimK;
1465 StorageIndex groupSizeK =
1467 ? std::max(std::min(totalTilesK,
1468 (StorageIndex)(device().getPowerOfTwo(device().getNumSyclMultiProcessors(),
true) * 4) /
1469 (groupSizeM * groupSizeN)),
1473 const StorageIndex numTilesPerGroup = Eigen::TensorSycl::internal::roundUp(totalTilesK, groupSizeK) / groupSizeK;
1475 const StorageIndex totalGroupSize = groupSizeM * groupSizeN * groupSizeK;
1477 const StorageIndex localRange = Properties::LocalThreadSizeM * Properties::LocalThreadSizeN;
1478 const StorageIndex globalRange = totalGroupSize * localRange;
1480 const StorageIndex scratchSize = (ct == TensorSycl::internal::contraction_type::local)
1481 ? ((Properties::DoubleBuffer + 1) *
1482 (Properties::TileSizeDimM + Properties::BC) * (Properties::TileSizeDimK)) +
1483 ((Properties::DoubleBuffer + 1) * (Properties::TileSizeDimK) *
1484 (Properties::TileSizeDimN + Properties::BC))
1487 auto thread_range = cl::sycl::nd_range<1>(cl::sycl::range<1>(globalRange), cl::sycl::range<1>(localRange));
1488 if (groupSizeK == 1) {
1489 typedef TensorSycl::internal::TensorContractionKernel<CoeffReturnType, LhsScalar, RhsScalar, EvaluatorPointerType,
1490 LhsMapper, RhsMapper, StorageIndex, Properties, TripleDim,
1491 PacketAccess, input_mapper_properties,
true, ct>
1494 .template binary_kernel_launcher<CoeffReturnType, ContractKernelName>(
1495 lhs, rhs, buffer, thread_range, scratchSize, groupSizeM, groupSizeN, numTilesPerGroup, triple_dim)
1498 typedef TensorSycl::internal::TensorContractionKernel<CoeffReturnType, LhsScalar, RhsScalar, EvaluatorPointerType,
1499 LhsMapper, RhsMapper, StorageIndex, Properties, TripleDim,
1500 PacketAccess, input_mapper_properties,
false, ct>
1502 CoeffReturnType *temp_pointer =
static_cast<CoeffReturnType *
>(
1503 device().allocate_temp(triple_dim.M * triple_dim.N * groupSizeK *
sizeof(CoeffReturnType)));
1504 EvaluatorPointerType tmp_global_accessor = device().get(temp_pointer);
1507 .template binary_kernel_launcher<CoeffReturnType, ContractKernelName>(
1508 lhs, rhs, tmp_global_accessor, thread_range, scratchSize, groupSizeM, groupSizeN, numTilesPerGroup,
1512 typedef Eigen::internal::SumReducer<CoeffReturnType> Op;
1514 typedef TensorSycl::internal::SecondStepPartialReduction<CoeffReturnType, StorageIndex, EvaluatorPointerType,
1515 EvaluatorPointerType, Op>
1519 .template unary_kernel_launcher<CoeffReturnType, ReductionKernel>(
1520 tmp_global_accessor, buffer,
1521 cl::sycl::nd_range<1>(cl::sycl::range<1>(StorageIndex(
1522 Eigen::TensorSycl::internal::roundUp(triple_dim.M * triple_dim.N, localRange))),
1523 cl::sycl::range<1>(localRange)),
1524 StorageIndex(1), op, StorageIndex(triple_dim.M * triple_dim.N), groupSizeK)
1526 device().deallocate_temp(temp_pointer);
1530#ifndef EIGEN_SYCL_DISABLE_GEMV
1531 template <
bool is_lhs_vec,
typename VectorMapper,
typename TensorMapper,
typename StorageIndex>
1532 void EIGEN_ALWAYS_INLINE LaunchVT(EvaluatorPointerType buffer,
const VectorMapper &vec,
const TensorMapper &mat,
1533 StorageIndex NC, StorageIndex C)
const {
1534 const StorageIndex nonContractDim = NC;
1535 constexpr StorageIndex NCFactor = 1;
1536 constexpr StorageIndex CFactor = 1;
1537 constexpr StorageIndex NCWindow = 16;
1538 typedef Eigen::TensorSycl::internal::TVPanelSize<CoeffReturnType, StorageIndex, NCWindow, CFactor, NCFactor>
1540 const StorageIndex roundUpC = Eigen::TensorSycl::internal::roundUp(C, Properties::TileSizeDimC);
1541 const StorageIndex cNumGroups = roundUpC / (Properties::LocalThreadSizeC * Properties::WorkLoadPerThreadC);
1542 const StorageIndex roundUpNC = Eigen::TensorSycl::internal::roundUp(nonContractDim, Properties::TileSizeDimNC);
1543 const StorageIndex nCNumGroups = roundUpNC / (Properties::LocalThreadSizeNC * Properties::WorkLoadPerThreadNC);
1544 const StorageIndex globalRange =
1545 (roundUpNC / (Properties::WorkLoadPerThreadNC)) * (roundUpC / (Properties::WorkLoadPerThreadC));
1546 const StorageIndex localRange = Properties::LocalThreadSizeNC * Properties::LocalThreadSizeC;
1547 const StorageIndex scratchSize =
1548 (Properties::WorkLoadPerThreadNC + CFactor) * Properties::LocalThreadSizeC * Properties::LocalThreadSizeNC;
1549 auto thread_range = cl::sycl::nd_range<1>(cl::sycl::range<1>(globalRange), cl::sycl::range<1>(localRange));
1550 if (cNumGroups > 1) {
1551 typedef Eigen::TensorSycl::internal::GeneralVectorTensor<CoeffReturnType, EvaluatorPointerType, VectorMapper,
1552 TensorMapper, StorageIndex, Properties, CFactor,
false,
1555 CoeffReturnType *temp_pointer =
1556 static_cast<CoeffReturnType *
>(device().allocate_temp(nonContractDim * cNumGroups *
sizeof(CoeffReturnType)));
1557 EvaluatorPointerType tmp_global_accessor = device().get(temp_pointer);
1560 .template binary_kernel_launcher<CoeffReturnType, ContractKernelName>(
1561 vec, mat, tmp_global_accessor, thread_range, scratchSize, nCNumGroups, nonContractDim, C)
1564 typedef Eigen::internal::SumReducer<CoeffReturnType> Op;
1565 typedef TensorSycl::internal::SecondStepPartialReduction<CoeffReturnType, StorageIndex, EvaluatorPointerType,
1566 EvaluatorPointerType, Op>
1570 .template unary_kernel_launcher<CoeffReturnType, ReductionKernel>(
1571 tmp_global_accessor, buffer,
1572 cl::sycl::nd_range<1>(
1573 cl::sycl::range<1>(Eigen::TensorSycl::internal::roundUp(nonContractDim, localRange)),
1574 cl::sycl::range<1>(localRange)),
1575 StorageIndex(1), Op(), nonContractDim, cNumGroups)
1577 device().deallocate_temp(temp_pointer);
1579 typedef Eigen::TensorSycl::internal::GeneralVectorTensor<CoeffReturnType, EvaluatorPointerType, VectorMapper,
1580 TensorMapper, StorageIndex, Properties, CFactor,
false,
1584 .template binary_kernel_launcher<CoeffReturnType, ContractKernelName>(
1585 vec, mat, buffer, thread_range, scratchSize, nCNumGroups, nonContractDim, C)
1591#ifndef EIGEN_SYCL_DISABLE_SCALAR
1592 template <
typename LhsMapper,
typename RhsMapper>
1593 EIGEN_ALWAYS_INLINE
void launchSC(EvaluatorPointerType buffer,
const LhsMapper &lhs,
const RhsMapper &rhs,
1594 StorageIndex K)
const {
1595 EIGEN_STATIC_ASSERT(!((EIGEN_SYCL_LOCAL_THREAD_DIM0 * EIGEN_SYCL_LOCAL_THREAD_DIM1) &
1596 (EIGEN_SYCL_LOCAL_THREAD_DIM0 * EIGEN_SYCL_LOCAL_THREAD_DIM1 - 1)),
1597 "The Local thread size must be a power of 2 for the reduction "
1599 constexpr StorageIndex local_range = EIGEN_SYCL_LOCAL_THREAD_DIM0 * EIGEN_SYCL_LOCAL_THREAD_DIM1;
1603 const StorageIndex num_work_group = ((K + (512 * local_range - 1)) / (512 * local_range) > 1 ? local_range : 1);
1604 const StorageIndex global_range = num_work_group * local_range;
1606 typedef Eigen::TensorSycl::internal::GeneralScalarContraction<
1607 CoeffReturnType, LhsScalar, RhsScalar, EvaluatorPointerType, LhsMapper, RhsMapper, StorageIndex,
false>
1609 auto thread_range = cl::sycl::nd_range<1>(cl::sycl::range<1>(global_range), cl::sycl::range<1>(local_range));
1610 if (num_work_group > 1) {
1611 CoeffReturnType *temp_pointer =
1612 static_cast<CoeffReturnType *
>(device().allocate_temp(num_work_group *
sizeof(CoeffReturnType)));
1613 EvaluatorPointerType tmp_global_accessor = device().get(temp_pointer);
1615 .template binary_kernel_launcher<CoeffReturnType, ContractKernelName>(lhs, rhs, tmp_global_accessor,
1616 thread_range, local_range, K)
1618 typedef Eigen::internal::SumReducer<CoeffReturnType> Op;
1619 typedef TensorSycl::internal::SecondStepFullReducer<CoeffReturnType, Op, EvaluatorPointerType,
1620 EvaluatorPointerType, StorageIndex, local_range>
1623 .template unary_kernel_launcher<CoeffReturnType, GenericRKernel>(
1624 tmp_global_accessor, buffer,
1625 cl::sycl::nd_range<1>(cl::sycl::range<1>(local_range), cl::sycl::range<1>(local_range)), local_range,
1628 device().deallocate_temp(temp_pointer);
1631 .template binary_kernel_launcher<CoeffReturnType, ContractKernelName>(lhs, rhs, buffer, thread_range,
1638 EIGEN_STRONG_INLINE
void cleanup() {
1639 this->m_leftImpl.cleanup();
1640 this->m_rightImpl.cleanup();
1642 if (this->m_result) {
1643 this->m_device.deallocate_temp(this->m_result);
1644 this->m_result =
nullptr;
Definition TensorContraction.h:335
TensorContractionKernel is a template class that provides Tensor -Tensor contraction operation.
Definition TensorContractionSycl.h:457
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47
BlockProperties is a template class that provides different characteristic of a block of each Tensor ...
Definition TensorContractionSycl.h:322
TTPanelSize, a template class used for setting the panel size required for launching General Tensor T...
Definition TensorContractionSycl.h:84
TVPanelSize, a template class used for setting the panel size required for launching General TensorVe...
Definition TensorContractionSycl.h:48
MemHolder this is a place holder struct for creating memory hierarchy in SYCL. Inside SYCL kernel it ...
Definition TensorContractionSycl.h:504
TiledMemory: contains required memory pointer for loading each tile of the TensorContraction panel fr...
Definition TensorContractionSycl.h:537
ThreadProperties is a template class that provides each thread's properties within a workgroup....
Definition TensorContractionSycl.h:373