已合并
update TColReduceIdx to support synchronous output of value and index #928
سقط 落创建于 5月15日
update TColReduceIdx to support synchronous output of value and index #928
已合并
共 9 个文件变更+2051-431
| @@ -1201,6 +1201,26 @@ PTO_INST RecordEvent TCOLARGMIN(TileDataOut &dst, TileDataIn &src, TileDataTmp & | |||
| 1201 | return {}; | 1201 | return {}; |
| 1202 | } | 1202 | } |
| 1203 | 1203 | ||
| 1204 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp, | ||
| 1205 | + typename... WaitEvents, std::enable_if_t<is_tile_data_v<TileDataTmp> && all_events_v<WaitEvents...>, int> = 0> | ||
| 1206 | +PTO_INST RecordEvent TCOLARGMAX(TileDataOutVal &dstVal, TileDataOutIdx &dstIdx, TileDataIn &src, TileDataTmp &tmp, | ||
| 1207 | + WaitEvents &...events) | ||
| 1208 | +{ | ||
| 1209 | + TSYNC(events...); | ||
| 1210 | + MAP_INSTR_IMPL(TCOLARGMAX, dstVal, dstIdx, src, tmp); | ||
| 1211 | + return {}; | ||
| 1212 | +} | ||
| 1213 | + | ||
| 1214 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp, | ||
| 1215 | + typename... WaitEvents, std::enable_if_t<is_tile_data_v<TileDataTmp> && all_events_v<WaitEvents...>, int> = 0> | ||
| 1216 | +PTO_INST RecordEvent TCOLARGMIN(TileDataOutVal &dstVal, TileDataOutIdx &dstIdx, TileDataIn &src, TileDataTmp &tmp, | ||
| 1217 | + WaitEvents &...events) | ||
| 1218 | +{ | ||
| 1219 | + TSYNC(events...); | ||
| 1220 | + MAP_INSTR_IMPL(TCOLARGMIN, dstVal, dstIdx, src, tmp); | ||
| 1221 | + return {}; | ||
| 1222 | +} | ||
| 1223 | + | ||
| 1204 | template <typename TileDataOut, typename TileDataIn, typename TileDataTmp, typename... WaitEvents> | 1224 | template <typename TileDataOut, typename TileDataIn, typename TileDataTmp, typename... WaitEvents> |
| 1205 | PTO_INST RecordEvent TROWMAX(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp, WaitEvents &...events) | 1225 | PTO_INST RecordEvent TROWMAX(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp, WaitEvents &...events) |
| 1206 | { | 1226 | { |
| @@ -8,242 +8,314 @@ INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A | |||
| 8 | See LICENSE in the root of the software repository for the full text of the License. | 8 | See LICENSE in the root of the software repository for the full text of the License. |
| 9 | */ | 9 | */ |
| 10 | 10 | ||
| 11 | -#ifndef TCOLREDUCEIDX_HPP | 11 | +#ifndef T_COL_REDUCE_IDX_OPS_HPP |
| 12 | -#define TCOLREDUCEIDX_HPP | 12 | +#define T_COL_REDUCE_IDX_OPS_HPP |
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | namespace pto { | 17 | namespace pto { |
| 18 | -template <typename TileDataOut, typename TileDataIn, typename TileDataTmp> | 18 | + |
| 19 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp, | ||
| 20 | + bool WithVal = false> | ||
| 19 | PTO_INTERNAL void TColReduceIdxCheck(unsigned srcValidRow, unsigned srcValidCol, unsigned dstValidRow, | 21 | PTO_INTERNAL void TColReduceIdxCheck(unsigned srcValidRow, unsigned srcValidCol, unsigned dstValidRow, |
| 20 | - unsigned dstValidCol) | 22 | + unsigned dstValidCol, unsigned dstValValidRow = 0, unsigned dstValValidCol = 0) |
| 21 | { | 23 | { |
| 24 | + // 输入数据类型检查 | ||
| 22 | static_assert( | 25 | static_assert( |
| 23 | std::is_same_v<typename TileDataIn::DType, uint32_t> || std::is_same_v<typename TileDataIn::DType, uint16_t> || | 26 | std::is_same_v<typename TileDataIn::DType, uint32_t> || std::is_same_v<typename TileDataIn::DType, uint16_t> || |
| 24 | std::is_same_v<typename TileDataIn::DType, half> || std::is_same_v<typename TileDataIn::DType, float>, | 27 | std::is_same_v<typename TileDataIn::DType, half> || std::is_same_v<typename TileDataIn::DType, float>, |
| 25 | - "Fix: TCOLARGMAX input data type must be f16/u16/f32/u32"); | 28 | + "Fix: TColReduceIdx input data type must be f16/u16/f32/u32"); |
| 26 | - static_assert(TileDataIn::Loc == pto::TileType::Vec, "Fix: TCOLARGMAX Src TileType must be Vec Tile!"); | 29 | + |
| 27 | - static_assert(TileDataOut::Loc == pto::TileType::Vec, "Fix: TCOLARGMAX Dst TileType must be Vec Tile!"); | 30 | + // 输入Tile类型检查 |
| 28 | - static_assert(TileDataIn::SFractal == SLayout::NoneBox, "Fix: TCOLARGMAX only support Nd or Dn fractal Tile"); | 31 | + static_assert(TileDataIn::Loc == pto::TileType::Vec, "Fix: TColReduceIdx Src TileType must be Vec Tile"); |
| 29 | - static_assert(TileDataOut::isRowMajor && TileDataOut::SFractal == SLayout::NoneBox, | 32 | + static_assert(TileDataIn::SFractal == SLayout::NoneBox, "Fix: TColReduceIdx only support Nd or Dn fractal Tile"); |
| 30 | - "Fix: TCOLARGMAX only support Nd fractal Tile"); | 33 | + |
| 31 | - static_assert( | 34 | + // 输出索引Tile类型检查 |
| 32 | - std::is_same_v<typename TileDataOut::DType, uint32_t> || std::is_same_v<typename TileDataOut::DType, int32_t>, | 35 | + static_assert(TileDataOutIdx::Loc == pto::TileType::Vec, "Fix: TColReduceIdx DstIdx TileType must be Vec Tile"); |
| 33 | - "Fix: TCOLARGMAX output data type must be s32 or u32."); | 36 | + static_assert(TileDataOutIdx::isRowMajor && TileDataOutIdx::SFractal == SLayout::NoneBox, |
| 37 | + "Fix: TColReduceIdx DstIdx only supports Nd fractal Tile"); | ||
| 38 | + | ||
| 39 | + // 临时Tile类型检查 | ||
| 34 | static_assert(std::is_same_v<typename TileDataIn::DType, typename TileDataTmp::DType>, | 40 | static_assert(std::is_same_v<typename TileDataIn::DType, typename TileDataTmp::DType>, |
| 35 | - "Fix: TCOLARGMAX input type must be consistent with the tmp type"); | 41 | + "Fix: TColReduceIdx input type must be consistent with tmp type"); |
| 42 | + | ||
| 43 | + // 基础输入维度检查 | ||
| 36 | PTO_ASSERT(srcValidRow != 0 && srcValidCol != 0, | 44 | PTO_ASSERT(srcValidRow != 0 && srcValidCol != 0, |
| 37 | - "Fix: TCOLARGMAX input shape is invalid, validCol or validRow is 0."); | 45 | + "Fix: TColReduceIdx input shape is invalid, validCol or validRow is 0"); |
| 38 | - PTO_ASSERT(dstValidRow == 1, "Fix: TCOLARGMAX output validRow must be 1"); | 46 | + PTO_ASSERT(dstValidRow == 1, "Fix: TColReduceIdx output idx validRow must be 1"); |
| 39 | - PTO_ASSERT(srcValidCol == dstValidCol, | 47 | + PTO_ASSERT(srcValidCol == dstValidCol, "Fix: TColReduceIdx input validCol must equal idx output validCol"); |
| 40 | - "Fix: TCOLARGMAX input validCol must be consistent with the output validCol"); | 48 | + |
| 49 | + if constexpr (WithVal) { | ||
| 50 | + // 值输出Tile类型检查 | ||
| 51 | + static_assert(TileDataOutVal::Loc == pto::TileType::Vec, | ||
| 52 | + "Fix: TColReduceOpsIdx DstVal TileType must be Vec Tile"); | ||
| 53 | + static_assert(TileDataOutVal::isRowMajor && TileDataOutVal::SFractal == SLayout::NoneBox, | ||
| 54 | + "Fix: TColReduceOpsIdx DstVal only supports Nd fractal Tile"); | ||
| 55 | + static_assert(std::is_same_v<typename TileDataOutVal::DType, typename TileDataIn::DType>, | ||
| 56 | + "Fix: TColReduceOpsIdx DstVal data type must match input type"); | ||
| 57 | + | ||
| 58 | + // 值输出维度检查 | ||
| 59 | + PTO_ASSERT(dstValValidRow == 1, "Fix: TColReduceOpsIdx output value validRow must be 1"); | ||
| 60 | + PTO_ASSERT(dstValValidCol != 0, "Fix: TColReduceOpsIdx output value validCol must be non-zero"); | ||
| 61 | + PTO_ASSERT(srcValidCol == dstValValidCol, | ||
| 62 | + "Fix: TColReduceOpsIdx input validCol must equal value output validCol"); | ||
| 63 | + PTO_ASSERT(dstValValidRow == dstValidRow, | ||
| 64 | + "Fix: TColReduceOpsIdx value and idx output tiles must have same validRow"); | ||
| 65 | + PTO_ASSERT(dstValValidCol == dstValidCol, | ||
| 66 | + "Fix: TColReduceOpsIdx value and idx output tiles must have same validCol"); | ||
| 67 | + | ||
| 68 | + if constexpr (sizeof(typename TileDataIn::DType) == 2) { | ||
| 69 | + static_assert(std::is_same_v<typename TileDataOutIdx::DType, uint16_t> || | ||
| 70 | + std::is_same_v<typename TileDataOutIdx::DType, int16_t>, | ||
| 71 | + "Fix: TColReduceOpsIdx DstIdx data type must be s16 or u16 when input type size is 2 bytes"); | ||
| 72 | + } else { | ||
| 73 | + static_assert(std::is_same_v<typename TileDataOutIdx::DType, uint32_t> || | ||
| 74 | + std::is_same_v<typename TileDataOutIdx::DType, int32_t>, | ||
| 75 | + "Fix: TColReduceOpsIdx DstIdx data type must be s32 or u32 when input type size is 4 bytes"); | ||
| 76 | + } | ||
| 77 | + } else { | ||
| 78 | + static_assert(std::is_same_v<typename TileDataOutIdx::DType, uint32_t> || | ||
| 79 | + std::is_same_v<typename TileDataOutIdx::DType, int32_t>, | ||
| 80 | + "Fix: TColReduceIdx DstIdx data type must be s32 or u32"); | ||
| 81 | + } | ||
| 41 | } | 82 | } |
| 42 | 83 | ||
| 43 | -template <typename TileDataOut, typename TileDataIn, typename TileDataTmp, bool IsArgMax> | 84 | +template <typename TVal, bool IsArgMax> |
| 44 | -__tf__ PTO_INTERNAL void TColReduceIdx16(typename TileDataOut::TileDType __out__ dst, | 85 | +struct TColIdxCompareOp { |
| 45 | - typename TileDataIn::TileDType __in__ src, | 86 | + PTO_INTERNAL static void VCmp(__ubuf__ TVal *src1, __ubuf__ TVal *src2) |
| 46 | - typename TileDataTmp::TileDType __in__ tmp, unsigned srcValidRow, | 87 | + { |
| 47 | - unsigned srcValidCol) | 88 | + if constexpr (IsArgMax) { |
| 48 | -{ | 89 | + vcmp_ge(src1, src2, 1, 1, 1, 1, 0, 0, 0); |
| 49 | - using TOUT = typename TileDataOut::DType; | 90 | + } else { |
| 50 | - using T = typename TileDataIn::DType; | 91 | + vcmp_le(src1, src2, 1, 1, 1, 1, 0, 0, 0); |
| 51 | - constexpr uint32_t srcRowStride = TileDataIn::Cols; | ||
| 52 | - constexpr uint32_t elemPerRpt = REPEAT_BYTE / sizeof(T); | ||
| 53 | - constexpr uint32_t elemPerBlock = BLOCK_BYTE_SIZE / sizeof(T); | ||
| 54 | - __ubuf__ TOUT *dstPtr = (__ubuf__ TOUT *)__cce_get_tile_ptr(dst); | ||
| 55 | - __ubuf__ T *srcPtr = (__ubuf__ T *)__cce_get_tile_ptr(src); | ||
| 56 | - __ubuf__ T *tmpPtr = (__ubuf__ T *)__cce_get_tile_ptr(tmp); | ||
| 57 | - | ||
| 58 | - uint16_t numLoop = srcValidCol / elemPerRpt; | ||
| 59 | - uint16_t remainAfterLoop = srcValidCol % elemPerRpt; | ||
| 60 | - uint32_t tmpGapEles = numLoop > 0 ? elemPerRpt : CeilDivision(srcValidCol, elemPerBlock) * elemPerBlock; | ||
| 61 | - | ||
| 62 | - for (uint16_t j = 0; j < numLoop; j++) { | ||
| 63 | - pipe_barrier(PIPE_V); | ||
| 64 | - vector_dup((__ubuf__ int16_t *)tmpPtr, 0, 1, 1, 1, 0, 0); // cur index | ||
| 65 | - vector_dup((__ubuf__ int16_t *)tmpPtr + 2 * tmpGapEles, 0, 1, 1, 1, 0, 0); // argmin index | ||
| 66 | - pto_copy_ubuf_to_ubuf(tmpPtr + tmpGapEles, srcPtr + j * elemPerRpt, 1, 8, 0, 0); // min elements | ||
| 67 | - pipe_barrier(PIPE_V); | ||
| 68 | - for (uint16_t i = 1; i < srcValidRow; i++) { | ||
| 69 | - vadds((__ubuf__ int16_t *)tmpPtr, (__ubuf__ int16_t *)tmpPtr, 1, 1, 1, 1, 0, 0); | ||
| 70 | - if constexpr (IsArgMax) { | ||
| 71 | - vcmp_ge((__ubuf__ half *)tmpPtr + tmpGapEles, | ||
| 72 | - (__ubuf__ half *)srcPtr + i * srcRowStride + j * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 73 | - pipe_barrier(PIPE_V); | ||
| 74 | - vsel((__ubuf__ half *)tmpPtr + 2 * tmpGapEles, (__ubuf__ half *)tmpPtr + 2 * tmpGapEles, | ||
| 75 | - (__ubuf__ half *)tmpPtr, 1, 1, 1, 1, 0, 0, 0, 0); | ||
| 76 | - vmax((__ubuf__ half *)tmpPtr + tmpGapEles, (__ubuf__ half *)tmpPtr + tmpGapEles, | ||
| 77 | - (__ubuf__ half *)srcPtr + i * srcRowStride + j * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 78 | - } else { | ||
| 79 | - vcmp_le((__ubuf__ half *)tmpPtr + tmpGapEles, | ||
| 80 | - (__ubuf__ half *)srcPtr + i * srcRowStride + j * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 81 | - pipe_barrier(PIPE_V); | ||
| 82 | - vsel((__ubuf__ half *)tmpPtr + 2 * tmpGapEles, (__ubuf__ half *)tmpPtr + 2 * tmpGapEles, | ||
| 83 | - (__ubuf__ half *)tmpPtr, 1, 1, 1, 1, 0, 0, 0, 0); | ||
| 84 | - vmin((__ubuf__ half *)tmpPtr + tmpGapEles, (__ubuf__ half *)tmpPtr + tmpGapEles, | ||
| 85 | - (__ubuf__ half *)srcPtr + i * srcRowStride + j * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 86 | - } | ||
| 87 | - pipe_barrier(PIPE_V); | ||
| 88 | } | 92 | } |
| 93 | + } | ||
| 89 | 94 | ||
| 90 | - vconv_s162f16a((__ubuf__ half *)tmpPtr + 2 * tmpGapEles, (__ubuf__ int16_t *)tmpPtr + 2 * tmpGapEles, 1, 1, 1, | 95 | + PTO_INTERNAL static void UpdateValue(__ubuf__ TVal *dst, __ubuf__ TVal *src1, __ubuf__ TVal *src2) |
| 91 | - 0, 0); | 96 | + { |
| 97 | + if constexpr (IsArgMax) { | ||
| 98 | + vmax(dst, src1, src2, 1, 1, 1, 1, 0, 0, 0); | ||
| 99 | + } else { | ||
| 100 | + vmin(dst, src1, src2, 1, 1, 1, 1, 0, 0, 0); | ||
| 101 | + } | ||
| 102 | + } | ||
| 103 | +}; | ||
| 104 | + | ||
| 105 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp, bool IsArgMax, | ||
| 106 | + bool IsRemainingBlock, bool WithVal> | ||
| 107 | +PTO_INTERNAL void ProcessSingleBlock(__ubuf__ typename TileDataOutVal::DType *dstVal, | ||
| 108 | + __ubuf__ typename TileDataOutIdx::DType *dstIdx, | ||
| 109 | + __ubuf__ typename TileDataIn::DType *src, | ||
| 110 | + __ubuf__ typename TileDataTmp::DType *tmp, unsigned srcValidRow, | ||
| 111 | + uint16_t blockIndex, uint32_t tmpGapElems, uint32_t elemPerRpt, | ||
| 112 | + uint16_t remainingElements = 0) | ||
| 113 | +{ | ||
| 114 | + using TIN = typename TileDataIn::DType; | ||
| 115 | + using TOUT = typename TileDataOutIdx::DType; | ||
| 116 | + constexpr uint32_t srcRowStride = TileDataIn::Cols; | ||
| 117 | + using TIdx = std::conditional_t<sizeof(TIN) == 2, int16_t, int32_t>; | ||
| 118 | + using TCmp = std::conditional_t<sizeof(TIN) == 2, half, float>; | ||
| 119 | + | ||
| 120 | + constexpr bool isHalfType = sizeof(TIN) == 2; | ||
| 121 | + | ||
| 122 | + __ubuf__ TIN *elemTmpPtr = tmp + tmpGapElems; | ||
| 123 | + __ubuf__ TOUT *currentDst; | ||
| 124 | + __ubuf__ TIdx *currentIndex; | ||
| 125 | + __ubuf__ TCmp *selectedCmpVals; | ||
| 126 | + | ||
| 127 | + if constexpr (isHalfType && !WithVal) { | ||
| 128 | + currentIndex = reinterpret_cast<__ubuf__ TIdx *>(tmp) + 2 * tmpGapElems; | ||
| 129 | + selectedCmpVals = reinterpret_cast<__ubuf__ TCmp *>(tmp) + 2 * tmpGapElems; | ||
| 130 | + currentDst = reinterpret_cast<__ubuf__ TOUT *>(tmp) + 2 * tmpGapElems; | ||
| 131 | + } else { | ||
| 132 | + currentDst = dstIdx + blockIndex * elemPerRpt; | ||
| 133 | + currentIndex = reinterpret_cast<__ubuf__ TIdx *>(dstIdx) + blockIndex * elemPerRpt; | ||
| 134 | + selectedCmpVals = reinterpret_cast<__ubuf__ TCmp *>(dstIdx) + blockIndex * elemPerRpt; | ||
| 135 | + } | ||
| 136 | + | ||
| 137 | + if constexpr (IsRemainingBlock) { | ||
| 138 | + set_mask_count(); | ||
| 139 | + set_vector_mask(0, remainingElements); | ||
| 140 | + } | ||
| 141 | + | ||
| 142 | + vector_dup(tmp, 0, 1, 1, 1, 0, 0); // 累积索引 | ||
| 143 | + vector_dup(currentIndex, 0, 1, 1, 1, 0, 0); // 当前索引 | ||
| 144 | + | ||
| 145 | + if constexpr (IsRemainingBlock) { | ||
| 146 | + vcopy(reinterpret_cast<__ubuf__ TIdx *>(elemTmpPtr), | ||
| 147 | + reinterpret_cast<__ubuf__ TIdx *>(src + blockIndex * elemPerRpt), 1, 1, 1, 0, 0); | ||
| 148 | + } else { | ||
| 149 | + pto_copy_ubuf_to_ubuf(elemTmpPtr, src + blockIndex * elemPerRpt, 1, 8, 0, 0); | ||
| 150 | + } | ||
| 151 | + pipe_barrier(PIPE_V); | ||
| 152 | + | ||
| 153 | + for (uint16_t rowIdx = 1; rowIdx < srcValidRow; ++rowIdx) { | ||
| 154 | + vadds(reinterpret_cast<__ubuf__ TIdx *>(tmp), reinterpret_cast<__ubuf__ TIdx *>(tmp), 1, 1, 1, 1, 0, 0); | ||
| 155 | + | ||
| 156 | + __ubuf__ TIN *currentSrc = src + rowIdx * srcRowStride + blockIndex * elemPerRpt; | ||
| 157 | + | ||
| 158 | + TColIdxCompareOp<TCmp, IsArgMax>::VCmp(reinterpret_cast<__ubuf__ TCmp *>(elemTmpPtr), | ||
| 159 | + reinterpret_cast<__ubuf__ TCmp *>(currentSrc)); | ||
| 92 | pipe_barrier(PIPE_V); | 160 | pipe_barrier(PIPE_V); |
| 93 | - vconv_f162s32a((__ubuf__ int32_t *)dstPtr + j * elemPerRpt, (__ubuf__ half *)tmpPtr + 2 * tmpGapEles, 2, 1, 1, | 161 | + |
| 94 | - 8, 4); | 162 | + vsel(reinterpret_cast<__ubuf__ TCmp *>(selectedCmpVals), reinterpret_cast<__ubuf__ TCmp *>(selectedCmpVals), |
| 163 | + reinterpret_cast<__ubuf__ TCmp *>(tmp), 1, 1, 1, 1, 0, 0, 0, 0); | ||
| 164 | + | ||
| 165 | + TColIdxCompareOp<TCmp, IsArgMax>::UpdateValue(reinterpret_cast<__ubuf__ TCmp *>(elemTmpPtr), | ||
| 166 | + reinterpret_cast<__ubuf__ TCmp *>(elemTmpPtr), | ||
| 167 | + reinterpret_cast<__ubuf__ TCmp *>(currentSrc)); | ||
| 95 | pipe_barrier(PIPE_V); | 168 | pipe_barrier(PIPE_V); |
| 96 | } | 169 | } |
| 97 | - if (remainAfterLoop > 0) { | 170 | + |
| 98 | - set_mask_count(); | 171 | + if constexpr (isHalfType && !WithVal) { |
| 99 | - set_vector_mask(0, remainAfterLoop); | 172 | + vconv_s162f16a(reinterpret_cast<__ubuf__ half *>(tmp) + 2 * tmpGapElems, |
| 100 | - vector_dup(tmpPtr, 0, 1, 1, 1, 0, 0); | 173 | + reinterpret_cast<__ubuf__ int16_t *>(tmp) + 2 * tmpGapElems, 1, 1, 1, 0, 0); |
| 101 | - vector_dup((__ubuf__ int16_t *)tmpPtr + 2 * tmpGapEles, 0, 1, 1, 1, 0, 0); | ||
| 102 | - vcopy((__ubuf__ int16_t *)tmpPtr + tmpGapEles, (__ubuf__ int16_t *)srcPtr + numLoop * elemPerRpt, 1, 1, 1, 0, | ||
| 103 | - 0); | ||
| 104 | pipe_barrier(PIPE_V); | 174 | pipe_barrier(PIPE_V); |
| 105 | 175 | ||
| 106 | - for (uint16_t i = 1; i < srcValidRow; i++) { | 176 | + vconv_f162s32a(reinterpret_cast<__ubuf__ int32_t *>(dstIdx) + blockIndex * elemPerRpt, |
| 107 | - vadds((__ubuf__ int16_t *)tmpPtr, (__ubuf__ int16_t *)tmpPtr, 1, 1, 1, 1, 0, 0); | 177 | + reinterpret_cast<__ubuf__ half *>(tmp) + 2 * tmpGapElems, 2, 1, 1, 8, 4); |
| 108 | - if constexpr (IsArgMax) { | ||
| 109 | - vcmp_ge((__ubuf__ half *)tmpPtr + tmpGapEles, | ||
| 110 | - (__ubuf__ half *)srcPtr + i * srcRowStride + numLoop * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 111 | - pipe_barrier(PIPE_V); | ||
| 112 | - vsel((__ubuf__ half *)tmpPtr + 2 * tmpGapEles, (__ubuf__ half *)tmpPtr + 2 * tmpGapEles, | ||
| 113 | - (__ubuf__ half *)tmpPtr, 1, 1, 1, 1, 0, 0, 0, 0); | ||
| 114 | - vmax((__ubuf__ half *)tmpPtr + tmpGapEles, (__ubuf__ half *)tmpPtr + tmpGapEles, | ||
| 115 | - (__ubuf__ half *)srcPtr + i * srcRowStride + numLoop * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 116 | - } else { | ||
| 117 | - vcmp_le((__ubuf__ half *)tmpPtr + tmpGapEles, | ||
| 118 | - (__ubuf__ half *)srcPtr + i * srcRowStride + numLoop * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 119 | - pipe_barrier(PIPE_V); | ||
| 120 | - vsel((__ubuf__ half *)tmpPtr + 2 * tmpGapEles, (__ubuf__ half *)tmpPtr + 2 * tmpGapEles, | ||
| 121 | - (__ubuf__ half *)tmpPtr, 1, 1, 1, 1, 0, 0, 0, 0); | ||
| 122 | - vmin((__ubuf__ half *)tmpPtr + tmpGapEles, (__ubuf__ half *)tmpPtr + tmpGapEles, | ||
| 123 | - (__ubuf__ half *)srcPtr + i * srcRowStride + numLoop * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 124 | - } | ||
| 125 | - pipe_barrier(PIPE_V); | ||
| 126 | - } | ||
| 127 | - | ||
| 128 | - vconv_s162f16a((__ubuf__ half *)tmpPtr + 2 * tmpGapEles, (__ubuf__ int16_t *)tmpPtr + 2 * tmpGapEles, 1, 1, 1, | ||
| 129 | - 0, 0); | ||
| 130 | pipe_barrier(PIPE_V); | 178 | pipe_barrier(PIPE_V); |
| 131 | - vconv_f162s32a((__ubuf__ int32_t *)dstPtr + numLoop * elemPerRpt, (__ubuf__ half *)tmpPtr + 2 * tmpGapEles, 2, | 179 | + } |
| 132 | - 1, 1, 8, 4); | 180 | + |
| 181 | + if constexpr (WithVal && !IsRemainingBlock) { | ||
| 182 | + pto_copy_ubuf_to_ubuf(dstVal + blockIndex * elemPerRpt, elemTmpPtr, 1, 8, 0, 0); | ||
| 183 | + pipe_barrier(PIPE_V); | ||
| 184 | + } else if constexpr (WithVal && IsRemainingBlock) { | ||
| 185 | + vcopy(reinterpret_cast<__ubuf__ TIdx *>(dstVal) + blockIndex * elemPerRpt, | ||
| 186 | + reinterpret_cast<__ubuf__ TIdx *>(elemTmpPtr), 1, 1, 1, 0, 0); | ||
| 187 | + pipe_barrier(PIPE_V); | ||
| 188 | + } | ||
| 189 | + | ||
| 190 | + if constexpr (IsRemainingBlock) { | ||
| 133 | set_mask_norm(); | 191 | set_mask_norm(); |
| 134 | set_vector_mask(-1, -1); | 192 | set_vector_mask(-1, -1); |
| 135 | pipe_barrier(PIPE_V); | 193 | pipe_barrier(PIPE_V); |
| 136 | } | 194 | } |
| 137 | } | 195 | } |
| 138 | 196 | ||
| 139 | -template <typename TileDataOut, typename TileDataIn, typename TileDataTmp, bool IsArgMax> | 197 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp, bool IsArgMax, |
| 140 | -__tf__ PTO_INTERNAL void TColReduceIdx32(typename TileDataOut::TileDType __out__ dst, | 198 | + bool WithVal> |
| 141 | - typename TileDataIn::TileDType __in__ src, | 199 | +PTO_INTERNAL void ProcessFullBlocks(__ubuf__ typename TileDataOutVal::DType *dstVal, |
| 142 | - typename TileDataTmp::TileDType __in__ tmp, unsigned srcValidRow, | 200 | + __ubuf__ typename TileDataOutIdx::DType *dstIdx, |
| 143 | - unsigned srcValidCol) | 201 | + __ubuf__ typename TileDataIn::DType *src, __ubuf__ typename TileDataTmp::DType *tmp, |
| 202 | + unsigned srcValidRow, unsigned srcValidCol, uint16_t numFullBlocks, | ||
| 203 | + uint32_t tmpGapElems, uint32_t elemPerRpt) | ||
| 144 | { | 204 | { |
| 145 | - using TOUT = typename TileDataOut::DType; | 205 | + if (numFullBlocks == 0) |
| 146 | - using T = typename TileDataIn::DType; | 206 | + return; |
| 147 | - __ubuf__ TOUT *dstPtr = (__ubuf__ TOUT *)__cce_get_tile_ptr(dst); | 207 | + for (uint16_t blockIdx = 0; blockIdx < numFullBlocks; ++blockIdx) { |
| 148 | - __ubuf__ T *srcPtr = (__ubuf__ T *)__cce_get_tile_ptr(src); | 208 | + ProcessSingleBlock<TileDataOutVal, TileDataOutIdx, TileDataIn, TileDataTmp, IsArgMax, false, WithVal>( |
| 149 | - __ubuf__ T *tmpPtr = (__ubuf__ T *)__cce_get_tile_ptr(tmp); | 209 | + dstVal, dstIdx, src, tmp, srcValidRow, blockIdx, tmpGapElems, elemPerRpt); |
| 150 | - | ||
| 151 | - constexpr uint32_t srcRowStride = TileDataIn::Cols; | ||
| 152 | - constexpr uint32_t elemPerRpt = REPEAT_BYTE / sizeof(T); | ||
| 153 | - constexpr uint32_t elemPerBlock = BLOCK_BYTE_SIZE / sizeof(T); | ||
| 154 | - uint16_t numLoop = srcValidCol / elemPerRpt; | ||
| 155 | - uint16_t remainAfterLoop = srcValidCol % elemPerRpt; | ||
| 156 | - uint32_t tmpGapEles = numLoop > 0 ? elemPerRpt : CeilDivision(srcValidCol, elemPerBlock) * elemPerBlock; | ||
| 157 | - | ||
| 158 | - for (uint16_t j = 0; j < numLoop; j++) { | ||
| 159 | - vector_dup(dstPtr + j * elemPerRpt, 0, 1, 1, 1, 0, 0); // argmin index | ||
| 160 | - vector_dup(tmpPtr, 0, 1, 1, 1, 0, 0); // cur index | ||
| 161 | - pto_copy_ubuf_to_ubuf(tmpPtr + tmpGapEles, srcPtr + j * elemPerRpt, 1, 8, 0, 0); // min elements | ||
| 162 | - pipe_barrier(PIPE_V); | ||
| 163 | - for (uint16_t i = 1; i < srcValidRow; i++) { | ||
| 164 | - vadds((__ubuf__ int32_t *)tmpPtr, (__ubuf__ int32_t *)tmpPtr, 1, 1, 1, 1, 0, 0); | ||
| 165 | - if constexpr (IsArgMax) { | ||
| 166 | - vcmp_ge((__ubuf__ float *)tmpPtr + tmpGapEles, | ||
| 167 | - (__ubuf__ float *)srcPtr + i * srcRowStride + j * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 168 | - pipe_barrier(PIPE_V); | ||
| 169 | - vsel((__ubuf__ float *)dstPtr + j * elemPerRpt, (__ubuf__ float *)dstPtr + j * elemPerRpt, | ||
| 170 | - (__ubuf__ float *)tmpPtr, 1, 1, 1, 1, 0, 0, 0, 0); | ||
| 171 | - vmax((__ubuf__ float *)tmpPtr + tmpGapEles, (__ubuf__ float *)tmpPtr + tmpGapEles, | ||
| 172 | - (__ubuf__ float *)srcPtr + i * srcRowStride + j * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 173 | - } else { | ||
| 174 | - vcmp_le((__ubuf__ float *)tmpPtr + tmpGapEles, | ||
| 175 | - (__ubuf__ float *)srcPtr + i * srcRowStride + j * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 176 | - pipe_barrier(PIPE_V); | ||
| 177 | - vsel((__ubuf__ float *)dstPtr + j * elemPerRpt, (__ubuf__ float *)dstPtr + j * elemPerRpt, | ||
| 178 | - (__ubuf__ float *)tmpPtr, 1, 1, 1, 1, 0, 0, 0, 0); | ||
| 179 | - vmin((__ubuf__ float *)tmpPtr + tmpGapEles, (__ubuf__ float *)tmpPtr + tmpGapEles, | ||
| 180 | - (__ubuf__ float *)srcPtr + i * srcRowStride + j * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 181 | - } | ||
| 182 | - pipe_barrier(PIPE_V); | ||
| 183 | - } | ||
| 184 | - } | ||
| 185 | - if (remainAfterLoop > 0) { | ||
| 186 | - set_mask_count(); | ||
| 187 | - set_vector_mask(0, remainAfterLoop); | ||
| 188 | - | ||
| 189 | - vector_dup(dstPtr + numLoop * elemPerRpt, 0, 1, 1, 1, 0, 0); | ||
| 190 | - vector_dup(tmpPtr, 0, 1, 1, 1, 0, 0); | ||
| 191 | - vcopy((__ubuf__ int32_t *)tmpPtr + tmpGapEles, (__ubuf__ int32_t *)srcPtr + numLoop * elemPerRpt, 1, 1, 1, 0, | ||
| 192 | - 0); | ||
| 193 | - pipe_barrier(PIPE_V); | ||
| 194 | - | ||
| 195 | - for (uint16_t i = 1; i < srcValidRow; i++) { | ||
| 196 | - vadds((__ubuf__ int32_t *)tmpPtr, (__ubuf__ int32_t *)tmpPtr, 1, 1, 1, 1, 0, 0); | ||
| 197 | - if constexpr (IsArgMax) { | ||
| 198 | - vcmp_ge((__ubuf__ float *)tmpPtr + tmpGapEles, | ||
| 199 | - (__ubuf__ float *)srcPtr + i * srcRowStride + numLoop * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 200 | - pipe_barrier(PIPE_V); | ||
| 201 | - vsel((__ubuf__ float *)dstPtr + numLoop * elemPerRpt, (__ubuf__ float *)dstPtr + numLoop * elemPerRpt, | ||
| 202 | - (__ubuf__ float *)tmpPtr, 1, 1, 1, 1, 0, 0, 0, 0); | ||
| 203 | - vmax((__ubuf__ float *)tmpPtr + tmpGapEles, (__ubuf__ float *)tmpPtr + tmpGapEles, | ||
| 204 | - (__ubuf__ float *)srcPtr + i * srcRowStride + numLoop * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 205 | - } else { | ||
| 206 | - vcmp_le((__ubuf__ float *)tmpPtr + tmpGapEles, | ||
| 207 | - (__ubuf__ float *)srcPtr + i * srcRowStride + numLoop * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 208 | - pipe_barrier(PIPE_V); | ||
| 209 | - vsel((__ubuf__ float *)dstPtr + numLoop * elemPerRpt, (__ubuf__ float *)dstPtr + numLoop * elemPerRpt, | ||
| 210 | - (__ubuf__ float *)tmpPtr, 1, 1, 1, 1, 0, 0, 0, 0); | ||
| 211 | - vmin((__ubuf__ float *)tmpPtr + tmpGapEles, (__ubuf__ float *)tmpPtr + tmpGapEles, | ||
| 212 | - (__ubuf__ float *)srcPtr + i * srcRowStride + numLoop * elemPerRpt, 1, 1, 1, 1, 0, 0, 0); | ||
| 213 | - } | ||
| 214 | - pipe_barrier(PIPE_V); | ||
| 215 | - } | ||
| 216 | - set_mask_norm(); | ||
| 217 | - set_vector_mask(-1, -1); | ||
| 218 | - pipe_barrier(PIPE_V); | ||
| 219 | } | 210 | } |
| 220 | } | 211 | } |
| 221 | 212 | ||
| 222 | -template <typename TileDataOut, typename TileDataIn, typename TileDataTmp, bool IsArgMax> | 213 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp, bool IsArgMax, |
| 223 | -PTO_INTERNAL void TCOLARG_DISPATCH(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp) | 214 | + bool WithVal> |
| 215 | +PTO_INTERNAL void ProcessRemainingElements(__ubuf__ typename TileDataOutVal::DType *dstVal, | ||
| 216 | + __ubuf__ typename TileDataOutIdx::DType *dstIdx, | ||
| 217 | + __ubuf__ typename TileDataIn::DType *src, | ||
| 218 | + __ubuf__ typename TileDataTmp::DType *tmp, unsigned srcValidRow, | ||
| 219 | + unsigned srcValidCol, uint16_t numFullBlocks, uint32_t tmpGapElems, | ||
| 220 | + uint32_t elemPerRpt) | ||
| 221 | +{ | ||
| 222 | + uint16_t remainingAfterLoop = srcValidCol % elemPerRpt; | ||
| 223 | + | ||
| 224 | + if (remainingAfterLoop == 0) | ||
| 225 | + return; | ||
| 226 | + | ||
| 227 | + ProcessSingleBlock<TileDataOutVal, TileDataOutIdx, TileDataIn, TileDataTmp, IsArgMax, true, WithVal>( | ||
| 228 | + dstVal, dstIdx, src, tmp, srcValidRow, numFullBlocks, tmpGapElems, elemPerRpt, remainingAfterLoop); | ||
| 229 | +} | ||
| 230 | + | ||
| 231 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp, bool IsArgMax, | ||
| 232 | + bool WithVal> | ||
| 233 | +PTO_INTERNAL void TColReduceIdxImpl(__ubuf__ typename TileDataOutVal::DType *dstVal, | ||
| 234 | + __ubuf__ typename TileDataOutIdx::DType *dstIdx, | ||
| 235 | + __ubuf__ typename TileDataIn::DType *src, __ubuf__ typename TileDataTmp::DType *tmp, | ||
| 236 | + unsigned srcValidRow, unsigned srcValidCol) | ||
| 237 | +{ | ||
| 238 | + using TIN = typename TileDataIn::DType; | ||
| 239 | + constexpr uint32_t elemPerRpt = REPEAT_BYTE / sizeof(TIN); | ||
| 240 | + constexpr uint32_t elemPerBlock = BLOCK_BYTE_SIZE / sizeof(TIN); | ||
| 241 | + | ||
| 242 | + uint16_t numFullBlocks = srcValidCol / elemPerRpt; | ||
| 243 | + uint32_t tmpGapElems = numFullBlocks > 0 ? elemPerRpt : CeilDivision(srcValidCol, elemPerBlock) * elemPerBlock; | ||
| 244 | + | ||
| 245 | + ProcessFullBlocks<TileDataOutVal, TileDataOutIdx, TileDataIn, TileDataTmp, IsArgMax, WithVal>( | ||
| 246 | + dstVal, dstIdx, src, tmp, srcValidRow, srcValidCol, numFullBlocks, tmpGapElems, elemPerRpt); | ||
| 247 | + | ||
| 248 | + ProcessRemainingElements<TileDataOutVal, TileDataOutIdx, TileDataIn, TileDataTmp, IsArgMax, WithVal>( | ||
| 249 | + dstVal, dstIdx, src, tmp, srcValidRow, srcValidCol, numFullBlocks, tmpGapElems, elemPerRpt); | ||
| 250 | +} | ||
| 251 | + | ||
| 252 | +template <bool IsArgMax, typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp, | ||
| 253 | + bool WithVal> | ||
| 254 | +__tf__ PTO_INTERNAL void TColReduceIdxInstr(typename TileDataOutVal::TileDType __out__ dstValData, | ||
| 255 | + typename TileDataOutIdx::TileDType __out__ dstIdxData, | ||
| 256 | + typename TileDataIn::TileDType __in__ srcData, | ||
| 257 | + typename TileDataTmp::TileDType __in__ tmpData, unsigned srcValidRow, | ||
| 258 | + unsigned srcValidCol, unsigned dstValidRow, unsigned dstValidCol) | ||
| 259 | +{ | ||
| 260 | + __ubuf__ typename TileDataOutVal::DType *dstVal = | ||
| 261 | + reinterpret_cast<__ubuf__ typename TileDataOutVal::DType *>(__cce_get_tile_ptr(dstValData)); | ||
| 262 | + __ubuf__ typename TileDataOutIdx::DType *dstIdx = | ||
| 263 | + reinterpret_cast<__ubuf__ typename TileDataOutIdx::DType *>(__cce_get_tile_ptr(dstIdxData)); | ||
| 264 | + __ubuf__ typename TileDataIn::DType *src = | ||
| 265 | + reinterpret_cast<__ubuf__ typename TileDataIn::DType *>(__cce_get_tile_ptr(srcData)); | ||
| 266 | + __ubuf__ typename TileDataTmp::DType *tmp = | ||
| 267 | + reinterpret_cast<__ubuf__ typename TileDataTmp::DType *>(__cce_get_tile_ptr(tmpData)); | ||
| 268 | + | ||
| 269 | + TColReduceIdxImpl<TileDataOutVal, TileDataOutIdx, TileDataIn, TileDataTmp, IsArgMax, WithVal>( | ||
| 270 | + dstVal, dstIdx, src, tmp, srcValidRow, srcValidCol); | ||
| 271 | +} | ||
| 272 | + | ||
| 273 | +template <bool IsArgMax, typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp, | ||
| 274 | + bool WithVal = false> | ||
| 275 | +PTO_INTERNAL void TColReduceIdxDispatch(TileDataOutVal &dstVal, TileDataOutIdx &dstIdx, TileDataIn &src, | ||
| 276 | + TileDataTmp &tmp) | ||
| 224 | { | 277 | { |
| 225 | unsigned srcValidRow = src.GetValidRow(); | 278 | unsigned srcValidRow = src.GetValidRow(); |
| 226 | unsigned srcValidCol = src.GetValidCol(); | 279 | unsigned srcValidCol = src.GetValidCol(); |
| 227 | - TColReduceIdxCheck<TileDataOut, TileDataIn, TileDataTmp>(srcValidRow, srcValidCol, dst.GetValidRow(), | ||
| 228 | - dst.GetValidCol()); | ||
| 229 | 280 | ||
| 230 | - if (sizeof(typename TileDataIn::DType) == 2) { | 281 | + unsigned dstIdxValidRow = dstIdx.GetValidRow(); |
| 231 | - TColReduceIdx16<TileDataOut, TileDataIn, TileDataTmp, IsArgMax>(dst.data(), src.data(), tmp.data(), srcValidRow, | 282 | + unsigned dstIdxValidCol = dstIdx.GetValidCol(); |
| 232 | - srcValidCol); | 283 | + |
| 233 | - } else if (sizeof(typename TileDataIn::DType) == 4) { | 284 | + unsigned dstValValidRow = dstVal.GetValidRow(); |
| 234 | - TColReduceIdx32<TileDataOut, TileDataIn, TileDataTmp, IsArgMax>(dst.data(), src.data(), tmp.data(), srcValidRow, | 285 | + unsigned dstValValidCol = dstVal.GetValidCol(); |
| 235 | - srcValidCol); | 286 | + |
| 236 | - } | 287 | + TColReduceIdxCheck<TileDataOutVal, TileDataOutIdx, TileDataIn, TileDataTmp, WithVal>( |
| 288 | + srcValidRow, srcValidCol, dstIdxValidRow, dstIdxValidCol, dstValValidRow, dstValValidCol); | ||
| 289 | + | ||
| 290 | + TColReduceIdxInstr<IsArgMax, TileDataOutVal, TileDataOutIdx, TileDataIn, TileDataTmp, WithVal>( | ||
| 291 | + dstVal.data(), dstIdx.data(), src.data(), tmp.data(), srcValidRow, srcValidCol, dstIdxValidRow, dstIdxValidCol); | ||
| 237 | } | 292 | } |
| 238 | -template <typename TileDataOut, typename TileDataIn, typename TileDataTmp> | 293 | + |
| 239 | -PTO_INTERNAL void TCOLARGMIN_IMPL(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp) | 294 | +template <typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp> |
| 295 | +PTO_INTERNAL void TCOLARGMAX_IMPL(TileDataOutIdx &dst, TileDataIn &src, TileDataTmp &tmp) | ||
| 240 | { | 296 | { |
| 241 | - TCOLARG_DISPATCH<TileDataOut, TileDataIn, TileDataTmp, false>(dst, src, tmp); // Min | 297 | + TColReduceIdxDispatch<true, TileDataIn, TileDataOutIdx, TileDataIn, TileDataTmp>(src, dst, src, tmp); |
| 242 | } | 298 | } |
| 243 | -template <typename TileDataOut, typename TileDataIn, typename TileDataTmp> | 299 | + |
| 244 | -PTO_INTERNAL void TCOLARGMAX_IMPL(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp) | 300 | +template <typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp> |
| 301 | +PTO_INTERNAL void TCOLARGMIN_IMPL(TileDataOutIdx &dst, TileDataIn &src, TileDataTmp &tmp) | ||
| 245 | { | 302 | { |
| 246 | - TCOLARG_DISPATCH<TileDataOut, TileDataIn, TileDataTmp, true>(dst, src, tmp); // Max | 303 | + TColReduceIdxDispatch<false, TileDataIn, TileDataOutIdx, TileDataIn, TileDataTmp>(src, dst, src, tmp); |
| 247 | } | 304 | } |
| 305 | + | ||
| 306 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp> | ||
| 307 | +PTO_INTERNAL void TCOLARGMAX_IMPL(TileDataOutVal &dstVal, TileDataOutIdx &dstIdx, TileDataIn &src, TileDataTmp &tmp) | ||
| 308 | +{ | ||
| 309 | + TColReduceIdxDispatch<true, TileDataOutVal, TileDataOutIdx, TileDataIn, TileDataTmp, true>(dstVal, dstIdx, src, | ||
| 310 | + tmp); | ||
| 311 | +} | ||
| 312 | + | ||
| 313 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp> | ||
| 314 | +PTO_INTERNAL void TCOLARGMIN_IMPL(TileDataOutVal &dstVal, TileDataOutIdx &dstIdx, TileDataIn &src, TileDataTmp &tmp) | ||
| 315 | +{ | ||
| 316 | + TColReduceIdxDispatch<false, TileDataOutVal, TileDataOutIdx, TileDataIn, TileDataTmp, true>(dstVal, dstIdx, src, | ||
| 317 | + tmp); | ||
| 318 | +} | ||
| 319 | + | ||
| 248 | } // namespace pto | 320 | } // namespace pto |
| 249 | 321 | ||
| @@ -17,248 +17,313 @@ See LICENSE in the root of the software repository for the full text of the Lice | |||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | namespace pto { | 19 | namespace pto { |
| 20 | -template <typename TileDataOut, typename TileDataIn> | 20 | + |
| 21 | +// ---------------------------------------------------------------------------- | ||
| 22 | +// Helper: comparison + select for one pair of old/new value and index registers | ||
| 23 | +// ---------------------------------------------------------------------------- | ||
| 24 | +template <typename VregOldT, typename VregNewT, typename VregIdxT, bool IsArgMax> | ||
| 25 | +PTO_INTERNAL void TColReduceIdxCompare(VregOldT &vregOld, VregNewT &vregNew, VregIdxT &vregIdxOld, VregIdxT &vregIdxNew, | ||
| 26 | + MaskReg &pregSelect, MaskReg &pregMask) | ||
| 27 | +{ | ||
| 28 | + if constexpr (IsArgMax) { | ||
| 29 | + vcmp_gt(pregSelect, vregNew, vregOld, pregMask); | ||
| 30 | + } else { | ||
| 31 | + vcmp_lt(pregSelect, vregNew, vregOld, pregMask); | ||
| 32 | + } | ||
| 33 | + vsel(vregIdxOld, vregIdxNew, vregIdxOld, pregSelect); | ||
| 34 | + vsel(vregOld, vregNew, vregOld, pregSelect); | ||
| 35 | +} | ||
| 36 | + | ||
| 37 | +// ---------------------------------------------------------------------------- | ||
| 38 | +// Helper: convert 16-bit index vector to 32-bit and store with interleave | ||
| 39 | +// Used by both 8-bit (calls twice for even/odd halves) and 16-bit (calls once) | ||
| 40 | +// ---------------------------------------------------------------------------- | ||
| 41 | +template <typename VregSrcT, typename TOUT> | ||
| 42 | +PTO_INTERNAL void TColReduceIdxStoreParts(VregSrcT &vregSrc, __ubuf__ TOUT *dstPtr, unsigned offset, MaskReg &pregPart0, | ||
| 43 | + MaskReg &pregPart1, MaskReg &pregAll) | ||
| 44 | +{ | ||
| 45 | + RegTensor<TOUT> vregOutEven; | ||
| 46 | + RegTensor<TOUT> vregOutOdd; | ||
| 47 | + RegTensor<TOUT> vregOut0; | ||
| 48 | + RegTensor<TOUT> vregOut1; | ||
| 49 | + vcvt(vregOutEven, vregSrc, pregAll, PART_EVEN); | ||
| 50 | + vcvt(vregOutOdd, vregSrc, pregAll, PART_ODD); | ||
| 51 | + vintlv(vregOut0, vregOut1, vregOutEven, vregOutOdd); | ||
| 52 | + vsts(vregOut0, dstPtr, offset, NORM_B32, pregPart0); | ||
| 53 | + vsts(vregOut1, dstPtr, offset + ELE_CNT_B32, NORM_B32, pregPart1); | ||
| 54 | +} | ||
| 55 | + | ||
| 56 | +// ---------------------------------------------------------------------------- | ||
| 57 | +// Check function (fixed assertion on line dstValidRow == 1) | ||
| 58 | +// ---------------------------------------------------------------------------- | ||
| 59 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, bool WithVal = false> | ||
| 21 | PTO_INTERNAL void TColReduceIdxCheck(unsigned srcValidRow, unsigned srcValidCol, unsigned dstValidRow, | 60 | PTO_INTERNAL void TColReduceIdxCheck(unsigned srcValidRow, unsigned srcValidCol, unsigned dstValidRow, |
| 22 | - unsigned dstValidCol) | 61 | + unsigned dstValidCol, unsigned dstValValidRow = 0, unsigned dstValValidCol = 0) |
| 23 | { | 62 | { |
| 24 | static_assert((sizeof(typename TileDataIn::DType) == 1) || (sizeof(typename TileDataIn::DType) == 2) || | 63 | static_assert((sizeof(typename TileDataIn::DType) == 1) || (sizeof(typename TileDataIn::DType) == 2) || |
| 25 | (sizeof(typename TileDataIn::DType) == 4), | 64 | (sizeof(typename TileDataIn::DType) == 4), |
| 26 | "Fix: TCOLREDUCEIDX data type must be b8/b16/b32"); | 65 | "Fix: TCOLREDUCEIDX data type must be b8/b16/b32"); |
| 66 | + | ||
| 27 | static_assert(TileDataIn::Loc == pto::TileType::Vec, "Fix: TCOLREDUCEIDX Src TileType must be Vec Tile!"); | 67 | static_assert(TileDataIn::Loc == pto::TileType::Vec, "Fix: TCOLREDUCEIDX Src TileType must be Vec Tile!"); |
| 28 | - static_assert(TileDataOut::Loc == pto::TileType::Vec, "Fix: TCOLREDUCEIDX Dst TileType must be Vec Tile!"); | 68 | + |
| 69 | + static_assert(TileDataOutIdx::Loc == pto::TileType::Vec, "Fix: TCOLREDUCEIDX DstIdx TileType must be Vec Tile!"); | ||
| 70 | + | ||
| 29 | static_assert(TileDataIn::SFractal == SLayout::NoneBox, "Fix: TCOLREDUCEIDX only support Nd or Dn fractal Tile"); | 71 | static_assert(TileDataIn::SFractal == SLayout::NoneBox, "Fix: TCOLREDUCEIDX only support Nd or Dn fractal Tile"); |
| 30 | - static_assert(TileDataOut::isRowMajor && TileDataOut::SFractal == SLayout::NoneBox, | 72 | + |
| 31 | - "Fix: TCOLREDUCEIDX only support Nd fractal Tile"); | 73 | + static_assert(TileDataOutIdx::isRowMajor && TileDataOutIdx::SFractal == SLayout::NoneBox, |
| 32 | - static_assert( | 74 | + "Fix: TCOLREDUCEIDX DstIdx only supports Nd fractal Tile"); |
| 33 | - std::is_same_v<typename TileDataOut::DType, uint32_t> || std::is_same_v<typename TileDataOut::DType, int32_t>, | 75 | + |
| 34 | - "Fix: TCOLREDUCEIDX output data type must be s32 or u32."); | ||
| 35 | PTO_ASSERT(srcValidRow != 0 && srcValidCol != 0, | 76 | PTO_ASSERT(srcValidRow != 0 && srcValidCol != 0, |
| 36 | - "Fix: TCOLREDUCEIDX input shape is invalid, validCol or validRow is 0."); | 77 | + "Fix: TCOLREDUCEIDX input shape is invalid, validCol or validRow is 0"); |
| 37 | - PTO_ASSERT(dstValidRow != 1, "Fix: TCOLREDUCEIDX output validRow must be 1"); | 78 | + PTO_ASSERT(dstValidRow == 1, "Fix: TCOLREDUCEIDX output validRow must be 1"); |
| 38 | - PTO_ASSERT(srcValidCol != dstValidCol, | 79 | + PTO_ASSERT(srcValidCol == dstValidCol, "Fix: TCOLREDUCEIDX input validCol must equal idx output validCol"); |
| 39 | - "Fix: TCOLREDUCEIDX input validCol must be consistent with the output validCol"); | 80 | + |
| 81 | + if constexpr (WithVal) { | ||
| 82 | + static_assert((sizeof(typename TileDataIn::DType) != 1), "Fix: TCOLREDUCEOPSIDX not support b8"); | ||
| 83 | + static_assert(TileDataOutVal::Loc == pto::TileType::Vec, | ||
| 84 | + "Fix: TCOLREDUCEOPSIDX DstVal TileType must be Vec Tile"); | ||
| 85 | + static_assert(TileDataOutVal::isRowMajor && TileDataOutVal::SFractal == SLayout::NoneBox, | ||
| 86 | + "Fix: TCOLREDUCEOPSIDX DstVal only supports Nd fractal Tile"); | ||
| 87 | + static_assert(std::is_same_v<typename TileDataOutVal::DType, typename TileDataIn::DType>, | ||
| 88 | + "Fix: TCOLREDUCEOPSIDX DstVal data type must match input type"); | ||
| 89 | + | ||
| 90 | + PTO_ASSERT(dstValValidRow == 1, "Fix: TCOLREDUCEOPSIDX output value validRow must be 1"); | ||
| 91 | + PTO_ASSERT(dstValValidCol != 0, "Fix: TCOLREDUCEOPSIDX output value validCol must be non-zero"); | ||
| 92 | + PTO_ASSERT(srcValidCol == dstValValidCol, | ||
| 93 | + "Fix: TCOLREDUCEOPSIDX input validCol must equal value output validCol"); | ||
| 94 | + PTO_ASSERT(dstValValidRow == dstValidRow, | ||
| 95 | + "Fix: TCOLREDUCEOPSIDX value and idx output tiles must have same validRow"); | ||
| 96 | + PTO_ASSERT(dstValValidCol == dstValidCol, | ||
| 97 | + "Fix: TCOLREDUCEOPSIDX value and idx output tiles must have same validCol"); | ||
| 98 | + | ||
| 99 | + if constexpr (sizeof(typename TileDataIn::DType) == 2) { | ||
| 100 | + static_assert(std::is_same_v<typename TileDataOutIdx::DType, uint16_t> || | ||
| 101 | + std::is_same_v<typename TileDataOutIdx::DType, int16_t>, | ||
| 102 | + "Fix: TCOLREDUCEOPSIDX DstIdx data type must be s16 or u16 when input type size <= 2 bytes"); | ||
| 103 | + } else { | ||
| 104 | + static_assert(std::is_same_v<typename TileDataOutIdx::DType, uint32_t> || | ||
| 105 | + std::is_same_v<typename TileDataOutIdx::DType, int32_t>, | ||
| 106 | + "Fix: TCOLREDUCEOPSIDX DstIdx data type must be s32 or u32 when input type size is 4 bytes"); | ||
| 107 | + } | ||
| 108 | + } else { | ||
| 109 | + static_assert(std::is_same_v<typename TileDataOutIdx::DType, uint32_t> || | ||
| 110 | + std::is_same_v<typename TileDataOutIdx::DType, int32_t>, | ||
| 111 | + "Fix: TCOLREDUCEOPSIDX DstIdx data type must be s32 or u32"); | ||
| 112 | + } | ||
| 40 | } | 113 | } |
| 114 | + | ||
| 115 | +// ---------------------------------------------------------------------------- | ||
| 116 | +// Per-chunk processing for 8-bit input | ||
| 117 | +// ---------------------------------------------------------------------------- | ||
| 41 | template <typename TileDataOut, typename TileDataIn, bool IsArgMax> | 118 | template <typename TileDataOut, typename TileDataIn, bool IsArgMax> |
| 42 | -__tf__ PTO_INTERNAL void TColReduceIdx8(typename TileDataOut::TileDType __out__ dst, | 119 | +PTO_INTERNAL void TColReduceIdxChunk8(__ubuf__ typename TileDataOut::DType *dstPtr, |
| 43 | - typename TileDataIn::TileDType __in__ src, unsigned srcValidRow, | 120 | + __ubuf__ typename TileDataIn::DType *srcPtr, unsigned srcValidRow, |
| 44 | - unsigned srcValidCol) | 121 | + uint16_t repeatTimes, unsigned srcRowStride, unsigned elementsPerRepeat, |
| 122 | + uint32_t &sregValidCol, MaskReg &pregAll) | ||
| 45 | { | 123 | { |
| 46 | - using TOUT = typename TileDataOut::DType; | ||
| 47 | using TIN = typename TileDataIn::DType; | 124 | using TIN = typename TileDataIn::DType; |
| 48 | - using T = std::conditional_t<std::is_same_v<TIN, int8_t>, vector_s16, | 125 | + using TOUT = typename TileDataOut::DType; |
| 49 | - std::conditional_t<std::is_same_v<TIN, uint8_t>, vector_u16, void>>; | 126 | + using T16 = std::conditional_t<std::is_same_v<TIN, int8_t>, vector_s16, |
| 127 | + std::conditional_t<std::is_same_v<TIN, uint8_t>, vector_u16, void>>; | ||
| 128 | + | ||
| 129 | + vector_s16 vregIndexOldEven; | ||
| 130 | + vector_s16 vregIndexOldOdd; | ||
| 131 | + vector_s16 vregIndexNewEven; | ||
| 132 | + vector_s16 vregIndexNewOdd; | ||
| 133 | + RegTensor<TIN> vregOld; | ||
| 134 | + RegTensor<TIN> vregNew; | ||
| 135 | + T16 vregOldEven; | ||
| 136 | + T16 vregOldOdd; | ||
| 137 | + T16 vregNewEven; | ||
| 138 | + T16 vregNewOdd; | ||
| 139 | + MaskReg preg0; | ||
| 140 | + MaskReg preg1; | ||
| 141 | + MaskReg preg2; | ||
| 142 | + MaskReg preg3; | ||
| 143 | + MaskReg pregSelectEven; | ||
| 144 | + MaskReg pregSelectOdd; | ||
| 145 | + | ||
| 146 | + for (uint16_t j = 0; j < repeatTimes; j++) { | ||
| 147 | + preg0 = plt_b32(sregValidCol, POST_UPDATE); | ||
| 148 | + preg1 = plt_b32(sregValidCol, POST_UPDATE); | ||
| 149 | + preg2 = plt_b32(sregValidCol, POST_UPDATE); | ||
| 150 | + preg3 = plt_b32(sregValidCol, POST_UPDATE); | ||
| 151 | + vdup(vregIndexOldEven, 0, pregAll, MODE_ZEROING); | ||
| 152 | + vdup(vregIndexOldOdd, 0, pregAll, MODE_ZEROING); | ||
| 153 | + vdup(vregIndexNewEven, 0, pregAll, MODE_ZEROING); | ||
| 154 | + vdup(vregIndexNewOdd, 0, pregAll, MODE_ZEROING); | ||
| 155 | + vlds(vregOld, srcPtr, j * elementsPerRepeat, NORM); | ||
| 156 | + vcvt(vregOldEven, vregOld, pregAll, PART_EVEN); | ||
| 157 | + vcvt(vregOldOdd, vregOld, pregAll, PART_ODD); | ||
| 158 | + | ||
| 159 | + for (uint16_t i = 1; i < (uint16_t)srcValidRow; i++) { | ||
| 160 | + vadds(vregIndexNewEven, vregIndexNewEven, 1, pregAll, MODE_ZEROING); | ||
| 161 | + vadds(vregIndexNewOdd, vregIndexNewOdd, 1, pregAll, MODE_ZEROING); | ||
| 162 | + vlds(vregNew, srcPtr, i * srcRowStride + j * elementsPerRepeat, NORM); | ||
| 163 | + vcvt(vregNewEven, vregNew, pregAll, PART_EVEN); | ||
| 164 | + vcvt(vregNewOdd, vregNew, pregAll, PART_ODD); | ||
| 165 | + TColReduceIdxCompare<T16, T16, vector_s16, IsArgMax>(vregOldEven, vregNewEven, vregIndexOldEven, | ||
| 166 | + vregIndexNewEven, pregSelectEven, pregAll); | ||
| 167 | + TColReduceIdxCompare<T16, T16, vector_s16, IsArgMax>(vregOldOdd, vregNewOdd, vregIndexOldOdd, | ||
| 168 | + vregIndexNewOdd, pregSelectOdd, pregAll); | ||
| 169 | + } | ||
| 170 | + | ||
| 171 | + vector_s16 vregTmp0; | ||
| 172 | + vector_s16 vregTmp1; | ||
| 173 | + vintlv(vregTmp0, vregTmp1, vregIndexOldEven, vregIndexOldOdd); | ||
| 174 | + TColReduceIdxStoreParts(vregTmp0, dstPtr, j * elementsPerRepeat, preg0, preg1, pregAll); | ||
| 175 | + TColReduceIdxStoreParts(vregTmp1, dstPtr, j * elementsPerRepeat + 2 * ELE_CNT_B32, preg2, preg3, pregAll); | ||
| 176 | + } | ||
| 177 | +} | ||
| 178 | + | ||
| 179 | +// ---------------------------------------------------------------------------- | ||
| 180 | +// Per-chunk processing for 16-bit and 32-bit input (unified template) | ||
| 181 | +// Differences handled by if constexpr (sizeof(TIN)): | ||
| 182 | +// 16-bit: IdxRegT = vector_s16, compute with pregAll, store via interleave | ||
| 183 | +// 32-bit: IdxRegT = RegTensor<TOUT>, compute with plt_b32, store direct | ||
| 184 | +// ---------------------------------------------------------------------------- | ||
| 185 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, bool IsArgMax, bool WithVal> | ||
| 186 | +PTO_INTERNAL void TColReduceIdxChunk16_32(__ubuf__ typename TileDataOutVal::DType *dstValPtr, | ||
| 187 | + __ubuf__ typename TileDataOutIdx::DType *dstIdxPtr, | ||
| 188 | + __ubuf__ typename TileDataIn::DType *srcPtr, unsigned srcValidRow, | ||
| 189 | + uint16_t repeatTimes, unsigned srcRowStride, unsigned elementsPerRepeat, | ||
| 190 | + uint32_t &sregValidCol, MaskReg &pregAll) | ||
| 191 | +{ | ||
| 192 | + using TIN = typename TileDataIn::DType; | ||
| 193 | + using TIDX = typename TileDataOutIdx::DType; | ||
| 194 | + using IdxRegT = std::conditional_t<sizeof(TIN) == 2, vector_s16, RegTensor<TIDX>>; | ||
| 195 | + | ||
| 196 | + IdxRegT vregIndexOld; | ||
| 197 | + IdxRegT vregIndexNew; | ||
| 198 | + RegTensor<TIN> vregOld; | ||
| 199 | + RegTensor<TIN> vregNew; | ||
| 200 | + MaskReg pregSelect; | ||
| 201 | + | ||
| 202 | + constexpr auto distValue = | ||
| 203 | + std::integral_constant<::DistVST, static_cast<::DistVST>(GetDistVst<TIDX, DistVST::DIST_ONEPT>())>(); | ||
| 204 | + | ||
| 205 | + for (uint16_t j = 0; j < repeatTimes; j++) { | ||
| 206 | + MaskReg pregCmp; | ||
| 207 | + MaskReg pregStore0; | ||
| 208 | + MaskReg pregStore1; | ||
| 209 | + | ||
| 210 | + if constexpr (sizeof(TIN) == 2) { | ||
| 211 | + pregStore0 = plt_b32(sregValidCol, POST_UPDATE); | ||
| 212 | + pregStore1 = plt_b32(sregValidCol, POST_UPDATE); | ||
| 213 | + pregCmp = pregAll; | ||
| 214 | + } else { | ||
| 215 | + pregCmp = plt_b32(sregValidCol, POST_UPDATE); | ||
| 216 | + } | ||
| 217 | + | ||
| 218 | + vdup(vregIndexOld, 0, pregCmp, MODE_ZEROING); | ||
| 219 | + vdup(vregIndexNew, 0, pregCmp, MODE_ZEROING); | ||
| 220 | + vlds(vregOld, srcPtr, j * elementsPerRepeat, NORM); | ||
| 221 | + for (uint16_t i = 1; i < (uint16_t)srcValidRow; i++) { | ||
| 222 | + vadds(vregIndexNew, vregIndexNew, (uint32_t)1, pregCmp, MODE_ZEROING); | ||
| 223 | + vlds(vregNew, srcPtr, i * srcRowStride + j * elementsPerRepeat, NORM); | ||
| 224 | + TColReduceIdxCompare<RegTensor<TIN>, RegTensor<TIN>, IdxRegT, IsArgMax>(vregOld, vregNew, vregIndexOld, | ||
| 225 | + vregIndexNew, pregSelect, pregCmp); | ||
| 226 | + } | ||
| 227 | + | ||
| 228 | + if constexpr (sizeof(TIN) == 2 && !WithVal) { | ||
| 229 | + TColReduceIdxStoreParts(vregIndexOld, dstIdxPtr, j * elementsPerRepeat, pregStore0, pregStore1, pregAll); | ||
| 230 | + } else { | ||
| 231 | + vsts(vregIndexOld, dstIdxPtr, j * elementsPerRepeat, NORM_B32, pregCmp); | ||
| 232 | + } | ||
| 233 | + | ||
| 234 | + if constexpr (WithVal) { | ||
| 235 | + vsts(vregOld, dstValPtr, j * elementsPerRepeat, NORM_B32, pregCmp); | ||
| 236 | + } | ||
| 237 | + } | ||
| 238 | +} | ||
| 239 | + | ||
| 240 | +// ---------------------------------------------------------------------------- | ||
| 241 | +// Unified implementation (thin dispatcher by sizeof(TIN)) | ||
| 242 | +// ---------------------------------------------------------------------------- | ||
| 243 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, bool IsArgMax, bool WithVal = false> | ||
| 244 | +__tf__ PTO_INTERNAL void TColReduceIdxImpl(typename TileDataOutVal::TileDType __out__ dstValData, | ||
| 245 | + typename TileDataOutIdx::TileDType __out__ dstIdxData, | ||
| 246 | + typename TileDataIn::TileDType __in__ src, unsigned srcValidRow, | ||
| 247 | + unsigned srcValidCol) | ||
| 248 | +{ | ||
| 249 | + using TIN = typename TileDataIn::DType; | ||
| 250 | + using TOUT = typename TileDataOutIdx::DType; | ||
| 50 | 251 | ||
| 51 | constexpr unsigned srcRowStride = TileDataIn::Cols; | 252 | constexpr unsigned srcRowStride = TileDataIn::Cols; |
| 52 | constexpr unsigned elementsPerRepeat = REPEAT_BYTE / sizeof(TIN); | 253 | constexpr unsigned elementsPerRepeat = REPEAT_BYTE / sizeof(TIN); |
| 53 | uint16_t repeatTimes = CeilDivision(srcValidCol, elementsPerRepeat); | 254 | uint16_t repeatTimes = CeilDivision(srcValidCol, elementsPerRepeat); |
| 54 | 255 | ||
| 55 | - __ubuf__ TOUT *dstPtr = (__ubuf__ TOUT *)__cce_get_tile_ptr(dst); | 256 | + __ubuf__ typename TileDataOutVal::DType *dstValPtr = |
| 56 | - __ubuf__ TIN *srcPtr = (__ubuf__ TIN *)__cce_get_tile_ptr(src); | 257 | + (__ubuf__ typename TileDataOutVal::DType *)__cce_get_tile_ptr(dstValData); |
| 57 | - __VEC_SCOPE__ | 258 | + __ubuf__ TOUT *dstIdxPtr = (__ubuf__ TOUT *)__cce_get_tile_ptr(dstIdxData); |
| 58 | - { | ||
| 59 | - vector_s16 vregIndexOldEven; | ||
| 60 | - vector_s16 vregIndexOldOdd; | ||
| 61 | - vector_s16 vregIndexOutput0; | ||
| 62 | - vector_s16 vregIndexOutput1; | ||
| 63 | - vector_s16 vregIndexNewEven; | ||
| 64 | - vector_s16 vregIndexNewOdd; | ||
| 65 | - RegTensor<TOUT> outputIndexEven; | ||
| 66 | - RegTensor<TOUT> outputIndexOdd; | ||
| 67 | - RegTensor<TOUT> outputIndex0; | ||
| 68 | - RegTensor<TOUT> outputIndex1; | ||
| 69 | - RegTensor<TIN> vregOld; | ||
| 70 | - RegTensor<TIN> vregNew; | ||
| 71 | - T vregOldEven; | ||
| 72 | - T vregOldOdd; | ||
| 73 | - T vregNewEven; | ||
| 74 | - T vregNewOdd; | ||
| 75 | - MaskReg preg = pset_b8(PAT_ALL); | ||
| 76 | - MaskReg selectEven; | ||
| 77 | - MaskReg selectOdd; | ||
| 78 | - MaskReg preg0; | ||
| 79 | - MaskReg preg1; | ||
| 80 | - MaskReg preg2; | ||
| 81 | - MaskReg preg3; | ||
| 82 | - uint32_t sreg = srcValidCol; | ||
| 83 | - | ||
| 84 | - for (uint16_t j = 0; j < repeatTimes; j++) { | ||
| 85 | - preg0 = plt_b32(sreg, POST_UPDATE); | ||
| 86 | - preg1 = plt_b32(sreg, POST_UPDATE); | ||
| 87 | - preg2 = plt_b32(sreg, POST_UPDATE); | ||
| 88 | - preg3 = plt_b32(sreg, POST_UPDATE); | ||
| 89 | - vdup(vregIndexOldEven, 0, preg, MODE_ZEROING); | ||
| 90 | - vdup(vregIndexOldOdd, 0, preg, MODE_ZEROING); | ||
| 91 | - vdup(vregIndexNewEven, 0, preg, MODE_ZEROING); | ||
| 92 | - vdup(vregIndexNewOdd, 0, preg, MODE_ZEROING); | ||
| 93 | - vlds(vregOld, srcPtr, j * elementsPerRepeat, NORM); | ||
| 94 | - vcvt(vregOldEven, vregOld, preg, PART_EVEN); | ||
| 95 | - vcvt(vregOldOdd, vregOld, preg, PART_ODD); | ||
| 96 | - | ||
| 97 | - for (uint16_t i = 1; i < (uint16_t)srcValidRow; i++) { | ||
| 98 | - vadds(vregIndexNewEven, vregIndexNewEven, 1, preg, MODE_ZEROING); | ||
| 99 | - vadds(vregIndexNewOdd, vregIndexNewOdd, 1, preg, MODE_ZEROING); | ||
| 100 | - vlds(vregNew, srcPtr, i * srcRowStride + j * elementsPerRepeat, NORM); | ||
| 101 | - vcvt(vregNewEven, vregNew, preg, PART_EVEN); | ||
| 102 | - vcvt(vregNewOdd, vregNew, preg, PART_ODD); | ||
| 103 | - if constexpr (IsArgMax) { | ||
| 104 | - vcmp_gt(selectEven, vregNewEven, vregOldEven, preg); | ||
| 105 | - vcmp_gt(selectOdd, vregNewOdd, vregOldOdd, preg); | ||
| 106 | - vsel(vregIndexOldEven, vregIndexNewEven, vregIndexOldEven, selectEven); | ||
| 107 | - vsel(vregIndexOldOdd, vregIndexNewOdd, vregIndexOldOdd, selectOdd); | ||
| 108 | - vmax(vregOldEven, vregOldEven, vregNewEven, preg, MODE_ZEROING); | ||
| 109 | - vmax(vregOldOdd, vregOldOdd, vregNewOdd, preg, MODE_ZEROING); | ||
| 110 | - } else { | ||
| 111 | - vcmp_lt(selectEven, vregNewEven, vregOldEven, preg); | ||
| 112 | - vcmp_lt(selectOdd, vregNewOdd, vregOldOdd, preg); | ||
| 113 | - vsel(vregIndexOldEven, vregIndexNewEven, vregIndexOldEven, selectEven); | ||
| 114 | - vsel(vregIndexOldOdd, vregIndexNewOdd, vregIndexOldOdd, selectOdd); | ||
| 115 | - vmin(vregOldEven, vregOldEven, vregNewEven, preg, MODE_ZEROING); | ||
| 116 | - vmin(vregOldOdd, vregOldOdd, vregNewOdd, preg, MODE_ZEROING); | ||
| 117 | - } | ||
| 118 | - } | ||
| 119 | - vintlv(vregIndexOutput0, vregIndexOutput1, vregIndexOldEven, vregIndexOldOdd); | ||
| 120 | - vcvt(outputIndexEven, vregIndexOutput0, preg, PART_EVEN); | ||
| 121 | - vcvt(outputIndexOdd, vregIndexOutput0, preg, PART_ODD); | ||
| 122 | - vintlv(outputIndex0, outputIndex1, outputIndexEven, outputIndexOdd); | ||
| 123 | - vsts(outputIndex0, dst, j * elementsPerRepeat, NORM_B32, preg0); | ||
| 124 | - vsts(outputIndex1, dst, j * elementsPerRepeat + ELE_CNT_B32, NORM_B32, preg1); | ||
| 125 | - | ||
| 126 | - vcvt(outputIndexEven, vregIndexOutput1, preg, PART_EVEN); | ||
| 127 | - vcvt(outputIndexOdd, vregIndexOutput1, preg, PART_ODD); | ||
| 128 | - vintlv(outputIndex0, outputIndex1, outputIndexEven, outputIndexOdd); | ||
| 129 | - vsts(outputIndex0, dst, j * elementsPerRepeat + 2 * ELE_CNT_B32, NORM_B32, preg2); | ||
| 130 | - vsts(outputIndex1, dst, j * elementsPerRepeat + 3 * ELE_CNT_B32, NORM_B32, preg3); | ||
| 131 | - } | ||
| 132 | - } | ||
| 133 | -} | ||
| 134 | - | ||
| 135 | -template <typename TileDataOut, typename TileDataIn, bool IsArgMax> | ||
| 136 | -__tf__ PTO_INTERNAL void TColReduceIdx16(typename TileDataOut::TileDType __out__ dst, | ||
| 137 | - typename TileDataIn::TileDType __in__ src, unsigned srcValidRow, | ||
| 138 | - unsigned srcValidCol) | ||
| 139 | -{ | ||
| 140 | - using TIN = typename TileDataIn::DType; | ||
| 141 | - using TOUT = typename TileDataOut::DType; | ||
| 142 | - constexpr unsigned srcRowStride = TileDataIn::Cols; | ||
| 143 | - constexpr unsigned elementsPerRepeat = REPEAT_BYTE / sizeof(TIN); | ||
| 144 | - uint16_t repeatTimes = CeilDivision(srcValidCol, elementsPerRepeat); | ||
| 145 | - __ubuf__ TOUT *dstPtr = (__ubuf__ TOUT *)__cce_get_tile_ptr(dst); | ||
| 146 | __ubuf__ TIN *srcPtr = (__ubuf__ TIN *)__cce_get_tile_ptr(src); | 259 | __ubuf__ TIN *srcPtr = (__ubuf__ TIN *)__cce_get_tile_ptr(src); |
| 147 | 260 | ||
| 148 | __VEC_SCOPE__ | 261 | __VEC_SCOPE__ |
| 149 | { | 262 | { |
| 150 | - vector_s16 vregIndexOld; | 263 | + MaskReg pregAll = pset_b8(PAT_ALL); |
| 151 | - vector_s16 vregIndexNew; | 264 | + uint32_t sregValidCol = srcValidCol; |
| 152 | - RegTensor<TOUT> outputIndexEven; | ||
| 153 | - RegTensor<TOUT> outputIndexOdd; | ||
| 154 | - RegTensor<TOUT> outputIndex0; | ||
| 155 | - RegTensor<TOUT> outputIndex1; | ||
| 156 | - MaskReg preg = pset_b8(PAT_ALL); | ||
| 157 | - MaskReg preg0; | ||
| 158 | - MaskReg preg1; | ||
| 159 | - MaskReg select; | ||
| 160 | - RegTensor<TIN> vregOld; | ||
| 161 | - RegTensor<TIN> vregNew; | ||
| 162 | - uint32_t sreg = srcValidCol; | ||
| 163 | 265 | ||
| 164 | - for (uint16_t j = 0; j < repeatTimes; j++) { | 266 | + if constexpr (sizeof(TIN) == 1) { |
| 165 | - preg0 = plt_b32(sreg, POST_UPDATE); | 267 | + TColReduceIdxChunk8<TileDataOutIdx, TileDataIn, IsArgMax>( |
| 166 | - preg1 = plt_b32(sreg, POST_UPDATE); | 268 | + dstIdxPtr, srcPtr, srcValidRow, repeatTimes, srcRowStride, elementsPerRepeat, sregValidCol, pregAll); |
| 167 | - vdup(vregIndexOld, 0, preg, MODE_ZEROING); | 269 | + } else { |
| 168 | - vdup(vregIndexNew, 0, preg, MODE_ZEROING); | 270 | + TColReduceIdxChunk16_32<TileDataOutVal, TileDataOutIdx, TileDataIn, IsArgMax, WithVal>( |
| 169 | - vlds(vregOld, srcPtr, j * elementsPerRepeat, NORM); | 271 | + dstValPtr, dstIdxPtr, srcPtr, srcValidRow, repeatTimes, srcRowStride, elementsPerRepeat, sregValidCol, |
| 170 | - for (uint16_t i = 1; i < (uint16_t)srcValidRow; i++) { | 272 | + pregAll); |
| 171 | - vadds(vregIndexNew, vregIndexNew, 1, preg, MODE_ZEROING); | ||
| 172 | - vlds(vregNew, srcPtr, i * srcRowStride + j * elementsPerRepeat, NORM); | ||
| 173 | - if constexpr (IsArgMax) { | ||
| 174 | - vcmp_gt(select, vregNew, vregOld, preg); | ||
| 175 | - vsel(vregIndexOld, vregIndexNew, vregIndexOld, select); | ||
| 176 | - vmax(vregOld, vregOld, vregNew, preg, MODE_ZEROING); | ||
| 177 | - } else { | ||
| 178 | - vcmp_lt(select, vregNew, vregOld, preg); | ||
| 179 | - vsel(vregIndexOld, vregIndexNew, vregIndexOld, select); | ||
| 180 | - vmin(vregOld, vregOld, vregNew, preg, MODE_ZEROING); | ||
| 181 | - } | ||
| 182 | - } | ||
| 183 | - vcvt(outputIndexEven, vregIndexOld, preg, PART_EVEN); | ||
| 184 | - vcvt(outputIndexOdd, vregIndexOld, preg, PART_ODD); | ||
| 185 | - vintlv(outputIndex0, outputIndex1, outputIndexEven, outputIndexOdd); | ||
| 186 | - vsts(outputIndex0, dst, j * elementsPerRepeat, NORM_B32, preg0); | ||
| 187 | - vsts(outputIndex1, dst, j * elementsPerRepeat + ELE_CNT_B32, NORM_B32, preg1); | ||
| 188 | } | 273 | } |
| 189 | } | 274 | } |
| 190 | } | 275 | } |
| 191 | 276 | ||
| 192 | -template <typename TileDataOut, typename TileDataIn, bool IsArgMax> | 277 | +// ---------------------------------------------------------------------------- |
| 193 | -__tf__ PTO_INTERNAL void TColReduceIdx32(typename TileDataOut::TileDType __out__ dst, | 278 | +// Public dispatch (unchanged interface — keeps backward compatibility) |
| 194 | - typename TileDataIn::TileDType __in__ src, unsigned srcValidRow, | 279 | +// ---------------------------------------------------------------------------- |
| 195 | - unsigned srcValidCol) | ||
| 196 | -{ | ||
| 197 | - using TIN = typename TileDataIn::DType; | ||
| 198 | - using TOUT = typename TileDataOut::DType; | ||
| 199 | - | ||
| 200 | - constexpr unsigned srcRowStride = TileDataIn::Cols; | ||
| 201 | - constexpr unsigned elementsPerRepeat = REPEAT_BYTE / sizeof(TIN); | ||
| 202 | - __ubuf__ TOUT *dstPtr = (__ubuf__ TOUT *)__cce_get_tile_ptr(dst); | ||
| 203 | - __ubuf__ TIN *srcPtr = (__ubuf__ TIN *)__cce_get_tile_ptr(src); | ||
| 204 | - uint16_t repeatTimes = CeilDivision(srcValidCol, elementsPerRepeat); | ||
| 205 | - | ||
| 206 | - __VEC_SCOPE__ | ||
| 207 | - { | ||
| 208 | - RegTensor<TOUT> vregIndexOld; | ||
| 209 | - RegTensor<TOUT> vregIndexNew; | ||
| 210 | - RegTensor<TIN> vregOld; | ||
| 211 | - RegTensor<TIN> vregNew; | ||
| 212 | - MaskReg pg; | ||
| 213 | - MaskReg select; | ||
| 214 | - uint32_t sreg = srcValidCol; | ||
| 215 | - | ||
| 216 | - for (uint16_t j = 0; j < repeatTimes; j++) { | ||
| 217 | - pg = plt_b32(sreg, POST_UPDATE); | ||
| 218 | - vdup(vregIndexOld, 0, pg, MODE_ZEROING); | ||
| 219 | - vdup(vregIndexNew, 0, pg, MODE_ZEROING); | ||
| 220 | - vlds(vregOld, srcPtr, j * elementsPerRepeat, NORM); | ||
| 221 | - for (uint16_t i = 1; i < (uint16_t)srcValidRow; i++) { | ||
| 222 | - vadds(vregIndexNew, vregIndexNew, (uint32_t)1, pg, MODE_ZEROING); | ||
| 223 | - vlds(vregNew, srcPtr, i * srcRowStride + j * elementsPerRepeat, NORM); | ||
| 224 | - if constexpr (IsArgMax) { | ||
| 225 | - vcmp_gt(select, vregNew, vregOld, pg); | ||
| 226 | - vsel(vregIndexOld, vregIndexNew, vregIndexOld, select); | ||
| 227 | - vmax(vregOld, vregOld, vregNew, pg, MODE_ZEROING); | ||
| 228 | - } else { | ||
| 229 | - vcmp_lt(select, vregNew, vregOld, pg); | ||
| 230 | - vsel(vregIndexOld, vregIndexNew, vregIndexOld, select); | ||
| 231 | - vmin(vregOld, vregOld, vregNew, pg, MODE_ZEROING); | ||
| 232 | - } | ||
| 233 | - } | ||
| 234 | - vsts(vregIndexOld, dstPtr, j * elementsPerRepeat, NORM_B32, pg); | ||
| 235 | - } | ||
| 236 | - } | ||
| 237 | -} | ||
| 238 | template <typename TileDataOut, typename TileDataIn, bool IsArgMax> | 280 | template <typename TileDataOut, typename TileDataIn, bool IsArgMax> |
| 239 | PTO_INTERNAL void TCOLARG_DISPATCH(TileDataOut &dst, TileDataIn &src) | 281 | PTO_INTERNAL void TCOLARG_DISPATCH(TileDataOut &dst, TileDataIn &src) |
| 240 | { | 282 | { |
| 241 | unsigned srcValidRow = src.GetValidRow(); | 283 | unsigned srcValidRow = src.GetValidRow(); |
| 242 | unsigned srcValidCol = src.GetValidCol(); | 284 | unsigned srcValidCol = src.GetValidCol(); |
| 243 | - TColReduceIdxCheck<TileDataOut, TileDataIn>(srcValidRow, srcValidCol, dst.GetValidRow(), dst.GetValidCol()); | 285 | + TColReduceIdxCheck<TileDataIn, TileDataOut, TileDataIn>(srcValidRow, srcValidCol, dst.GetValidRow(), |
| 244 | - | 286 | + dst.GetValidCol()); |
| 245 | - if constexpr (sizeof(typename TileDataIn::DType) == 1) { | 287 | + TColReduceIdxImpl<TileDataIn, TileDataOut, TileDataIn, IsArgMax>(src.data(), dst.data(), src.data(), srcValidRow, |
| 246 | - TColReduceIdx8<TileDataOut, TileDataIn, IsArgMax>(dst.data(), src.data(), srcValidRow, srcValidCol); | 288 | + srcValidCol); |
| 247 | - } else if (sizeof(typename TileDataIn::DType) == 2) { | ||
| 248 | - TColReduceIdx16<TileDataOut, TileDataIn, IsArgMax>(dst.data(), src.data(), srcValidRow, srcValidCol); | ||
| 249 | - } else if (sizeof(typename TileDataIn::DType) == 4) { | ||
| 250 | - TColReduceIdx32<TileDataOut, TileDataIn, IsArgMax>(dst.data(), src.data(), srcValidRow, srcValidCol); | ||
| 251 | - } | ||
| 252 | } | 289 | } |
| 290 | + | ||
| 291 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, bool IsArgMax> | ||
| 292 | +PTO_INTERNAL void TCOLARG_DISPATCH(TileDataOutVal &dstVal, TileDataOutIdx &dstIdx, TileDataIn &src) | ||
| 293 | +{ | ||
| 294 | + unsigned srcValidRow = src.GetValidRow(); | ||
| 295 | + unsigned srcValidCol = src.GetValidCol(); | ||
| 296 | + TColReduceIdxCheck<TileDataOutVal, TileDataOutIdx, TileDataIn, true>(srcValidRow, srcValidCol, dstIdx.GetValidRow(), | ||
| 297 | + dstIdx.GetValidCol(), dstVal.GetValidRow(), | ||
| 298 | + dstVal.GetValidCol()); | ||
| 299 | + TColReduceIdxImpl<TileDataOutVal, TileDataOutIdx, TileDataIn, IsArgMax, true>(dstVal.data(), dstIdx.data(), | ||
| 300 | + src.data(), srcValidRow, srcValidCol); | ||
| 301 | +} | ||
| 302 | + | ||
| 303 | +// ========================================================================================== | ||
| 253 | template <typename TileDataOut, typename TileDataIn, typename TileDataTmp> | 304 | template <typename TileDataOut, typename TileDataIn, typename TileDataTmp> |
| 254 | PTO_INTERNAL void TCOLARGMIN_IMPL(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp) | 305 | PTO_INTERNAL void TCOLARGMIN_IMPL(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp) |
| 255 | { | 306 | { |
| 256 | - TCOLARG_DISPATCH<TileDataOut, TileDataIn, false>(dst, src); // Min | 307 | + TCOLARG_DISPATCH<TileDataOut, TileDataIn, false>(dst, src); |
| 257 | } | 308 | } |
| 309 | + | ||
| 258 | template <typename TileDataOut, typename TileDataIn, typename TileDataTmp> | 310 | template <typename TileDataOut, typename TileDataIn, typename TileDataTmp> |
| 259 | PTO_INTERNAL void TCOLARGMAX_IMPL(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp) | 311 | PTO_INTERNAL void TCOLARGMAX_IMPL(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp) |
| 260 | { | 312 | { |
| 261 | - TCOLARG_DISPATCH<TileDataOut, TileDataIn, true>(dst, src); // Max | 313 | + TCOLARG_DISPATCH<TileDataOut, TileDataIn, true>(dst, src); |
| 262 | } | 314 | } |
| 315 | + | ||
| 316 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp> | ||
| 317 | +PTO_INTERNAL void TCOLARGMAX_IMPL(TileDataOutVal &dstVal, TileDataOutIdx &dstIdx, TileDataIn &src, TileDataTmp &tmp) | ||
| 318 | +{ | ||
| 319 | + TCOLARG_DISPATCH<TileDataOutVal, TileDataOutIdx, TileDataIn, true>(dstVal, dstIdx, src); | ||
| 320 | +} | ||
| 321 | + | ||
| 322 | +template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp> | ||
| 323 | +PTO_INTERNAL void TCOLARGMIN_IMPL(TileDataOutVal &dstVal, TileDataOutIdx &dstIdx, TileDataIn &src, TileDataTmp &tmp) | ||
| 324 | +{ | ||
| 325 | + TCOLARG_DISPATCH<TileDataOutVal, TileDataOutIdx, TileDataIn, false>(dstVal, dstIdx, src); | ||
| 326 | +} | ||
| 327 | + | ||
| 263 | } // namespace pto | 328 | } // namespace pto |
| 264 | 329 | ||
| @@ -13,6 +13,7 @@ | |||
| 13 | import os | 13 | import os |
| 14 | import numpy as np | 14 | import numpy as np |
| 15 | import math | 15 | import math |
| 16 | + | ||
| 16 | np.random.seed(19) | 17 | np.random.seed(19) |
| 17 | 18 | ||
| 18 | 19 | ||
| @@ -33,22 +34,36 @@ def gen_golden_data(param): | |||
| 33 | input_arr = np.random.uniform(low=value_min, high=value_max, size=(row, col)).astype(data_type) | 34 | input_arr = np.random.uniform(low=value_min, high=value_max, size=(row, col)).astype(data_type) |
| 34 | output_arr = np.argmax(input_arr[0:valid_row], axis=0) | 35 | output_arr = np.argmax(input_arr[0:valid_row], axis=0) |
| 35 | output_arr[valid_col:] = 0 | 36 | output_arr[valid_col:] = 0 |
| 36 | - dst_col = math.ceil(valid_col / 8) * 8 | 37 | + if not param.idx: |
| 37 | - output_arr = output_arr[:dst_col] | 38 | + dst_col = math.ceil(valid_col / 8) * 8 |
| 38 | - # 先计算, 再强转类型, 保证结果精度不裂化 | 39 | + output_arr = output_arr[:dst_col] |
| 39 | - output_arr = output_arr.astype(np.int32) | 40 | + # 先计算, 再强转类型, 保证结果精度不裂化 |
| 40 | - input_arr.tofile('input.bin') | 41 | + output_arr = output_arr.astype(np.int32) |
| 41 | - output_arr.tofile('golden.bin') | 42 | + input_arr.tofile("input.bin") |
| 43 | + output_arr.tofile("golden.bin") | ||
| 44 | + else: | ||
| 45 | + input_arr.tofile("input.bin") | ||
| 46 | + output_idx = output_arr[:col] | ||
| 47 | + output_arr = np.max(input_arr[0:valid_row], axis=0) | ||
| 48 | + if input_arr.itemsize == 2: | ||
| 49 | + output_idx = output_idx.astype(np.int16) | ||
| 50 | + else: | ||
| 51 | + output_idx = output_idx.astype(np.int32) | ||
| 52 | + output_idx.tofile("idx.bin") | ||
| 53 | + output_arr[valid_col:] = 0 | ||
| 54 | + output_arr.tofile("golden.bin") | ||
| 42 | 55 | ||
| 43 | 56 | ||
| 44 | class TColCMaxParams: | 57 | class TColCMaxParams: |
| 45 | - def __init__(self, name, data_type, row, valid_row, col, valid_col): | 58 | + def __init__(self, name, data_type, row, valid_row, col, valid_col, idx=False): |
| 46 | self.name = name | 59 | self.name = name |
| 47 | self.data_type = data_type | 60 | self.data_type = data_type |
| 48 | self.row = row | 61 | self.row = row |
| 49 | self.valid_row = valid_row | 62 | self.valid_row = valid_row |
| 50 | self.col = col | 63 | self.col = col |
| 51 | self.valid_col = valid_col | 64 | self.valid_col = valid_col |
| 65 | + self.idx = idx | ||
| 66 | + | ||
| 52 | 67 | ||
| 53 | if __name__ == "__main__": | 68 | if __name__ == "__main__": |
| 54 | case_params_list = [ | 69 | case_params_list = [ |
| @@ -70,7 +85,26 @@ if __name__ == "__main__": | |||
| 70 | TColCMaxParams("TCOLCMAXTest.case84", np.float32, 16, 16, 32, 31), | 85 | TColCMaxParams("TCOLCMAXTest.case84", np.float32, 16, 16, 32, 31), |
| 71 | TColCMaxParams("TCOLCMAXTest.case91", np.uint16, 16, 16, 128, 120), | 86 | TColCMaxParams("TCOLCMAXTest.case91", np.uint16, 16, 16, 128, 120), |
| 72 | TColCMaxParams("TCOLCMAXTest.case92", np.float16, 16, 16, 96, 88), | 87 | TColCMaxParams("TCOLCMAXTest.case92", np.float16, 16, 16, 96, 88), |
| 73 | - TColCMaxParams("TCOLCMAXTest.case93", np.uint16, 4, 4, 48, 34) | 88 | + TColCMaxParams("TCOLCMAXTest.case93", np.uint16, 4, 4, 48, 34), |
| 89 | + TColCMaxParams("TCOLCMAXTest.case001", np.float32, 1, 1, 256, 255, True), | ||
| 90 | + TColCMaxParams("TCOLCMAXTest.case002", np.float32, 16, 16, 128, 127, True), | ||
| 91 | + TColCMaxParams("TCOLCMAXTest.case003", np.float32, 16, 15, 256, 255, True), | ||
| 92 | + TColCMaxParams("TCOLCMAXTest.case011", np.float16, 1, 1, 256, 255, True), | ||
| 93 | + TColCMaxParams("TCOLCMAXTest.case012", np.float16, 16, 16, 128, 127, True), | ||
| 94 | + TColCMaxParams("TCOLCMAXTest.case013", np.float16, 16, 15, 256, 255, True), | ||
| 95 | + TColCMaxParams("TCOLCMAXTest.case051", np.uint16, 1, 1, 256, 255, True), | ||
| 96 | + TColCMaxParams("TCOLCMAXTest.case052", np.uint16, 16, 16, 128, 127, True), | ||
| 97 | + TColCMaxParams("TCOLCMAXTest.case053", np.uint16, 16, 15, 256, 255, True), | ||
| 98 | + TColCMaxParams("TCOLCMAXTest.case071", np.uint32, 1, 1, 256, 255, True), | ||
| 99 | + TColCMaxParams("TCOLCMAXTest.case072", np.uint32, 16, 16, 128, 127, True), | ||
| 100 | + TColCMaxParams("TCOLCMAXTest.case073", np.uint32, 16, 15, 256, 255, True), | ||
| 101 | + TColCMaxParams("TCOLCMAXTest.case081", np.float16, 16, 16, 32, 32, True), | ||
| 102 | + TColCMaxParams("TCOLCMAXTest.case082", np.uint16, 16, 16, 32, 32, True), | ||
| 103 | + TColCMaxParams("TCOLCMAXTest.case083", np.uint32, 16, 16, 32, 31, True), | ||
| 104 | + TColCMaxParams("TCOLCMAXTest.case084", np.float32, 16, 16, 32, 31, True), | ||
| 105 | + TColCMaxParams("TCOLCMAXTest.case091", np.uint16, 16, 16, 128, 120, True), | ||
| 106 | + TColCMaxParams("TCOLCMAXTest.case092", np.float16, 16, 16, 96, 88, True), | ||
| 107 | + TColCMaxParams("TCOLCMAXTest.case093", np.uint16, 4, 4, 48, 34, True), | ||
| 74 | ] | 108 | ] |
| 75 | 109 | ||
| 76 | for _, case in enumerate(case_params_list): | 110 | for _, case in enumerate(case_params_list): |
| @@ -79,4 +113,4 @@ if __name__ == "__main__": | |||
| 79 | original_dir = os.getcwd() | 113 | original_dir = os.getcwd() |
| 80 | os.chdir(case.name) | 114 | os.chdir(case.name) |
| 81 | gen_golden_data(case) | 115 | gen_golden_data(case) |
| 82 | - os.chdir(original_dir) | 116 | + os.chdir(original_dir) |
| @@ -18,6 +18,9 @@ using namespace PtoTestCommon; | |||
| 18 | template <uint32_t caseId> | 18 | template <uint32_t caseId> |
| 19 | void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream); | 19 | void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream); |
| 20 | 20 | ||
| 21 | +template <uint32_t caseId> | ||
| 22 | +void launchTCOLIDXVALMAXCase(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 23 | + | ||
| 21 | std::string GetGoldenDir() | 24 | std::string GetGoldenDir() |
| 22 | { | 25 | { |
| 23 | const testing::TestInfo *testInfo = testing::UnitTest::GetInstance()->current_test_info(); | 26 | const testing::TestInfo *testInfo = testing::UnitTest::GetInstance()->current_test_info(); |
| @@ -35,6 +38,11 @@ public: | |||
| 35 | void *dstDevice; | 38 | void *dstDevice; |
| 36 | void *srcDevice; | 39 | void *srcDevice; |
| 37 | 40 | ||
| 41 | + void *dstHostVal; | ||
| 42 | + void *dstHostIdx; | ||
| 43 | + void *dstDeviceVal; | ||
| 44 | + void *dstDeviceIdx; | ||
| 45 | + | ||
| 38 | protected: | 46 | protected: |
| 39 | void SetUp() override | 47 | void SetUp() override |
| 40 | { | 48 | { |
| @@ -64,6 +72,27 @@ protected: | |||
| 64 | return ResultCmp(golden, result, eps, 0, 1000, false, true); | 72 | return ResultCmp(golden, result, eps, 0, 1000, false, true); |
| 65 | } | 73 | } |
| 66 | 74 | ||
| 75 | + template <typename TVal, typename TIdx> | ||
| 76 | + bool CompareGoldenValIdx(size_t dstByteSize, bool printAllEn = false) | ||
| 77 | + { | ||
| 78 | + std::vector<TIdx> goldenIdx(dstByteSize); | ||
| 79 | + std::vector<TIdx> resultIdx(dstByteSize); | ||
| 80 | + std::vector<TVal> goldenVal(dstByteSize); | ||
| 81 | + std::vector<TVal> resultVal(dstByteSize); | ||
| 82 | + | ||
| 83 | + float eps = 0.001f; | ||
| 84 | + ReadFile(GetGoldenDir() + "/golden.bin", dstByteSize, goldenVal.data(), dstByteSize); | ||
| 85 | + ReadFile(GetGoldenDir() + "/output_val.bin", dstByteSize, resultVal.data(), dstByteSize); | ||
| 86 | + ReadFile(GetGoldenDir() + "/idx.bin", dstByteSize, goldenIdx.data(), dstByteSize); | ||
| 87 | + ReadFile(GetGoldenDir() + "/output_idx.bin", dstByteSize, resultIdx.data(), dstByteSize); | ||
| 88 | + if (printAllEn) { | ||
| 89 | + return ResultCmp(goldenVal, resultVal, eps, 0, 1000, true) && | ||
| 90 | + ResultCmp(goldenIdx, resultIdx, eps, 0, 1000, true); | ||
| 91 | + } | ||
| 92 | + return ResultCmp(goldenVal, resultVal, eps, 0, 1000, false, true) && | ||
| 93 | + ResultCmp(goldenIdx, resultIdx, eps, 0, 1000, false, true); | ||
| 94 | + } | ||
| 95 | + | ||
| 67 | template <uint32_t caseId, typename T, int srcRow, int srcValidRow, int dstRow, int col, int validCol> | 96 | template <uint32_t caseId, typename T, int srcRow, int srcValidRow, int dstRow, int col, int validCol> |
| 68 | bool TCOLCMAXTestFramework() | 97 | bool TCOLCMAXTestFramework() |
| 69 | { | 98 | { |
| @@ -91,6 +120,43 @@ protected: | |||
| 91 | 120 | ||
| 92 | return CompareGolden(dstByteSize); | 121 | return CompareGolden(dstByteSize); |
| 93 | } | 122 | } |
| 123 | + | ||
| 124 | + template <uint32_t caseId, typename TVal, typename TIdx, int srcRow, int srcValidRow, int dstRow, int col, | ||
| 125 | + int validCol> | ||
| 126 | + bool TCOLCMAXTestFramework() | ||
| 127 | + { | ||
| 128 | + // int dstCol = (validCol + 7) / 8 * 8; | ||
| 129 | + size_t dstIdxByteSize = dstRow * col * sizeof(TIdx); | ||
| 130 | + size_t dstValByteSize = dstRow * col * sizeof(TVal); | ||
| 131 | + size_t srcByteSize = srcRow * col * sizeof(TVal); | ||
| 132 | + aclrtMallocHost(&dstHostIdx, dstIdxByteSize); | ||
| 133 | + aclrtMallocHost(&dstHostVal, dstValByteSize); | ||
| 134 | + aclrtMallocHost(&srcHost, srcByteSize); | ||
| 135 | + aclrtMalloc(&dstDeviceIdx, dstIdxByteSize, ACL_MEM_MALLOC_HUGE_FIRST); | ||
| 136 | + aclrtMalloc(&dstDeviceVal, dstValByteSize, ACL_MEM_MALLOC_HUGE_FIRST); | ||
| 137 | + aclrtMalloc(&srcDevice, srcByteSize, ACL_MEM_MALLOC_HUGE_FIRST); | ||
| 138 | + | ||
| 139 | + ReadFile(GetGoldenDir() + "/input.bin", srcByteSize, srcHost, srcByteSize); | ||
| 140 | + aclrtMemcpy(srcDevice, srcByteSize, srcHost, srcByteSize, ACL_MEMCPY_HOST_TO_DEVICE); | ||
| 141 | + | ||
| 142 | + launchTCOLIDXVALMAXCase<caseId>(dstDeviceVal, dstDeviceIdx, srcDevice, stream); | ||
| 143 | + aclrtSynchronizeStream(stream); | ||
| 144 | + | ||
| 145 | + aclrtMemcpy(dstHostIdx, dstIdxByteSize, dstDeviceIdx, dstIdxByteSize, ACL_MEMCPY_DEVICE_TO_HOST); | ||
| 146 | + aclrtMemcpy(dstHostVal, dstValByteSize, dstDeviceVal, dstValByteSize, ACL_MEMCPY_DEVICE_TO_HOST); | ||
| 147 | + | ||
| 148 | + WriteFile(GetGoldenDir() + "/output_val.bin", dstHostVal, dstValByteSize); | ||
| 149 | + WriteFile(GetGoldenDir() + "/output_idx.bin", dstHostIdx, dstIdxByteSize); | ||
| 150 | + | ||
| 151 | + aclrtFree(dstDeviceIdx); | ||
| 152 | + aclrtFree(dstDeviceVal); | ||
| 153 | + aclrtFree(srcDevice); | ||
| 154 | + aclrtFreeHost(dstHostVal); | ||
| 155 | + aclrtFreeHost(dstHostIdx); | ||
| 156 | + aclrtFreeHost(srcHost); | ||
| 157 | + | ||
| 158 | + return CompareGoldenValIdx<TVal, TIdx>(dstValByteSize); | ||
| 159 | + } | ||
| 94 | }; | 160 | }; |
| 95 | 161 | ||
| 96 | TEST_F(TCOLCMAXTest, case01) | 162 | TEST_F(TCOLCMAXTest, case01) |
| @@ -188,3 +254,99 @@ TEST_F(TCOLCMAXTest, case93) | |||
| 188 | bool ret = TCOLCMAXTestFramework<93, uint16_t, 4, 4, 1, 48, 34>(); | 254 | bool ret = TCOLCMAXTestFramework<93, uint16_t, 4, 4, 1, 48, 34>(); |
| 189 | EXPECT_TRUE(ret); | 255 | EXPECT_TRUE(ret); |
| 190 | } | 256 | } |
| 257 | + | ||
| 258 | +TEST_F(TCOLCMAXTest, case001) | ||
| 259 | +{ | ||
| 260 | + bool ret = TCOLCMAXTestFramework<1, float, uint32_t, 1, 1, 1, 256, 255>(); | ||
| 261 | + EXPECT_TRUE(ret); | ||
| 262 | +} | ||
| 263 | +TEST_F(TCOLCMAXTest, case002) | ||
| 264 | +{ | ||
| 265 | + bool ret = TCOLCMAXTestFramework<2, float, uint32_t, 16, 16, 1, 128, 127>(); | ||
| 266 | + EXPECT_TRUE(ret); | ||
| 267 | +} | ||
| 268 | +TEST_F(TCOLCMAXTest, case003) | ||
| 269 | +{ | ||
| 270 | + bool ret = TCOLCMAXTestFramework<3, float, uint32_t, 16, 15, 1, 256, 255>(); | ||
| 271 | + EXPECT_TRUE(ret); | ||
| 272 | +} | ||
| 273 | +TEST_F(TCOLCMAXTest, case011) | ||
| 274 | +{ | ||
| 275 | + bool ret = TCOLCMAXTestFramework<11, aclFloat16, uint16_t, 1, 1, 1, 256, 255>(); | ||
| 276 | + EXPECT_TRUE(ret); | ||
| 277 | +} | ||
| 278 | +TEST_F(TCOLCMAXTest, case012) | ||
| 279 | +{ | ||
| 280 | + bool ret = TCOLCMAXTestFramework<12, aclFloat16, uint16_t, 16, 16, 1, 128, 127>(); | ||
| 281 | + EXPECT_TRUE(ret); | ||
| 282 | +} | ||
| 283 | +TEST_F(TCOLCMAXTest, case013) | ||
| 284 | +{ | ||
| 285 | + bool ret = TCOLCMAXTestFramework<13, aclFloat16, uint16_t, 16, 15, 1, 256, 255>(); | ||
| 286 | + EXPECT_TRUE(ret); | ||
| 287 | +} | ||
| 288 | +TEST_F(TCOLCMAXTest, case051) | ||
| 289 | +{ | ||
| 290 | + bool ret = TCOLCMAXTestFramework<51, uint16_t, uint16_t, 1, 1, 1, 256, 255>(); | ||
| 291 | + EXPECT_TRUE(ret); | ||
| 292 | +} | ||
| 293 | +TEST_F(TCOLCMAXTest, case052) | ||
| 294 | +{ | ||
| 295 | + bool ret = TCOLCMAXTestFramework<52, uint16_t, uint16_t, 16, 16, 1, 128, 127>(); | ||
| 296 | + EXPECT_TRUE(ret); | ||
| 297 | +} | ||
| 298 | +TEST_F(TCOLCMAXTest, case053) | ||
| 299 | +{ | ||
| 300 | + bool ret = TCOLCMAXTestFramework<53, uint16_t, uint16_t, 16, 15, 1, 256, 255>(); | ||
| 301 | + EXPECT_TRUE(ret); | ||
| 302 | +} | ||
| 303 | +TEST_F(TCOLCMAXTest, case071) | ||
| 304 | +{ | ||
| 305 | + bool ret = TCOLCMAXTestFramework<71, uint32_t, uint32_t, 1, 1, 1, 256, 255>(); | ||
| 306 | + EXPECT_TRUE(ret); | ||
| 307 | +} | ||
| 308 | +TEST_F(TCOLCMAXTest, case072) | ||
| 309 | +{ | ||
| 310 | + bool ret = TCOLCMAXTestFramework<72, uint32_t, uint32_t, 16, 16, 1, 128, 127>(); | ||
| 311 | + EXPECT_TRUE(ret); | ||
| 312 | +} | ||
| 313 | +TEST_F(TCOLCMAXTest, case073) | ||
| 314 | +{ | ||
| 315 | + bool ret = TCOLCMAXTestFramework<73, uint32_t, uint32_t, 16, 15, 1, 256, 255>(); | ||
| 316 | + EXPECT_TRUE(ret); | ||
| 317 | +} | ||
| 318 | +TEST_F(TCOLCMAXTest, case081) | ||
| 319 | +{ | ||
| 320 | + bool ret = TCOLCMAXTestFramework<81, aclFloat16, uint16_t, 16, 16, 1, 32, 32>(); | ||
| 321 | + EXPECT_TRUE(ret); | ||
| 322 | +} | ||
| 323 | +TEST_F(TCOLCMAXTest, case082) | ||
| 324 | +{ | ||
| 325 | + bool ret = TCOLCMAXTestFramework<82, uint16_t, uint16_t, 16, 16, 1, 32, 32>(); | ||
| 326 | + EXPECT_TRUE(ret); | ||
| 327 | +} | ||
| 328 | +TEST_F(TCOLCMAXTest, case083) | ||
| 329 | +{ | ||
| 330 | + bool ret = TCOLCMAXTestFramework<83, uint32_t, uint32_t, 16, 16, 1, 32, 31>(); | ||
| 331 | + EXPECT_TRUE(ret); | ||
| 332 | +} | ||
| 333 | +TEST_F(TCOLCMAXTest, case084) | ||
| 334 | +{ | ||
| 335 | + bool ret = TCOLCMAXTestFramework<84, float, uint32_t, 16, 16, 1, 32, 31>(); | ||
| 336 | + EXPECT_TRUE(ret); | ||
| 337 | +} | ||
| 338 | +TEST_F(TCOLCMAXTest, case091) | ||
| 339 | +{ | ||
| 340 | + bool ret = TCOLCMAXTestFramework<91, uint16_t, uint16_t, 16, 16, 1, 128, 120>(); | ||
| 341 | + EXPECT_TRUE(ret); | ||
| 342 | +} | ||
| 343 | +TEST_F(TCOLCMAXTest, case092) | ||
| 344 | +{ | ||
| 345 | + bool ret = TCOLCMAXTestFramework<92, aclFloat16, uint16_t, 16, 16, 1, 96, 88>(); | ||
| 346 | + EXPECT_TRUE(ret); | ||
| 347 | +} | ||
| 348 | +TEST_F(TCOLCMAXTest, case093) | ||
| 349 | +{ | ||
| 350 | + bool ret = TCOLCMAXTestFramework<93, uint16_t, uint16_t, 4, 4, 1, 48, 34>(); | ||
| 351 | + EXPECT_TRUE(ret); | ||
| 352 | +} | ||
| @@ -48,6 +48,51 @@ PTO_INTERNAL void runTColCMax(__gm__ uint32_t __out__ *out, __gm__ T __in__ *src | |||
| 48 | out = dstGlobal.data(); | 48 | out = dstGlobal.data(); |
| 49 | } | 49 | } |
| 50 | 50 | ||
| 51 | +template <typename TVal, typename TIdx, int srcRow, int srcValidRow, int dstRow, int col, int validCol> | ||
| 52 | +PTO_INTERNAL void runTColIdxValMax(__gm__ TVal __out__ *outVal, __gm__ TIdx __out__ *outIdx, __gm__ TVal __in__ *src) | ||
| 53 | +{ | ||
| 54 | + using SrcShapeDim5 = Shape<1, 1, 1, -1, -1>; | ||
| 55 | + using srcStridDim5 = Stride<1, 1, -1, -1, 1>; | ||
| 56 | + using dstShapeDim5 = Shape<1, 1, 1, -1, -1>; | ||
| 57 | + using dstStridDim5 = Stride<1, 1, -1, -1, 1>; | ||
| 58 | + | ||
| 59 | + using GlobalDataSrc = GlobalTensor<TVal, SrcShapeDim5, srcStridDim5>; | ||
| 60 | + using GlobalDataDstVal = GlobalTensor<TVal, dstShapeDim5, dstStridDim5>; | ||
| 61 | + using GlobalDataDstIdx = GlobalTensor<TIdx, dstShapeDim5, dstStridDim5>; | ||
| 62 | + | ||
| 63 | + GlobalDataSrc srcGlobal(src, SrcShapeDim5(srcValidRow, validCol), srcStridDim5(srcRow, col)); | ||
| 64 | + GlobalDataDstIdx dstIdxGlobal(outIdx, dstShapeDim5(dstRow, validCol), dstStridDim5(srcRow, col)); | ||
| 65 | + GlobalDataDstVal dstValGlobal(outVal, dstShapeDim5(dstRow, validCol), dstStridDim5(srcRow, col)); | ||
| 66 | + | ||
| 67 | + using SrcTileData = Tile<TileType::Vec, TVal, srcRow, col, BLayout::RowMajor, -1, -1>; | ||
| 68 | + using DstIdxTileData = Tile<TileType::Vec, TIdx, dstRow, col, BLayout::RowMajor, -1, -1>; | ||
| 69 | + using DstValTileData = Tile<TileType::Vec, TVal, dstRow, col, BLayout::RowMajor, -1, -1>; | ||
| 70 | + using TmpTile = Tile<TileType::Vec, TVal, 1, 32, BLayout::RowMajor, -1, -1>; | ||
| 71 | + | ||
| 72 | + SrcTileData srcTile(srcValidRow, validCol); | ||
| 73 | + DstIdxTileData idxTile(dstRow, validCol); | ||
| 74 | + DstValTileData valTile(dstRow, validCol); | ||
| 75 | + TmpTile tmpTile(1, 32); | ||
| 76 | + | ||
| 77 | + TASSIGN(srcTile, 0x0); | ||
| 78 | + TASSIGN(idxTile, srcRow * col * sizeof(TVal)); | ||
| 79 | + TASSIGN(valTile, (srcRow + 1) * col * sizeof(TVal)); | ||
| 80 | + TASSIGN(tmpTile, (srcRow + 2) * col * sizeof(TVal)); | ||
| 81 | + | ||
| 82 | + TLOAD(srcTile, srcGlobal); | ||
| 83 | + | ||
| 84 | + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); | ||
| 85 | + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); | ||
| 86 | + TCOLARGMAX(valTile, idxTile, srcTile, tmpTile); | ||
| 87 | + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); | ||
| 88 | + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); | ||
| 89 | + TSTORE(dstValGlobal, valTile); | ||
| 90 | + TSTORE(dstIdxGlobal, idxTile); | ||
| 91 | + | ||
| 92 | + outVal = dstValGlobal.data(); | ||
| 93 | + outIdx = dstIdxGlobal.data(); | ||
| 94 | +} | ||
| 95 | + | ||
| 51 | extern "C" __global__ AICORE void launchTCOLCMAXCase01(__gm__ uint32_t *out, __gm__ float *src) | 96 | extern "C" __global__ AICORE void launchTCOLCMAXCase01(__gm__ uint32_t *out, __gm__ float *src) |
| 52 | { | 97 | { |
| 53 | runTColCMax<float, 1, 1, 1, 256, 255>(out, src, false); | 98 | runTColCMax<float, 1, 1, 1, 256, 255>(out, src, false); |
| @@ -125,6 +170,102 @@ extern "C" __global__ AICORE void launchTCOLCMAXCase93(__gm__ uint32_t *out, __g | |||
| 125 | runTColCMax<uint16_t, 4, 4, 1, 48, 34>(out, src, false); | 170 | runTColCMax<uint16_t, 4, 4, 1, 48, 34>(out, src, false); |
| 126 | } | 171 | } |
| 127 | 172 | ||
| 173 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase01(__gm__ float *outVal, __gm__ uint32_t *outIdx, | ||
| 174 | + __gm__ float *src) | ||
| 175 | +{ | ||
| 176 | + runTColIdxValMax<float, uint32_t, 1, 1, 1, 256, 255>(outVal, outIdx, src); | ||
| 177 | +} | ||
| 178 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase02(__gm__ float *outVal, __gm__ uint32_t *outIdx, | ||
| 179 | + __gm__ float *src) | ||
| 180 | +{ | ||
| 181 | + runTColIdxValMax<float, uint32_t, 16, 16, 1, 128, 127>(outVal, outIdx, src); | ||
| 182 | +} | ||
| 183 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase03(__gm__ float *outVal, __gm__ uint32_t *outIdx, | ||
| 184 | + __gm__ float *src) | ||
| 185 | +{ | ||
| 186 | + runTColIdxValMax<float, uint32_t, 16, 15, 1, 256, 255>(outVal, outIdx, src); | ||
| 187 | +} | ||
| 188 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase11(__gm__ half *outVal, __gm__ uint16_t *outIdx, | ||
| 189 | + __gm__ half *src) | ||
| 190 | +{ | ||
| 191 | + runTColIdxValMax<half, uint16_t, 1, 1, 1, 256, 255>(outVal, outIdx, src); | ||
| 192 | +} | ||
| 193 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase12(__gm__ half *outVal, __gm__ uint16_t *outIdx, | ||
| 194 | + __gm__ half *src) | ||
| 195 | +{ | ||
| 196 | + runTColIdxValMax<half, uint16_t, 16, 16, 1, 128, 127>(outVal, outIdx, src); | ||
| 197 | +} | ||
| 198 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase13(__gm__ half *outVal, __gm__ uint16_t *outIdx, | ||
| 199 | + __gm__ half *src) | ||
| 200 | +{ | ||
| 201 | + runTColIdxValMax<half, uint16_t, 16, 15, 1, 256, 255>(outVal, outIdx, src); | ||
| 202 | +} | ||
| 203 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase51(__gm__ uint16_t *outVal, __gm__ uint16_t *outIdx, | ||
| 204 | + __gm__ uint16_t *src) | ||
| 205 | +{ | ||
| 206 | + runTColIdxValMax<uint16_t, uint16_t, 1, 1, 1, 256, 255>(outVal, outIdx, src); | ||
| 207 | +} | ||
| 208 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase52(__gm__ uint16_t *outVal, __gm__ uint16_t *outIdx, | ||
| 209 | + __gm__ uint16_t *src) | ||
| 210 | +{ | ||
| 211 | + runTColIdxValMax<uint16_t, uint16_t, 16, 16, 1, 128, 127>(outVal, outIdx, src); | ||
| 212 | +} | ||
| 213 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase53(__gm__ uint16_t *outVal, __gm__ uint16_t *outIdx, | ||
| 214 | + __gm__ uint16_t *src) | ||
| 215 | +{ | ||
| 216 | + runTColIdxValMax<uint16_t, uint16_t, 16, 15, 1, 256, 255>(outVal, outIdx, src); | ||
| 217 | +} | ||
| 218 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase71(__gm__ uint32_t *outVal, __gm__ uint32_t *outIdx, | ||
| 219 | + __gm__ uint32_t *src) | ||
| 220 | +{ | ||
| 221 | + runTColIdxValMax<uint32_t, uint32_t, 1, 1, 1, 256, 255>(outVal, outIdx, src); | ||
| 222 | +} | ||
| 223 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase72(__gm__ uint32_t *outVal, __gm__ uint32_t *outIdx, | ||
| 224 | + __gm__ uint32_t *src) | ||
| 225 | +{ | ||
| 226 | + runTColIdxValMax<uint32_t, uint32_t, 16, 16, 1, 128, 127>(outVal, outIdx, src); | ||
| 227 | +} | ||
| 228 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase73(__gm__ uint32_t *outVal, __gm__ uint32_t *outIdx, | ||
| 229 | + __gm__ uint32_t *src) | ||
| 230 | +{ | ||
| 231 | + runTColIdxValMax<uint32_t, uint32_t, 16, 15, 1, 256, 255>(outVal, outIdx, src); | ||
| 232 | +} | ||
| 233 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase81(__gm__ half *outVal, __gm__ uint16_t *outIdx, | ||
| 234 | + __gm__ half *src) | ||
| 235 | +{ | ||
| 236 | + runTColIdxValMax<half, uint16_t, 16, 16, 1, 32, 32>(outVal, outIdx, src); | ||
| 237 | +} | ||
| 238 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase82(__gm__ uint16_t *outVal, __gm__ uint16_t *outIdx, | ||
| 239 | + __gm__ uint16_t *src) | ||
| 240 | +{ | ||
| 241 | + runTColIdxValMax<uint16_t, uint16_t, 16, 16, 1, 32, 32>(outVal, outIdx, src); | ||
| 242 | +} | ||
| 243 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase83(__gm__ uint32_t *outVal, __gm__ uint32_t *outIdx, | ||
| 244 | + __gm__ uint32_t *src) | ||
| 245 | +{ | ||
| 246 | + runTColIdxValMax<uint32_t, uint32_t, 16, 16, 1, 32, 31>(outVal, outIdx, src); | ||
| 247 | +} | ||
| 248 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase84(__gm__ float *outVal, __gm__ uint32_t *outIdx, | ||
| 249 | + __gm__ float *src) | ||
| 250 | +{ | ||
| 251 | + runTColIdxValMax<float, uint32_t, 16, 16, 1, 32, 31>(outVal, outIdx, src); | ||
| 252 | +} | ||
| 253 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase91(__gm__ uint16_t *outVal, __gm__ uint16_t *outIdx, | ||
| 254 | + __gm__ uint16_t *src) | ||
| 255 | +{ | ||
| 256 | + runTColIdxValMax<uint16_t, uint16_t, 16, 16, 1, 128, 120>(outVal, outIdx, src); | ||
| 257 | +} | ||
| 258 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase92(__gm__ half *outVal, __gm__ uint16_t *outIdx, | ||
| 259 | + __gm__ half *src) | ||
| 260 | +{ | ||
| 261 | + runTColIdxValMax<half, uint16_t, 16, 16, 1, 96, 88>(outVal, outIdx, src); | ||
| 262 | +} | ||
| 263 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase93(__gm__ uint16_t *outVal, __gm__ uint16_t *outIdx, | ||
| 264 | + __gm__ uint16_t *src) | ||
| 265 | +{ | ||
| 266 | + runTColIdxValMax<uint16_t, uint16_t, 4, 4, 1, 48, 34>(outVal, outIdx, src); | ||
| 267 | +} | ||
| 268 | + | ||
| 128 | template <uint32_t caseId> | 269 | template <uint32_t caseId> |
| 129 | void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream) | 270 | void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream) |
| 130 | { | 271 | { |
| @@ -210,6 +351,91 @@ void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream) | |||
| 210 | } | 351 | } |
| 211 | } | 352 | } |
| 212 | 353 | ||
| 354 | +template <uint32_t caseId> | ||
| 355 | +void launchTCOLIDXVALMAXCase(void *outVal, void *outIdx, void *src, aclrtStream stream) | ||
| 356 | +{ | ||
| 357 | + switch (caseId) { | ||
| 358 | + case 1: { | ||
| 359 | + launchTCOLIDXVALMAXCase01<<<1, nullptr, stream>>>((float *)outVal, (uint32_t *)outIdx, (float *)src); | ||
| 360 | + break; | ||
| 361 | + } | ||
| 362 | + case 2: { | ||
| 363 | + launchTCOLIDXVALMAXCase02<<<1, nullptr, stream>>>((float *)outVal, (uint32_t *)outIdx, (float *)src); | ||
| 364 | + break; | ||
| 365 | + } | ||
| 366 | + case 3: { | ||
| 367 | + launchTCOLIDXVALMAXCase03<<<1, nullptr, stream>>>((float *)outVal, (uint32_t *)outIdx, (float *)src); | ||
| 368 | + break; | ||
| 369 | + } | ||
| 370 | + case 11: { | ||
| 371 | + launchTCOLIDXVALMAXCase11<<<1, nullptr, stream>>>((half *)outVal, (uint16_t *)outIdx, (half *)src); | ||
| 372 | + break; | ||
| 373 | + } | ||
| 374 | + case 12: { | ||
| 375 | + launchTCOLIDXVALMAXCase12<<<1, nullptr, stream>>>((half *)outVal, (uint16_t *)outIdx, (half *)src); | ||
| 376 | + break; | ||
| 377 | + } | ||
| 378 | + case 13: { | ||
| 379 | + launchTCOLIDXVALMAXCase13<<<1, nullptr, stream>>>((half *)outVal, (uint16_t *)outIdx, (half *)src); | ||
| 380 | + break; | ||
| 381 | + } | ||
| 382 | + case 51: { | ||
| 383 | + launchTCOLIDXVALMAXCase51<<<1, nullptr, stream>>>((uint16_t *)outVal, (uint16_t *)outIdx, (uint16_t *)src); | ||
| 384 | + break; | ||
| 385 | + } | ||
| 386 | + case 52: { | ||
| 387 | + launchTCOLIDXVALMAXCase52<<<1, nullptr, stream>>>((uint16_t *)outVal, (uint16_t *)outIdx, (uint16_t *)src); | ||
| 388 | + break; | ||
| 389 | + } | ||
| 390 | + case 53: { | ||
| 391 | + launchTCOLIDXVALMAXCase53<<<1, nullptr, stream>>>((uint16_t *)outVal, (uint16_t *)outIdx, (uint16_t *)src); | ||
| 392 | + break; | ||
| 393 | + } | ||
| 394 | + case 71: { | ||
| 395 | + launchTCOLIDXVALMAXCase71<<<1, nullptr, stream>>>((uint32_t *)outVal, (uint32_t *)outIdx, (uint32_t *)src); | ||
| 396 | + break; | ||
| 397 | + } | ||
| 398 | + case 72: { | ||
| 399 | + launchTCOLIDXVALMAXCase72<<<1, nullptr, stream>>>((uint32_t *)outVal, (uint32_t *)outIdx, (uint32_t *)src); | ||
| 400 | + break; | ||
| 401 | + } | ||
| 402 | + case 73: { | ||
| 403 | + launchTCOLIDXVALMAXCase73<<<1, nullptr, stream>>>((uint32_t *)outVal, (uint32_t *)outIdx, (uint32_t *)src); | ||
| 404 | + break; | ||
| 405 | + } | ||
| 406 | + case 81: { | ||
| 407 | + launchTCOLIDXVALMAXCase81<<<1, nullptr, stream>>>((half *)outVal, (uint16_t *)outIdx, (half *)src); | ||
| 408 | + break; | ||
| 409 | + } | ||
| 410 | + case 82: { | ||
| 411 | + launchTCOLIDXVALMAXCase82<<<1, nullptr, stream>>>((uint16_t *)outVal, (uint16_t *)outIdx, (uint16_t *)src); | ||
| 412 | + break; | ||
| 413 | + } | ||
| 414 | + case 83: { | ||
| 415 | + launchTCOLIDXVALMAXCase83<<<1, nullptr, stream>>>((uint32_t *)outVal, (uint32_t *)outIdx, (uint32_t *)src); | ||
| 416 | + break; | ||
| 417 | + } | ||
| 418 | + case 84: { | ||
| 419 | + launchTCOLIDXVALMAXCase84<<<1, nullptr, stream>>>((float *)outVal, (uint32_t *)outIdx, (float *)src); | ||
| 420 | + break; | ||
| 421 | + } | ||
| 422 | + case 91: { | ||
| 423 | + launchTCOLIDXVALMAXCase91<<<1, nullptr, stream>>>((uint16_t *)outVal, (uint16_t *)outIdx, (uint16_t *)src); | ||
| 424 | + break; | ||
| 425 | + } | ||
| 426 | + case 92: { | ||
| 427 | + launchTCOLIDXVALMAXCase92<<<1, nullptr, stream>>>((half *)outVal, (uint16_t *)outIdx, (half *)src); | ||
| 428 | + break; | ||
| 429 | + } | ||
| 430 | + case 93: { | ||
| 431 | + launchTCOLIDXVALMAXCase93<<<1, nullptr, stream>>>((uint16_t *)outVal, (uint16_t *)outIdx, (uint16_t *)src); | ||
| 432 | + break; | ||
| 433 | + } | ||
| 434 | + default: { | ||
| 435 | + } | ||
| 436 | + } | ||
| 437 | +} | ||
| 438 | + | ||
| 213 | template void launchTCOLCMAXTestCase<1>(void *out, void *src, aclrtStream stream); | 439 | template void launchTCOLCMAXTestCase<1>(void *out, void *src, aclrtStream stream); |
| 214 | template void launchTCOLCMAXTestCase<2>(void *out, void *src, aclrtStream stream); | 440 | template void launchTCOLCMAXTestCase<2>(void *out, void *src, aclrtStream stream); |
| 215 | template void launchTCOLCMAXTestCase<3>(void *out, void *src, aclrtStream stream); | 441 | template void launchTCOLCMAXTestCase<3>(void *out, void *src, aclrtStream stream); |
| @@ -228,4 +454,24 @@ template void launchTCOLCMAXTestCase<83>(void *out, void *src, aclrtStream strea | |||
| 228 | template void launchTCOLCMAXTestCase<84>(void *out, void *src, aclrtStream stream); | 454 | template void launchTCOLCMAXTestCase<84>(void *out, void *src, aclrtStream stream); |
| 229 | template void launchTCOLCMAXTestCase<91>(void *out, void *src, aclrtStream stream); | 455 | template void launchTCOLCMAXTestCase<91>(void *out, void *src, aclrtStream stream); |
| 230 | template void launchTCOLCMAXTestCase<92>(void *out, void *src, aclrtStream stream); | 456 | template void launchTCOLCMAXTestCase<92>(void *out, void *src, aclrtStream stream); |
| 231 | -template void launchTCOLCMAXTestCase<93>(void *out, void *src, aclrtStream stream); | 457 | +template void launchTCOLCMAXTestCase<93>(void *out, void *src, aclrtStream stream); |
| 458 | + | ||
| 459 | +template void launchTCOLIDXVALMAXCase<1>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 460 | +template void launchTCOLIDXVALMAXCase<2>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 461 | +template void launchTCOLIDXVALMAXCase<3>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 462 | +template void launchTCOLIDXVALMAXCase<11>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 463 | +template void launchTCOLIDXVALMAXCase<12>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 464 | +template void launchTCOLIDXVALMAXCase<13>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 465 | +template void launchTCOLIDXVALMAXCase<51>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 466 | +template void launchTCOLIDXVALMAXCase<52>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 467 | +template void launchTCOLIDXVALMAXCase<53>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 468 | +template void launchTCOLIDXVALMAXCase<71>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 469 | +template void launchTCOLIDXVALMAXCase<72>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 470 | +template void launchTCOLIDXVALMAXCase<73>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 471 | +template void launchTCOLIDXVALMAXCase<81>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 472 | +template void launchTCOLIDXVALMAXCase<82>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 473 | +template void launchTCOLIDXVALMAXCase<83>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 474 | +template void launchTCOLIDXVALMAXCase<84>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 475 | +template void launchTCOLIDXVALMAXCase<91>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 476 | +template void launchTCOLIDXVALMAXCase<92>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 477 | +template void launchTCOLIDXVALMAXCase<93>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| @@ -1,17 +1,19 @@ | |||
| 1 | #!/usr/bin/python3 | 1 | #!/usr/bin/python3 |
| 2 | # coding=utf-8 | 2 | # coding=utf-8 |
| 3 | -# -------------------------------------------------------------------------------- | 3 | +""" |
| 4 | -# Copyright (c) 2025 Huawei Technologies Co., Ltd. | 4 | +Copyright (c) 2025 Huawei Technologies Co., Ltd. |
| 5 | -# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | 5 | +This program is free software, you can redistribute it and/or modify it under the terms and conditions of |
| 6 | -# CANN Open Software License Agreement Version 2.0 (the "License"). | 6 | +CANN Open Software License Agreement Version 2.0 (the "License"). |
| 7 | -# Please refer to the License for details. You may not use this file except in compliance with the License. | 7 | +Please refer to the License for details. You may not use this file except in compliance with the License. |
| 8 | -# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | 8 | +THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, |
| 9 | -# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | 9 | +INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. |
| 10 | -# See LICENSE in the root of the software repository for the full text of the License. | 10 | +See LICENSE in the root of the software repository for the full text of the License. |
| 11 | -# -------------------------------------------------------------------------------- | 11 | +""" |
| 12 | 12 | ||
| 13 | import os | 13 | import os |
| 14 | +import math | ||
| 14 | import numpy as np | 15 | import numpy as np |
| 16 | + | ||
| 15 | np.random.seed(19) | 17 | np.random.seed(19) |
| 16 | 18 | ||
| 17 | 19 | ||
| @@ -23,44 +25,146 @@ def gen_golden_data(param): | |||
| 23 | valid_col = param.valid_col | 25 | valid_col = param.valid_col |
| 24 | value_max = 100 | 26 | value_max = 100 |
| 25 | value_min = -100 | 27 | value_min = -100 |
| 26 | - if data_type == np.uint16 or data_type == np.uint32: | 28 | + if data_type in (np.uint16, np.uint32, np.uint8): |
| 27 | value_max = 200 | 29 | value_max = 200 |
| 28 | value_min = 0 | 30 | value_min = 0 |
| 31 | + if data_type == np.int8: | ||
| 32 | + value_max = 10 | ||
| 33 | + value_min = -10 | ||
| 29 | if data_type == np.uint8: | 34 | if data_type == np.uint8: |
| 30 | value_max = 10 | 35 | value_max = 10 |
| 31 | value_min = 0 | 36 | value_min = 0 |
| 32 | input_arr = np.random.uniform(low=value_min, high=value_max, size=(row, col)).astype(data_type) | 37 | input_arr = np.random.uniform(low=value_min, high=value_max, size=(row, col)).astype(data_type) |
| 33 | - output_arr = np.argmax(input_arr[0:valid_row], axis=0) | 38 | + |
| 34 | - output_arr[valid_col:] = 0 | 39 | + if not param.idx: |
| 35 | - # 先计算, 再强转类型, 保证结果精度不裂化 | 40 | + # Pure index output mode |
| 36 | - output_arr = output_arr.astype(np.int32) | 41 | + output_arr = np.argmax(input_arr[0:valid_row], axis=0) |
| 37 | - input_arr.tofile('input.bin') | 42 | + output_arr[valid_col:] = 0 |
| 38 | - output_arr.tofile('golden.bin') | 43 | + dst_col = math.ceil(valid_col / 8) * 8 |
| 44 | + output_arr = output_arr[:dst_col] | ||
| 45 | + output_arr = output_arr.astype(np.int32) | ||
| 46 | + input_arr.tofile("input.bin") | ||
| 47 | + output_arr.tofile("golden.bin") | ||
| 48 | + else: | ||
| 49 | + # Value + index output mode | ||
| 50 | + input_arr.tofile("input.bin") | ||
| 51 | + dst_col = math.ceil(valid_col / 8) * 8 | ||
| 52 | + output_idx = np.argmax(input_arr[0:valid_row], axis=0) | ||
| 53 | + output_val = np.max(input_arr[0:valid_row], axis=0) | ||
| 54 | + output_val[valid_col:] = 0 | ||
| 55 | + output_idx[valid_col:] = 0 | ||
| 56 | + output_val = output_val[:dst_col] | ||
| 57 | + output_idx = output_idx[:dst_col] | ||
| 58 | + if input_arr.itemsize == 2: | ||
| 59 | + output_idx = output_idx.astype(np.int16) | ||
| 60 | + else: | ||
| 61 | + output_idx = output_idx.astype(np.int32) | ||
| 62 | + output_idx.tofile("idx.bin") | ||
| 63 | + output_val.tofile("golden.bin") | ||
| 39 | 64 | ||
| 40 | 65 | ||
| 41 | class TColCMaxParams: | 66 | class TColCMaxParams: |
| 42 | - def __init__(self, name, data_type, row, valid_row, col, valid_col): | 67 | + def __init__(self, name, data_type, row, valid_row, col, valid_col, idx=False): |
| 43 | self.name = name | 68 | self.name = name |
| 44 | self.data_type = data_type | 69 | self.data_type = data_type |
| 45 | self.row = row | 70 | self.row = row |
| 46 | self.valid_row = valid_row | 71 | self.valid_row = valid_row |
| 47 | self.col = col | 72 | self.col = col |
| 48 | self.valid_col = valid_col | 73 | self.valid_col = valid_col |
| 74 | + self.idx = idx | ||
| 75 | + | ||
| 49 | 76 | ||
| 50 | if __name__ == "__main__": | 77 | if __name__ == "__main__": |
| 51 | case_params_list = [ | 78 | case_params_list = [ |
| 79 | + # ========================================================================= | ||
| 80 | + # Pure index mode (TCOLARGMAX with 3 args): all 8 supported types x 3 dims | ||
| 81 | + # ========================================================================= | ||
| 82 | + # float32 | ||
| 52 | TColCMaxParams("TCOLCMAXTest.case01", np.float32, 1, 1, 256, 255), | 83 | TColCMaxParams("TCOLCMAXTest.case01", np.float32, 1, 1, 256, 255), |
| 53 | TColCMaxParams("TCOLCMAXTest.case02", np.float32, 16, 16, 128, 127), | 84 | TColCMaxParams("TCOLCMAXTest.case02", np.float32, 16, 16, 128, 127), |
| 54 | TColCMaxParams("TCOLCMAXTest.case03", np.float32, 16, 15, 256, 255), | 85 | TColCMaxParams("TCOLCMAXTest.case03", np.float32, 16, 15, 256, 255), |
| 86 | + # float16 | ||
| 55 | TColCMaxParams("TCOLCMAXTest.case11", np.float16, 1, 1, 256, 255), | 87 | TColCMaxParams("TCOLCMAXTest.case11", np.float16, 1, 1, 256, 255), |
| 56 | TColCMaxParams("TCOLCMAXTest.case12", np.float16, 16, 16, 128, 127), | 88 | TColCMaxParams("TCOLCMAXTest.case12", np.float16, 16, 16, 128, 127), |
| 57 | TColCMaxParams("TCOLCMAXTest.case13", np.float16, 16, 15, 256, 255), | 89 | TColCMaxParams("TCOLCMAXTest.case13", np.float16, 16, 15, 256, 255), |
| 90 | + # int8 | ||
| 91 | + TColCMaxParams("TCOLCMAXTest.case21", np.int8, 1, 1, 256, 255), | ||
| 92 | + TColCMaxParams("TCOLCMAXTest.case22", np.int8, 16, 16, 128, 127), | ||
| 93 | + TColCMaxParams("TCOLCMAXTest.case23", np.int8, 16, 15, 256, 255), | ||
| 94 | + # uint8 | ||
| 95 | + TColCMaxParams("TCOLCMAXTest.case31", np.uint8, 1, 1, 256, 255), | ||
| 96 | + TColCMaxParams("TCOLCMAXTest.case32", np.uint8, 16, 16, 128, 127), | ||
| 97 | + TColCMaxParams("TCOLCMAXTest.case33", np.uint8, 16, 15, 256, 255), | ||
| 98 | + # int16 | ||
| 99 | + TColCMaxParams("TCOLCMAXTest.case41", np.int16, 1, 1, 256, 255), | ||
| 100 | + TColCMaxParams("TCOLCMAXTest.case42", np.int16, 16, 16, 128, 127), | ||
| 101 | + TColCMaxParams("TCOLCMAXTest.case43", np.int16, 16, 15, 256, 255), | ||
| 102 | + # uint16 | ||
| 58 | TColCMaxParams("TCOLCMAXTest.case51", np.uint16, 1, 1, 256, 255), | 103 | TColCMaxParams("TCOLCMAXTest.case51", np.uint16, 1, 1, 256, 255), |
| 59 | TColCMaxParams("TCOLCMAXTest.case52", np.uint16, 16, 16, 128, 127), | 104 | TColCMaxParams("TCOLCMAXTest.case52", np.uint16, 16, 16, 128, 127), |
| 60 | TColCMaxParams("TCOLCMAXTest.case53", np.uint16, 16, 15, 256, 255), | 105 | TColCMaxParams("TCOLCMAXTest.case53", np.uint16, 16, 15, 256, 255), |
| 106 | + # int32 | ||
| 107 | + TColCMaxParams("TCOLCMAXTest.case61", np.int32, 1, 1, 256, 255), | ||
| 108 | + TColCMaxParams("TCOLCMAXTest.case62", np.int32, 16, 16, 128, 127), | ||
| 109 | + TColCMaxParams("TCOLCMAXTest.case63", np.int32, 16, 15, 256, 255), | ||
| 110 | + # uint32 | ||
| 61 | TColCMaxParams("TCOLCMAXTest.case71", np.uint32, 1, 1, 256, 255), | 111 | TColCMaxParams("TCOLCMAXTest.case71", np.uint32, 1, 1, 256, 255), |
| 62 | TColCMaxParams("TCOLCMAXTest.case72", np.uint32, 16, 16, 128, 127), | 112 | TColCMaxParams("TCOLCMAXTest.case72", np.uint32, 16, 16, 128, 127), |
| 63 | TColCMaxParams("TCOLCMAXTest.case73", np.uint32, 16, 15, 256, 255), | 113 | TColCMaxParams("TCOLCMAXTest.case73", np.uint32, 16, 15, 256, 255), |
| 114 | + # ========================================================================= | ||
| 115 | + # Pure index mode -- small dimension edge cases | ||
| 116 | + # ========================================================================= | ||
| 117 | + TColCMaxParams("TCOLCMAXTest.case81", np.float16, 16, 16, 32, 32), | ||
| 118 | + TColCMaxParams("TCOLCMAXTest.case82", np.uint16, 16, 16, 32, 32), | ||
| 119 | + TColCMaxParams("TCOLCMAXTest.case83", np.uint32, 16, 16, 32, 31), | ||
| 120 | + TColCMaxParams("TCOLCMAXTest.case84", np.float32, 16, 16, 32, 31), | ||
| 121 | + TColCMaxParams("TCOLCMAXTest.case85", np.int8, 16, 16, 32, 31), | ||
| 122 | + TColCMaxParams("TCOLCMAXTest.case86", np.uint8, 16, 16, 32, 31), | ||
| 123 | + TColCMaxParams("TCOLCMAXTest.case87", np.int16, 16, 16, 32, 31), | ||
| 124 | + TColCMaxParams("TCOLCMAXTest.case88", np.int32, 16, 16, 32, 31), | ||
| 125 | + TColCMaxParams("TCOLCMAXTest.case91", np.uint16, 16, 16, 128, 120), | ||
| 126 | + TColCMaxParams("TCOLCMAXTest.case92", np.float16, 16, 16, 96, 88), | ||
| 127 | + TColCMaxParams("TCOLCMAXTest.case93", np.uint16, 4, 4, 48, 34), | ||
| 128 | + # ========================================================================= | ||
| 129 | + # Value + index mode (TCOLARGMAX with 4 args): 6 types x 3 dims | ||
| 130 | + # Not supported for 8-bit types (Chunk8 does not support WithVal) | ||
| 131 | + # ========================================================================= | ||
| 132 | + # float32 + uint32 index | ||
| 133 | + TColCMaxParams("TCOLCMAXTest.case001", np.float32, 1, 1, 256, 255, True), | ||
| 134 | + TColCMaxParams("TCOLCMAXTest.case002", np.float32, 16, 16, 128, 127, True), | ||
| 135 | + TColCMaxParams("TCOLCMAXTest.case003", np.float32, 16, 15, 256, 255, True), | ||
| 136 | + # float16 + int16 index | ||
| 137 | + TColCMaxParams("TCOLCMAXTest.case011", np.float16, 1, 1, 256, 255, True), | ||
| 138 | + TColCMaxParams("TCOLCMAXTest.case012", np.float16, 16, 16, 128, 127, True), | ||
| 139 | + TColCMaxParams("TCOLCMAXTest.case013", np.float16, 16, 15, 256, 255, True), | ||
| 140 | + # int16 + int16 index | ||
| 141 | + TColCMaxParams("TCOLCMAXTest.case041", np.int16, 1, 1, 256, 255, True), | ||
| 142 | + TColCMaxParams("TCOLCMAXTest.case042", np.int16, 16, 16, 128, 127, True), | ||
| 143 | + TColCMaxParams("TCOLCMAXTest.case043", np.int16, 16, 15, 256, 255, True), | ||
| 144 | + # uint16 + int16 index | ||
| 145 | + TColCMaxParams("TCOLCMAXTest.case051", np.uint16, 1, 1, 256, 255, True), | ||
| 146 | + TColCMaxParams("TCOLCMAXTest.case052", np.uint16, 16, 16, 128, 127, True), | ||
| 147 | + TColCMaxParams("TCOLCMAXTest.case053", np.uint16, 16, 15, 256, 255, True), | ||
| 148 | + # int32 + int32 index | ||
| 149 | + TColCMaxParams("TCOLCMAXTest.case061", np.int32, 1, 1, 256, 255, True), | ||
| 150 | + TColCMaxParams("TCOLCMAXTest.case062", np.int32, 16, 16, 128, 127, True), | ||
| 151 | + TColCMaxParams("TCOLCMAXTest.case063", np.int32, 16, 15, 256, 255, True), | ||
| 152 | + # uint32 + int32 index | ||
| 153 | + TColCMaxParams("TCOLCMAXTest.case071", np.uint32, 1, 1, 256, 255, True), | ||
| 154 | + TColCMaxParams("TCOLCMAXTest.case072", np.uint32, 16, 16, 128, 127, True), | ||
| 155 | + TColCMaxParams("TCOLCMAXTest.case073", np.uint32, 16, 15, 256, 255, True), | ||
| 156 | + # ========================================================================= | ||
| 157 | + # Value + index mode -- small dimension edge cases | ||
| 158 | + # ========================================================================= | ||
| 159 | + TColCMaxParams("TCOLCMAXTest.case081", np.float16, 16, 16, 32, 32, True), | ||
| 160 | + TColCMaxParams("TCOLCMAXTest.case082", np.uint16, 16, 16, 32, 32, True), | ||
| 161 | + TColCMaxParams("TCOLCMAXTest.case083", np.uint32, 16, 16, 32, 31, True), | ||
| 162 | + TColCMaxParams("TCOLCMAXTest.case084", np.float32, 16, 16, 32, 31, True), | ||
| 163 | + TColCMaxParams("TCOLCMAXTest.case085", np.int16, 16, 16, 32, 31, True), | ||
| 164 | + TColCMaxParams("TCOLCMAXTest.case086", np.int32, 16, 16, 32, 31, True), | ||
| 165 | + TColCMaxParams("TCOLCMAXTest.case091", np.uint16, 16, 16, 128, 120, True), | ||
| 166 | + TColCMaxParams("TCOLCMAXTest.case092", np.float16, 16, 16, 96, 88, True), | ||
| 167 | + TColCMaxParams("TCOLCMAXTest.case093", np.uint16, 4, 4, 48, 34, True), | ||
| 64 | ] | 168 | ] |
| 65 | 169 | ||
| 66 | for _, case in enumerate(case_params_list): | 170 | for _, case in enumerate(case_params_list): |
| @@ -69,4 +173,4 @@ if __name__ == "__main__": | |||
| 69 | original_dir = os.getcwd() | 173 | original_dir = os.getcwd() |
| 70 | os.chdir(case.name) | 174 | os.chdir(case.name) |
| 71 | gen_golden_data(case) | 175 | gen_golden_data(case) |
| 72 | - os.chdir(original_dir) | 176 | + os.chdir(original_dir) |
| @@ -15,9 +15,19 @@ See LICENSE in the root of the software repository for the full text of the Lice | |||
| 15 | using namespace std; | 15 | using namespace std; |
| 16 | using namespace PtoTestCommon; | 16 | using namespace PtoTestCommon; |
| 17 | 17 | ||
| 18 | +// ============================================================================= | ||
| 19 | +// Pure index dispatcher (3-arg TCOLARGMAX) | ||
| 20 | +// ============================================================================= | ||
| 18 | template <uint32_t caseId> | 21 | template <uint32_t caseId> |
| 19 | void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream); | 22 | void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream); |
| 20 | 23 | ||
| 24 | +// ============================================================================= | ||
| 25 | +// Value + index dispatcher (4-arg TCOLARGMAX) | ||
| 26 | +// ============================================================================= | ||
| 27 | +template <uint32_t caseId> | ||
| 28 | +void launchTCOLIDXVALMAXCase(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 29 | + | ||
| 30 | +// ============================================================================= | ||
| 21 | std::string GetGoldenDir() | 31 | std::string GetGoldenDir() |
| 22 | { | 32 | { |
| 23 | const testing::TestInfo *testInfo = testing::UnitTest::GetInstance()->current_test_info(); | 33 | const testing::TestInfo *testInfo = testing::UnitTest::GetInstance()->current_test_info(); |
| @@ -27,6 +37,9 @@ std::string GetGoldenDir() | |||
| 27 | return fullPath; | 37 | return fullPath; |
| 28 | } | 38 | } |
| 29 | 39 | ||
| 40 | +// ============================================================================= | ||
| 41 | +// Test fixture: supports both pure index (3-arg) and value+index (4-arg) | ||
| 42 | +// ============================================================================= | ||
| 30 | class TCOLCMAXTest : public testing::Test { | 43 | class TCOLCMAXTest : public testing::Test { |
| 31 | public: | 44 | public: |
| 32 | aclrtStream stream; | 45 | aclrtStream stream; |
| @@ -35,6 +48,11 @@ public: | |||
| 35 | void *dstDevice; | 48 | void *dstDevice; |
| 36 | void *srcDevice; | 49 | void *srcDevice; |
| 37 | 50 | ||
| 51 | + void *dstHostVal; | ||
| 52 | + void *dstHostIdx; | ||
| 53 | + void *dstDeviceVal; | ||
| 54 | + void *dstDeviceIdx; | ||
| 55 | + | ||
| 38 | protected: | 56 | protected: |
| 39 | void SetUp() override | 57 | void SetUp() override |
| 40 | { | 58 | { |
| @@ -50,7 +68,6 @@ protected: | |||
| 50 | aclFinalize(); | 68 | aclFinalize(); |
| 51 | } | 69 | } |
| 52 | 70 | ||
| 53 | - // template <typename T> | ||
| 54 | bool CompareGolden(size_t dstByteSize, bool printAllEn = false) | 71 | bool CompareGolden(size_t dstByteSize, bool printAllEn = false) |
| 55 | { | 72 | { |
| 56 | std::vector<uint32_t> golden(dstByteSize); | 73 | std::vector<uint32_t> golden(dstByteSize); |
| @@ -64,10 +81,35 @@ protected: | |||
| 64 | return ResultCmp(golden, result, eps, 0, 1000, false, true); | 81 | return ResultCmp(golden, result, eps, 0, 1000, false, true); |
| 65 | } | 82 | } |
| 66 | 83 | ||
| 84 | + template <typename TVal, typename TIdx> | ||
| 85 | + bool CompareGoldenValIdx(size_t dstByteSize, bool printAllEn = false) | ||
| 86 | + { | ||
| 87 | + std::vector<TIdx> goldenIdx(dstByteSize); | ||
| 88 | + std::vector<TIdx> resultIdx(dstByteSize); | ||
| 89 | + std::vector<TVal> goldenVal(dstByteSize); | ||
| 90 | + std::vector<TVal> resultVal(dstByteSize); | ||
| 91 | + | ||
| 92 | + float eps = 0.001f; | ||
| 93 | + ReadFile(GetGoldenDir() + "/golden.bin", dstByteSize, goldenVal.data(), dstByteSize); | ||
| 94 | + ReadFile(GetGoldenDir() + "/output_val.bin", dstByteSize, resultVal.data(), dstByteSize); | ||
| 95 | + ReadFile(GetGoldenDir() + "/idx.bin", dstByteSize, goldenIdx.data(), dstByteSize); | ||
| 96 | + ReadFile(GetGoldenDir() + "/output_idx.bin", dstByteSize, resultIdx.data(), dstByteSize); | ||
| 97 | + if (printAllEn) { | ||
| 98 | + return ResultCmp(goldenVal, resultVal, eps, 0, 1000, true) && | ||
| 99 | + ResultCmp(goldenIdx, resultIdx, eps, 0, 1000, true); | ||
| 100 | + } | ||
| 101 | + return ResultCmp(goldenVal, resultVal, eps, 0, 1000, false, true) && | ||
| 102 | + ResultCmp(goldenIdx, resultIdx, eps, 0, 1000, false, true); | ||
| 103 | + } | ||
| 104 | + | ||
| 105 | + // ------------------------------------------------------------------------- | ||
| 106 | + // Pure index framework (3-arg) | ||
| 107 | + // ------------------------------------------------------------------------- | ||
| 67 | template <uint32_t caseId, typename T, int srcRow, int srcValidRow, int dstRow, int col, int validCol> | 108 | template <uint32_t caseId, typename T, int srcRow, int srcValidRow, int dstRow, int col, int validCol> |
| 68 | bool TCOLCMAXTestFramework() | 109 | bool TCOLCMAXTestFramework() |
| 69 | { | 110 | { |
| 70 | - size_t dstByteSize = dstRow * col * sizeof(uint32_t); | 111 | + constexpr int dstCol = (validCol + 7) / 8 * 8; |
| 112 | + size_t dstByteSize = dstRow * dstCol * sizeof(uint32_t); | ||
| 71 | size_t srcByteSize = srcRow * col * sizeof(T); | 113 | size_t srcByteSize = srcRow * col * sizeof(T); |
| 72 | aclrtMallocHost(&dstHost, dstByteSize); | 114 | aclrtMallocHost(&dstHost, dstByteSize); |
| 73 | aclrtMallocHost(&srcHost, srcByteSize); | 115 | aclrtMallocHost(&srcHost, srcByteSize); |
| @@ -90,8 +132,51 @@ protected: | |||
| 90 | 132 | ||
| 91 | return CompareGolden(dstByteSize); | 133 | return CompareGolden(dstByteSize); |
| 92 | } | 134 | } |
| 135 | + | ||
| 136 | + // ------------------------------------------------------------------------- | ||
| 137 | + // Value + index framework (4-arg) | ||
| 138 | + // ------------------------------------------------------------------------- | ||
| 139 | + template <uint32_t caseId, typename TVal, typename TIdx, int srcRow, int srcValidRow, int dstRow, int col, | ||
| 140 | + int validCol> | ||
| 141 | + bool TCOLCMAXTestFramework() | ||
| 142 | + { | ||
| 143 | + constexpr int dstCol = (validCol + 7) / 8 * 8; | ||
| 144 | + size_t dstIdxByteSize = dstRow * dstCol * sizeof(TIdx); | ||
| 145 | + size_t dstValByteSize = dstRow * dstCol * sizeof(TVal); | ||
| 146 | + size_t srcByteSize = srcRow * col * sizeof(TVal); | ||
| 147 | + aclrtMallocHost(&dstHostIdx, dstIdxByteSize); | ||
| 148 | + aclrtMallocHost(&dstHostVal, dstValByteSize); | ||
| 149 | + aclrtMallocHost(&srcHost, srcByteSize); | ||
| 150 | + aclrtMalloc(&dstDeviceIdx, dstIdxByteSize, ACL_MEM_MALLOC_HUGE_FIRST); | ||
| 151 | + aclrtMalloc(&dstDeviceVal, dstValByteSize, ACL_MEM_MALLOC_HUGE_FIRST); | ||
| 152 | + aclrtMalloc(&srcDevice, srcByteSize, ACL_MEM_MALLOC_HUGE_FIRST); | ||
| 153 | + | ||
| 154 | + ReadFile(GetGoldenDir() + "/input.bin", srcByteSize, srcHost, srcByteSize); | ||
| 155 | + aclrtMemcpy(srcDevice, srcByteSize, srcHost, srcByteSize, ACL_MEMCPY_HOST_TO_DEVICE); | ||
| 156 | + | ||
| 157 | + launchTCOLIDXVALMAXCase<caseId>(dstDeviceVal, dstDeviceIdx, srcDevice, stream); | ||
| 158 | + aclrtSynchronizeStream(stream); | ||
| 159 | + | ||
| 160 | + aclrtMemcpy(dstHostIdx, dstIdxByteSize, dstDeviceIdx, dstIdxByteSize, ACL_MEMCPY_DEVICE_TO_HOST); | ||
| 161 | + aclrtMemcpy(dstHostVal, dstValByteSize, dstDeviceVal, dstValByteSize, ACL_MEMCPY_DEVICE_TO_HOST); | ||
| 162 | + | ||
| 163 | + WriteFile(GetGoldenDir() + "/output_val.bin", dstHostVal, dstValByteSize); | ||
| 164 | + WriteFile(GetGoldenDir() + "/output_idx.bin", dstHostIdx, dstIdxByteSize); | ||
| 165 | + | ||
| 166 | + aclrtFree(dstDeviceIdx); | ||
| 167 | + aclrtFree(dstDeviceVal); | ||
| 168 | + aclrtFree(srcDevice); | ||
| 169 | + aclrtFreeHost(dstHostVal); | ||
| 170 | + aclrtFreeHost(dstHostIdx); | ||
| 171 | + aclrtFreeHost(srcHost); | ||
| 172 | + | ||
| 173 | + return CompareGoldenValIdx<TVal, TIdx>(dstValByteSize); | ||
| 174 | + } | ||
| 93 | }; | 175 | }; |
| 94 | 176 | ||
| 177 | +// ============================================================================= | ||
| 178 | +// Pure index TEST_F macros (35 cases) | ||
| 179 | +// ============================================================================= | ||
| 95 | TEST_F(TCOLCMAXTest, case01) | 180 | TEST_F(TCOLCMAXTest, case01) |
| 96 | { | 181 | { |
| 97 | bool ret = TCOLCMAXTestFramework<1, float, 1, 1, 1, 256, 255>(); | 182 | bool ret = TCOLCMAXTestFramework<1, float, 1, 1, 1, 256, 255>(); |
| @@ -107,6 +192,7 @@ TEST_F(TCOLCMAXTest, case03) | |||
| 107 | bool ret = TCOLCMAXTestFramework<3, float, 16, 15, 1, 256, 255>(); | 192 | bool ret = TCOLCMAXTestFramework<3, float, 16, 15, 1, 256, 255>(); |
| 108 | EXPECT_TRUE(ret); | 193 | EXPECT_TRUE(ret); |
| 109 | } | 194 | } |
| 195 | + | ||
| 110 | TEST_F(TCOLCMAXTest, case11) | 196 | TEST_F(TCOLCMAXTest, case11) |
| 111 | { | 197 | { |
| 112 | bool ret = TCOLCMAXTestFramework<11, aclFloat16, 1, 1, 1, 256, 255>(); | 198 | bool ret = TCOLCMAXTestFramework<11, aclFloat16, 1, 1, 1, 256, 255>(); |
| @@ -122,6 +208,55 @@ TEST_F(TCOLCMAXTest, case13) | |||
| 122 | bool ret = TCOLCMAXTestFramework<13, aclFloat16, 16, 15, 1, 256, 255>(); | 208 | bool ret = TCOLCMAXTestFramework<13, aclFloat16, 16, 15, 1, 256, 255>(); |
| 123 | EXPECT_TRUE(ret); | 209 | EXPECT_TRUE(ret); |
| 124 | } | 210 | } |
| 211 | + | ||
| 212 | +TEST_F(TCOLCMAXTest, case21) | ||
| 213 | +{ | ||
| 214 | + bool ret = TCOLCMAXTestFramework<21, int8_t, 1, 1, 1, 256, 255>(); | ||
| 215 | + EXPECT_TRUE(ret); | ||
| 216 | +} | ||
| 217 | +TEST_F(TCOLCMAXTest, case22) | ||
| 218 | +{ | ||
| 219 | + bool ret = TCOLCMAXTestFramework<22, int8_t, 16, 16, 1, 128, 127>(); | ||
| 220 | + EXPECT_TRUE(ret); | ||
| 221 | +} | ||
| 222 | +TEST_F(TCOLCMAXTest, case23) | ||
| 223 | +{ | ||
| 224 | + bool ret = TCOLCMAXTestFramework<23, int8_t, 16, 15, 1, 256, 255>(); | ||
| 225 | + EXPECT_TRUE(ret); | ||
| 226 | +} | ||
| 227 | + | ||
| 228 | +TEST_F(TCOLCMAXTest, case31) | ||
| 229 | +{ | ||
| 230 | + bool ret = TCOLCMAXTestFramework<31, uint8_t, 1, 1, 1, 256, 255>(); | ||
| 231 | + EXPECT_TRUE(ret); | ||
| 232 | +} | ||
| 233 | +TEST_F(TCOLCMAXTest, case32) | ||
| 234 | +{ | ||
| 235 | + bool ret = TCOLCMAXTestFramework<32, uint8_t, 16, 16, 1, 128, 127>(); | ||
| 236 | + EXPECT_TRUE(ret); | ||
| 237 | +} | ||
| 238 | +TEST_F(TCOLCMAXTest, case33) | ||
| 239 | +{ | ||
| 240 | + bool ret = TCOLCMAXTestFramework<33, uint8_t, 16, 15, 1, 256, 255>(); | ||
| 241 | + EXPECT_TRUE(ret); | ||
| 242 | +} | ||
| 243 | + | ||
| 244 | +TEST_F(TCOLCMAXTest, case41) | ||
| 245 | +{ | ||
| 246 | + bool ret = TCOLCMAXTestFramework<41, int16_t, 1, 1, 1, 256, 255>(); | ||
| 247 | + EXPECT_TRUE(ret); | ||
| 248 | +} | ||
| 249 | +TEST_F(TCOLCMAXTest, case42) | ||
| 250 | +{ | ||
| 251 | + bool ret = TCOLCMAXTestFramework<42, int16_t, 16, 16, 1, 128, 127>(); | ||
| 252 | + EXPECT_TRUE(ret); | ||
| 253 | +} | ||
| 254 | +TEST_F(TCOLCMAXTest, case43) | ||
| 255 | +{ | ||
| 256 | + bool ret = TCOLCMAXTestFramework<43, int16_t, 16, 15, 1, 256, 255>(); | ||
| 257 | + EXPECT_TRUE(ret); | ||
| 258 | +} | ||
| 259 | + | ||
| 125 | TEST_F(TCOLCMAXTest, case51) | 260 | TEST_F(TCOLCMAXTest, case51) |
| 126 | { | 261 | { |
| 127 | bool ret = TCOLCMAXTestFramework<51, uint16_t, 1, 1, 1, 256, 255>(); | 262 | bool ret = TCOLCMAXTestFramework<51, uint16_t, 1, 1, 1, 256, 255>(); |
| @@ -137,6 +272,23 @@ TEST_F(TCOLCMAXTest, case53) | |||
| 137 | bool ret = TCOLCMAXTestFramework<53, uint16_t, 16, 15, 1, 256, 255>(); | 272 | bool ret = TCOLCMAXTestFramework<53, uint16_t, 16, 15, 1, 256, 255>(); |
| 138 | EXPECT_TRUE(ret); | 273 | EXPECT_TRUE(ret); |
| 139 | } | 274 | } |
| 275 | + | ||
| 276 | +TEST_F(TCOLCMAXTest, case61) | ||
| 277 | +{ | ||
| 278 | + bool ret = TCOLCMAXTestFramework<61, int32_t, 1, 1, 1, 256, 255>(); | ||
| 279 | + EXPECT_TRUE(ret); | ||
| 280 | +} | ||
| 281 | +TEST_F(TCOLCMAXTest, case62) | ||
| 282 | +{ | ||
| 283 | + bool ret = TCOLCMAXTestFramework<62, int32_t, 16, 16, 1, 128, 127>(); | ||
| 284 | + EXPECT_TRUE(ret); | ||
| 285 | +} | ||
| 286 | +TEST_F(TCOLCMAXTest, case63) | ||
| 287 | +{ | ||
| 288 | + bool ret = TCOLCMAXTestFramework<63, int32_t, 16, 15, 1, 256, 255>(); | ||
| 289 | + EXPECT_TRUE(ret); | ||
| 290 | +} | ||
| 291 | + | ||
| 140 | TEST_F(TCOLCMAXTest, case71) | 292 | TEST_F(TCOLCMAXTest, case71) |
| 141 | { | 293 | { |
| 142 | bool ret = TCOLCMAXTestFramework<71, uint32_t, 1, 1, 1, 256, 255>(); | 294 | bool ret = TCOLCMAXTestFramework<71, uint32_t, 1, 1, 1, 256, 255>(); |
| @@ -152,3 +304,206 @@ TEST_F(TCOLCMAXTest, case73) | |||
| 152 | bool ret = TCOLCMAXTestFramework<73, uint32_t, 16, 15, 1, 256, 255>(); | 304 | bool ret = TCOLCMAXTestFramework<73, uint32_t, 16, 15, 1, 256, 255>(); |
| 153 | EXPECT_TRUE(ret); | 305 | EXPECT_TRUE(ret); |
| 154 | } | 306 | } |
| 307 | + | ||
| 308 | +TEST_F(TCOLCMAXTest, case81) | ||
| 309 | +{ | ||
| 310 | + bool ret = TCOLCMAXTestFramework<81, aclFloat16, 16, 16, 1, 32, 32>(); | ||
| 311 | + EXPECT_TRUE(ret); | ||
| 312 | +} | ||
| 313 | +TEST_F(TCOLCMAXTest, case82) | ||
| 314 | +{ | ||
| 315 | + bool ret = TCOLCMAXTestFramework<82, uint16_t, 16, 16, 1, 32, 32>(); | ||
| 316 | + EXPECT_TRUE(ret); | ||
| 317 | +} | ||
| 318 | +TEST_F(TCOLCMAXTest, case83) | ||
| 319 | +{ | ||
| 320 | + bool ret = TCOLCMAXTestFramework<83, uint32_t, 16, 16, 1, 32, 31>(); | ||
| 321 | + EXPECT_TRUE(ret); | ||
| 322 | +} | ||
| 323 | +TEST_F(TCOLCMAXTest, case84) | ||
| 324 | +{ | ||
| 325 | + bool ret = TCOLCMAXTestFramework<84, float, 16, 16, 1, 32, 31>(); | ||
| 326 | + EXPECT_TRUE(ret); | ||
| 327 | +} | ||
| 328 | +TEST_F(TCOLCMAXTest, case85) | ||
| 329 | +{ | ||
| 330 | + bool ret = TCOLCMAXTestFramework<85, int8_t, 16, 16, 1, 32, 31>(); | ||
| 331 | + EXPECT_TRUE(ret); | ||
| 332 | +} | ||
| 333 | +TEST_F(TCOLCMAXTest, case86) | ||
| 334 | +{ | ||
| 335 | + bool ret = TCOLCMAXTestFramework<86, uint8_t, 16, 16, 1, 32, 31>(); | ||
| 336 | + EXPECT_TRUE(ret); | ||
| 337 | +} | ||
| 338 | +TEST_F(TCOLCMAXTest, case87) | ||
| 339 | +{ | ||
| 340 | + bool ret = TCOLCMAXTestFramework<87, int16_t, 16, 16, 1, 32, 31>(); | ||
| 341 | + EXPECT_TRUE(ret); | ||
| 342 | +} | ||
| 343 | +TEST_F(TCOLCMAXTest, case88) | ||
| 344 | +{ | ||
| 345 | + bool ret = TCOLCMAXTestFramework<88, int32_t, 16, 16, 1, 32, 31>(); | ||
| 346 | + EXPECT_TRUE(ret); | ||
| 347 | +} | ||
| 348 | + | ||
| 349 | +TEST_F(TCOLCMAXTest, case91) | ||
| 350 | +{ | ||
| 351 | + bool ret = TCOLCMAXTestFramework<91, uint16_t, 16, 16, 1, 128, 120>(); | ||
| 352 | + EXPECT_TRUE(ret); | ||
| 353 | +} | ||
| 354 | +TEST_F(TCOLCMAXTest, case92) | ||
| 355 | +{ | ||
| 356 | + bool ret = TCOLCMAXTestFramework<92, aclFloat16, 16, 16, 1, 96, 88>(); | ||
| 357 | + EXPECT_TRUE(ret); | ||
| 358 | +} | ||
| 359 | +TEST_F(TCOLCMAXTest, case93) | ||
| 360 | +{ | ||
| 361 | + bool ret = TCOLCMAXTestFramework<93, uint16_t, 4, 4, 1, 48, 34>(); | ||
| 362 | + EXPECT_TRUE(ret); | ||
| 363 | +} | ||
| 364 | + | ||
| 365 | +// ============================================================================= | ||
| 366 | +// Value + index TEST_F macros (27 cases) | ||
| 367 | +// ============================================================================= | ||
| 368 | +TEST_F(TCOLCMAXTest, case001) | ||
| 369 | +{ | ||
| 370 | + bool ret = TCOLCMAXTestFramework<1, float, int32_t, 1, 1, 1, 256, 255>(); | ||
| 371 | + EXPECT_TRUE(ret); | ||
| 372 | +} | ||
| 373 | +TEST_F(TCOLCMAXTest, case002) | ||
| 374 | +{ | ||
| 375 | + bool ret = TCOLCMAXTestFramework<2, float, int32_t, 16, 16, 1, 128, 127>(); | ||
| 376 | + EXPECT_TRUE(ret); | ||
| 377 | +} | ||
| 378 | +TEST_F(TCOLCMAXTest, case003) | ||
| 379 | +{ | ||
| 380 | + bool ret = TCOLCMAXTestFramework<3, float, int32_t, 16, 15, 1, 256, 255>(); | ||
| 381 | + EXPECT_TRUE(ret); | ||
| 382 | +} | ||
| 383 | + | ||
| 384 | +TEST_F(TCOLCMAXTest, case011) | ||
| 385 | +{ | ||
| 386 | + bool ret = TCOLCMAXTestFramework<11, aclFloat16, int16_t, 1, 1, 1, 256, 255>(); | ||
| 387 | + EXPECT_TRUE(ret); | ||
| 388 | +} | ||
| 389 | +TEST_F(TCOLCMAXTest, case012) | ||
| 390 | +{ | ||
| 391 | + bool ret = TCOLCMAXTestFramework<12, aclFloat16, int16_t, 16, 16, 1, 128, 127>(); | ||
| 392 | + EXPECT_TRUE(ret); | ||
| 393 | +} | ||
| 394 | +TEST_F(TCOLCMAXTest, case013) | ||
| 395 | +{ | ||
| 396 | + bool ret = TCOLCMAXTestFramework<13, aclFloat16, int16_t, 16, 15, 1, 256, 255>(); | ||
| 397 | + EXPECT_TRUE(ret); | ||
| 398 | +} | ||
| 399 | + | ||
| 400 | +TEST_F(TCOLCMAXTest, case041) | ||
| 401 | +{ | ||
| 402 | + bool ret = TCOLCMAXTestFramework<41, int16_t, int16_t, 1, 1, 1, 256, 255>(); | ||
| 403 | + EXPECT_TRUE(ret); | ||
| 404 | +} | ||
| 405 | +TEST_F(TCOLCMAXTest, case042) | ||
| 406 | +{ | ||
| 407 | + bool ret = TCOLCMAXTestFramework<42, int16_t, int16_t, 16, 16, 1, 128, 127>(); | ||
| 408 | + EXPECT_TRUE(ret); | ||
| 409 | +} | ||
| 410 | +TEST_F(TCOLCMAXTest, case043) | ||
| 411 | +{ | ||
| 412 | + bool ret = TCOLCMAXTestFramework<43, int16_t, int16_t, 16, 15, 1, 256, 255>(); | ||
| 413 | + EXPECT_TRUE(ret); | ||
| 414 | +} | ||
| 415 | + | ||
| 416 | +TEST_F(TCOLCMAXTest, case051) | ||
| 417 | +{ | ||
| 418 | + bool ret = TCOLCMAXTestFramework<51, uint16_t, int16_t, 1, 1, 1, 256, 255>(); | ||
| 419 | + EXPECT_TRUE(ret); | ||
| 420 | +} | ||
| 421 | +TEST_F(TCOLCMAXTest, case052) | ||
| 422 | +{ | ||
| 423 | + bool ret = TCOLCMAXTestFramework<52, uint16_t, int16_t, 16, 16, 1, 128, 127>(); | ||
| 424 | + EXPECT_TRUE(ret); | ||
| 425 | +} | ||
| 426 | +TEST_F(TCOLCMAXTest, case053) | ||
| 427 | +{ | ||
| 428 | + bool ret = TCOLCMAXTestFramework<53, uint16_t, int16_t, 16, 15, 1, 256, 255>(); | ||
| 429 | + EXPECT_TRUE(ret); | ||
| 430 | +} | ||
| 431 | + | ||
| 432 | +TEST_F(TCOLCMAXTest, case061) | ||
| 433 | +{ | ||
| 434 | + bool ret = TCOLCMAXTestFramework<61, int32_t, int32_t, 1, 1, 1, 256, 255>(); | ||
| 435 | + EXPECT_TRUE(ret); | ||
| 436 | +} | ||
| 437 | +TEST_F(TCOLCMAXTest, case062) | ||
| 438 | +{ | ||
| 439 | + bool ret = TCOLCMAXTestFramework<62, int32_t, int32_t, 16, 16, 1, 128, 127>(); | ||
| 440 | + EXPECT_TRUE(ret); | ||
| 441 | +} | ||
| 442 | +TEST_F(TCOLCMAXTest, case063) | ||
| 443 | +{ | ||
| 444 | + bool ret = TCOLCMAXTestFramework<63, int32_t, int32_t, 16, 15, 1, 256, 255>(); | ||
| 445 | + EXPECT_TRUE(ret); | ||
| 446 | +} | ||
| 447 | + | ||
| 448 | +TEST_F(TCOLCMAXTest, case071) | ||
| 449 | +{ | ||
| 450 | + bool ret = TCOLCMAXTestFramework<71, uint32_t, int32_t, 1, 1, 1, 256, 255>(); | ||
| 451 | + EXPECT_TRUE(ret); | ||
| 452 | +} | ||
| 453 | +TEST_F(TCOLCMAXTest, case072) | ||
| 454 | +{ | ||
| 455 | + bool ret = TCOLCMAXTestFramework<72, uint32_t, int32_t, 16, 16, 1, 128, 127>(); | ||
| 456 | + EXPECT_TRUE(ret); | ||
| 457 | +} | ||
| 458 | +TEST_F(TCOLCMAXTest, case073) | ||
| 459 | +{ | ||
| 460 | + bool ret = TCOLCMAXTestFramework<73, uint32_t, int32_t, 16, 15, 1, 256, 255>(); | ||
| 461 | + EXPECT_TRUE(ret); | ||
| 462 | +} | ||
| 463 | + | ||
| 464 | +TEST_F(TCOLCMAXTest, case081) | ||
| 465 | +{ | ||
| 466 | + bool ret = TCOLCMAXTestFramework<81, aclFloat16, int16_t, 16, 16, 1, 32, 32>(); | ||
| 467 | + EXPECT_TRUE(ret); | ||
| 468 | +} | ||
| 469 | +TEST_F(TCOLCMAXTest, case082) | ||
| 470 | +{ | ||
| 471 | + bool ret = TCOLCMAXTestFramework<82, uint16_t, int16_t, 16, 16, 1, 32, 32>(); | ||
| 472 | + EXPECT_TRUE(ret); | ||
| 473 | +} | ||
| 474 | +TEST_F(TCOLCMAXTest, case083) | ||
| 475 | +{ | ||
| 476 | + bool ret = TCOLCMAXTestFramework<83, uint32_t, int32_t, 16, 16, 1, 32, 31>(); | ||
| 477 | + EXPECT_TRUE(ret); | ||
| 478 | +} | ||
| 479 | +TEST_F(TCOLCMAXTest, case084) | ||
| 480 | +{ | ||
| 481 | + bool ret = TCOLCMAXTestFramework<84, float, int32_t, 16, 16, 1, 32, 31>(); | ||
| 482 | + EXPECT_TRUE(ret); | ||
| 483 | +} | ||
| 484 | +TEST_F(TCOLCMAXTest, case085) | ||
| 485 | +{ | ||
| 486 | + bool ret = TCOLCMAXTestFramework<85, int16_t, int16_t, 16, 16, 1, 32, 31>(); | ||
| 487 | + EXPECT_TRUE(ret); | ||
| 488 | +} | ||
| 489 | +TEST_F(TCOLCMAXTest, case086) | ||
| 490 | +{ | ||
| 491 | + bool ret = TCOLCMAXTestFramework<86, int32_t, int32_t, 16, 16, 1, 32, 31>(); | ||
| 492 | + EXPECT_TRUE(ret); | ||
| 493 | +} | ||
| 494 | + | ||
| 495 | +TEST_F(TCOLCMAXTest, case091) | ||
| 496 | +{ | ||
| 497 | + bool ret = TCOLCMAXTestFramework<91, uint16_t, int16_t, 16, 16, 1, 128, 120>(); | ||
| 498 | + EXPECT_TRUE(ret); | ||
| 499 | +} | ||
| 500 | +TEST_F(TCOLCMAXTest, case092) | ||
| 501 | +{ | ||
| 502 | + bool ret = TCOLCMAXTestFramework<92, aclFloat16, int16_t, 16, 16, 1, 96, 88>(); | ||
| 503 | + EXPECT_TRUE(ret); | ||
| 504 | +} | ||
| 505 | +TEST_F(TCOLCMAXTest, case093) | ||
| 506 | +{ | ||
| 507 | + bool ret = TCOLCMAXTestFramework<93, uint16_t, int16_t, 4, 4, 1, 48, 34>(); | ||
| 508 | + EXPECT_TRUE(ret); | ||
| 509 | +} | ||
| @@ -15,6 +15,9 @@ See LICENSE in the root of the software repository for the full text of the Lice | |||
| 15 | using namespace std; | 15 | using namespace std; |
| 16 | using namespace pto; | 16 | using namespace pto; |
| 17 | 17 | ||
| 18 | +// ============================================================================= | ||
| 19 | +// Pure index mode kernel (3-arg TCOLARGMAX) | ||
| 20 | +// ============================================================================= | ||
| 18 | template <typename T, int srcRow, int srcValidRow, int dstRow, int col, int validCol> | 21 | template <typename T, int srcRow, int srcValidRow, int dstRow, int col, int validCol> |
| 19 | PTO_INTERNAL void runTColCMax(__gm__ uint32_t __out__ *out, __gm__ T __in__ *src, bool isBinary) | 22 | PTO_INTERNAL void runTColCMax(__gm__ uint32_t __out__ *out, __gm__ T __in__ *src, bool isBinary) |
| 20 | { | 23 | { |
| @@ -35,7 +38,6 @@ PTO_INTERNAL void runTColCMax(__gm__ uint32_t __out__ *out, __gm__ T __in__ *src | |||
| 35 | TASSIGN(dstTile, srcRow * col * sizeof(T)); | 38 | TASSIGN(dstTile, srcRow * col * sizeof(T)); |
| 36 | TASSIGN(tmpTile, srcRow * col * sizeof(T) + col * sizeof(uint32_t)); | 39 | TASSIGN(tmpTile, srcRow * col * sizeof(T) + col * sizeof(uint32_t)); |
| 37 | 40 | ||
| 38 | - // 搬运数据 | ||
| 39 | TLOAD(srcTile, srcGlobal); | 41 | TLOAD(srcTile, srcGlobal); |
| 40 | 42 | ||
| 41 | set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); | 43 | set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); |
| @@ -47,6 +49,55 @@ PTO_INTERNAL void runTColCMax(__gm__ uint32_t __out__ *out, __gm__ T __in__ *src | |||
| 47 | out = dstGlobal.data(); | 49 | out = dstGlobal.data(); |
| 48 | } | 50 | } |
| 49 | 51 | ||
| 52 | +// ============================================================================= | ||
| 53 | +// Value + index mode kernel (4-arg TCOLARGMAX) | ||
| 54 | +// ============================================================================= | ||
| 55 | +template <typename TVal, typename TIdx, int srcRow, int srcValidRow, int dstRow, int col, int validCol> | ||
| 56 | +PTO_INTERNAL void runTColIdxValMax(__gm__ TVal __out__ *outVal, __gm__ TIdx __out__ *outIdx, __gm__ TVal __in__ *src) | ||
| 57 | +{ | ||
| 58 | + using DynDim2Shape = Shape<1, 1, 1, -1, -1>; | ||
| 59 | + using DynDim2Stride = pto::Stride<1, 1, -1, -1, 1>; | ||
| 60 | + | ||
| 61 | + using GlobalDataSrc = GlobalTensor<TVal, DynDim2Shape, DynDim2Stride>; | ||
| 62 | + using GlobalDataDstVal = GlobalTensor<TVal, DynDim2Shape, DynDim2Stride>; | ||
| 63 | + using GlobalDataDstIdx = GlobalTensor<TIdx, DynDim2Shape, DynDim2Stride>; | ||
| 64 | + | ||
| 65 | + GlobalDataSrc srcGlobal(src, DynDim2Shape(srcValidRow, validCol), DynDim2Stride(srcRow, col)); | ||
| 66 | + GlobalDataDstVal dstValGlobal(outVal, DynDim2Shape(dstRow, validCol), DynDim2Stride(dstRow, col)); | ||
| 67 | + GlobalDataDstIdx dstIdxGlobal(outIdx, DynDim2Shape(dstRow, validCol), DynDim2Stride(dstRow, col)); | ||
| 68 | + | ||
| 69 | + using SrcTileData = Tile<TileType::Vec, TVal, srcRow, col, BLayout::RowMajor, -1, -1>; | ||
| 70 | + using DstValTileData = Tile<TileType::Vec, TVal, dstRow, col, BLayout::RowMajor, -1, -1>; | ||
| 71 | + using DstIdxTileData = Tile<TileType::Vec, TIdx, dstRow, col, BLayout::RowMajor, -1, -1>; | ||
| 72 | + using TmpTile = Tile<TileType::Vec, TVal, 1, 32, BLayout::RowMajor, -1, -1>; | ||
| 73 | + | ||
| 74 | + SrcTileData srcTile(srcValidRow, validCol); | ||
| 75 | + DstValTileData valTile(dstRow, validCol); | ||
| 76 | + DstIdxTileData idxTile(dstRow, validCol); | ||
| 77 | + TmpTile tmpTile(1, 32); | ||
| 78 | + | ||
| 79 | + TASSIGN(srcTile, 0x0); | ||
| 80 | + TASSIGN(idxTile, srcRow * col * sizeof(TVal)); | ||
| 81 | + TASSIGN(valTile, (srcRow + 1) * col * sizeof(TVal)); | ||
| 82 | + TASSIGN(tmpTile, (srcRow + 2) * col * sizeof(TVal)); | ||
| 83 | + | ||
| 84 | + TLOAD(srcTile, srcGlobal); | ||
| 85 | + | ||
| 86 | + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); | ||
| 87 | + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); | ||
| 88 | + TCOLARGMAX(valTile, idxTile, srcTile, tmpTile); | ||
| 89 | + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); | ||
| 90 | + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); | ||
| 91 | + TSTORE(dstValGlobal, valTile); | ||
| 92 | + TSTORE(dstIdxGlobal, idxTile); | ||
| 93 | + | ||
| 94 | + outVal = dstValGlobal.data(); | ||
| 95 | + outIdx = dstIdxGlobal.data(); | ||
| 96 | +} | ||
| 97 | + | ||
| 98 | +// ============================================================================= | ||
| 99 | +// Pure index extern "C" entry points -- float32 | ||
| 100 | +// ============================================================================= | ||
| 50 | extern "C" __global__ AICORE void launchTCOLCMAXCase01(__gm__ uint32_t *out, __gm__ float *src) | 101 | extern "C" __global__ AICORE void launchTCOLCMAXCase01(__gm__ uint32_t *out, __gm__ float *src) |
| 51 | { | 102 | { |
| 52 | runTColCMax<float, 1, 1, 1, 256, 255>(out, src, false); | 103 | runTColCMax<float, 1, 1, 1, 256, 255>(out, src, false); |
| @@ -59,6 +110,10 @@ extern "C" __global__ AICORE void launchTCOLCMAXCase03(__gm__ uint32_t *out, __g | |||
| 59 | { | 110 | { |
| 60 | runTColCMax<float, 16, 15, 1, 256, 255>(out, src, false); | 111 | runTColCMax<float, 16, 15, 1, 256, 255>(out, src, false); |
| 61 | } | 112 | } |
| 113 | + | ||
| 114 | +// ============================================================================= | ||
| 115 | +// Pure index extern "C" entry points -- float16 | ||
| 116 | +// ============================================================================= | ||
| 62 | extern "C" __global__ AICORE void launchTCOLCMAXCase11(__gm__ uint32_t *out, __gm__ half *src) | 117 | extern "C" __global__ AICORE void launchTCOLCMAXCase11(__gm__ uint32_t *out, __gm__ half *src) |
| 63 | { | 118 | { |
| 64 | runTColCMax<half, 1, 1, 1, 256, 255>(out, src, false); | 119 | runTColCMax<half, 1, 1, 1, 256, 255>(out, src, false); |
| @@ -71,6 +126,10 @@ extern "C" __global__ AICORE void launchTCOLCMAXCase13(__gm__ uint32_t *out, __g | |||
| 71 | { | 126 | { |
| 72 | runTColCMax<half, 16, 15, 1, 256, 255>(out, src, false); | 127 | runTColCMax<half, 16, 15, 1, 256, 255>(out, src, false); |
| 73 | } | 128 | } |
| 129 | + | ||
| 130 | +// ============================================================================= | ||
| 131 | +// Pure index extern "C" entry points -- int8 | ||
| 132 | +// ============================================================================= | ||
| 74 | extern "C" __global__ AICORE void launchTCOLCMAXCase21(__gm__ uint32_t *out, __gm__ int8_t *src) | 133 | extern "C" __global__ AICORE void launchTCOLCMAXCase21(__gm__ uint32_t *out, __gm__ int8_t *src) |
| 75 | { | 134 | { |
| 76 | runTColCMax<int8_t, 1, 1, 1, 256, 255>(out, src, false); | 135 | runTColCMax<int8_t, 1, 1, 1, 256, 255>(out, src, false); |
| @@ -83,6 +142,10 @@ extern "C" __global__ AICORE void launchTCOLCMAXCase23(__gm__ uint32_t *out, __g | |||
| 83 | { | 142 | { |
| 84 | runTColCMax<int8_t, 16, 15, 1, 256, 255>(out, src, false); | 143 | runTColCMax<int8_t, 16, 15, 1, 256, 255>(out, src, false); |
| 85 | } | 144 | } |
| 145 | + | ||
| 146 | +// ============================================================================= | ||
| 147 | +// Pure index extern "C" entry points -- uint8 | ||
| 148 | +// ============================================================================= | ||
| 86 | extern "C" __global__ AICORE void launchTCOLCMAXCase31(__gm__ uint32_t *out, __gm__ uint8_t *src) | 149 | extern "C" __global__ AICORE void launchTCOLCMAXCase31(__gm__ uint32_t *out, __gm__ uint8_t *src) |
| 87 | { | 150 | { |
| 88 | runTColCMax<uint8_t, 1, 1, 1, 256, 255>(out, src, false); | 151 | runTColCMax<uint8_t, 1, 1, 1, 256, 255>(out, src, false); |
| @@ -95,6 +158,10 @@ extern "C" __global__ AICORE void launchTCOLCMAXCase33(__gm__ uint32_t *out, __g | |||
| 95 | { | 158 | { |
| 96 | runTColCMax<uint8_t, 16, 15, 1, 256, 255>(out, src, false); | 159 | runTColCMax<uint8_t, 16, 15, 1, 256, 255>(out, src, false); |
| 97 | } | 160 | } |
| 161 | + | ||
| 162 | +// ============================================================================= | ||
| 163 | +// Pure index extern "C" entry points -- int16 | ||
| 164 | +// ============================================================================= | ||
| 98 | extern "C" __global__ AICORE void launchTCOLCMAXCase41(__gm__ uint32_t *out, __gm__ int16_t *src) | 165 | extern "C" __global__ AICORE void launchTCOLCMAXCase41(__gm__ uint32_t *out, __gm__ int16_t *src) |
| 99 | { | 166 | { |
| 100 | runTColCMax<int16_t, 1, 1, 1, 256, 255>(out, src, false); | 167 | runTColCMax<int16_t, 1, 1, 1, 256, 255>(out, src, false); |
| @@ -107,6 +174,10 @@ extern "C" __global__ AICORE void launchTCOLCMAXCase43(__gm__ uint32_t *out, __g | |||
| 107 | { | 174 | { |
| 108 | runTColCMax<int16_t, 16, 15, 1, 256, 255>(out, src, false); | 175 | runTColCMax<int16_t, 16, 15, 1, 256, 255>(out, src, false); |
| 109 | } | 176 | } |
| 177 | + | ||
| 178 | +// ============================================================================= | ||
| 179 | +// Pure index extern "C" entry points -- uint16 | ||
| 180 | +// ============================================================================= | ||
| 110 | extern "C" __global__ AICORE void launchTCOLCMAXCase51(__gm__ uint32_t *out, __gm__ uint16_t *src) | 181 | extern "C" __global__ AICORE void launchTCOLCMAXCase51(__gm__ uint32_t *out, __gm__ uint16_t *src) |
| 111 | { | 182 | { |
| 112 | runTColCMax<uint16_t, 1, 1, 1, 256, 255>(out, src, false); | 183 | runTColCMax<uint16_t, 1, 1, 1, 256, 255>(out, src, false); |
| @@ -119,6 +190,10 @@ extern "C" __global__ AICORE void launchTCOLCMAXCase53(__gm__ uint32_t *out, __g | |||
| 119 | { | 190 | { |
| 120 | runTColCMax<uint16_t, 16, 15, 1, 256, 255>(out, src, false); | 191 | runTColCMax<uint16_t, 16, 15, 1, 256, 255>(out, src, false); |
| 121 | } | 192 | } |
| 193 | + | ||
| 194 | +// ============================================================================= | ||
| 195 | +// Pure index extern "C" entry points -- int32 | ||
| 196 | +// ============================================================================= | ||
| 122 | extern "C" __global__ AICORE void launchTCOLCMAXCase61(__gm__ uint32_t *out, __gm__ int32_t *src) | 197 | extern "C" __global__ AICORE void launchTCOLCMAXCase61(__gm__ uint32_t *out, __gm__ int32_t *src) |
| 123 | { | 198 | { |
| 124 | runTColCMax<int32_t, 1, 1, 1, 256, 255>(out, src, false); | 199 | runTColCMax<int32_t, 1, 1, 1, 256, 255>(out, src, false); |
| @@ -131,6 +206,10 @@ extern "C" __global__ AICORE void launchTCOLCMAXCase63(__gm__ uint32_t *out, __g | |||
| 131 | { | 206 | { |
| 132 | runTColCMax<int32_t, 16, 15, 1, 256, 255>(out, src, false); | 207 | runTColCMax<int32_t, 16, 15, 1, 256, 255>(out, src, false); |
| 133 | } | 208 | } |
| 209 | + | ||
| 210 | +// ============================================================================= | ||
| 211 | +// Pure index extern "C" entry points -- uint32 | ||
| 212 | +// ============================================================================= | ||
| 134 | extern "C" __global__ AICORE void launchTCOLCMAXCase71(__gm__ uint32_t *out, __gm__ uint32_t *src) | 213 | extern "C" __global__ AICORE void launchTCOLCMAXCase71(__gm__ uint32_t *out, __gm__ uint32_t *src) |
| 135 | { | 214 | { |
| 136 | runTColCMax<uint32_t, 1, 1, 1, 256, 255>(out, src, false); | 215 | runTColCMax<uint32_t, 1, 1, 1, 256, 255>(out, src, false); |
| @@ -144,6 +223,220 @@ extern "C" __global__ AICORE void launchTCOLCMAXCase73(__gm__ uint32_t *out, __g | |||
| 144 | runTColCMax<uint32_t, 16, 15, 1, 256, 255>(out, src, false); | 223 | runTColCMax<uint32_t, 16, 15, 1, 256, 255>(out, src, false); |
| 145 | } | 224 | } |
| 146 | 225 | ||
| 226 | +// ============================================================================= | ||
| 227 | +// Pure index extern "C" entry points -- small dim edge cases | ||
| 228 | +// ============================================================================= | ||
| 229 | +extern "C" __global__ AICORE void launchTCOLCMAXCase81(__gm__ uint32_t *out, __gm__ half *src) | ||
| 230 | +{ | ||
| 231 | + runTColCMax<half, 16, 16, 1, 32, 32>(out, src, false); | ||
| 232 | +} | ||
| 233 | +extern "C" __global__ AICORE void launchTCOLCMAXCase82(__gm__ uint32_t *out, __gm__ uint16_t *src) | ||
| 234 | +{ | ||
| 235 | + runTColCMax<uint16_t, 16, 16, 1, 32, 32>(out, src, false); | ||
| 236 | +} | ||
| 237 | +extern "C" __global__ AICORE void launchTCOLCMAXCase83(__gm__ uint32_t *out, __gm__ uint32_t *src) | ||
| 238 | +{ | ||
| 239 | + runTColCMax<uint32_t, 16, 16, 1, 32, 31>(out, src, false); | ||
| 240 | +} | ||
| 241 | +extern "C" __global__ AICORE void launchTCOLCMAXCase84(__gm__ uint32_t *out, __gm__ float *src) | ||
| 242 | +{ | ||
| 243 | + runTColCMax<float, 16, 16, 1, 32, 31>(out, src, false); | ||
| 244 | +} | ||
| 245 | +extern "C" __global__ AICORE void launchTCOLCMAXCase85(__gm__ uint32_t *out, __gm__ int8_t *src) | ||
| 246 | +{ | ||
| 247 | + runTColCMax<int8_t, 16, 16, 1, 32, 31>(out, src, false); | ||
| 248 | +} | ||
| 249 | +extern "C" __global__ AICORE void launchTCOLCMAXCase86(__gm__ uint32_t *out, __gm__ uint8_t *src) | ||
| 250 | +{ | ||
| 251 | + runTColCMax<uint8_t, 16, 16, 1, 32, 31>(out, src, false); | ||
| 252 | +} | ||
| 253 | +extern "C" __global__ AICORE void launchTCOLCMAXCase87(__gm__ uint32_t *out, __gm__ int16_t *src) | ||
| 254 | +{ | ||
| 255 | + runTColCMax<int16_t, 16, 16, 1, 32, 31>(out, src, false); | ||
| 256 | +} | ||
| 257 | +extern "C" __global__ AICORE void launchTCOLCMAXCase88(__gm__ uint32_t *out, __gm__ int32_t *src) | ||
| 258 | +{ | ||
| 259 | + runTColCMax<int32_t, 16, 16, 1, 32, 31>(out, src, false); | ||
| 260 | +} | ||
| 261 | +extern "C" __global__ AICORE void launchTCOLCMAXCase91(__gm__ uint32_t *out, __gm__ uint16_t *src) | ||
| 262 | +{ | ||
| 263 | + runTColCMax<uint16_t, 16, 16, 1, 128, 120>(out, src, false); | ||
| 264 | +} | ||
| 265 | +extern "C" __global__ AICORE void launchTCOLCMAXCase92(__gm__ uint32_t *out, __gm__ half *src) | ||
| 266 | +{ | ||
| 267 | + runTColCMax<half, 16, 16, 1, 96, 88>(out, src, false); | ||
| 268 | +} | ||
| 269 | +extern "C" __global__ AICORE void launchTCOLCMAXCase93(__gm__ uint32_t *out, __gm__ uint16_t *src) | ||
| 270 | +{ | ||
| 271 | + runTColCMax<uint16_t, 4, 4, 1, 48, 34>(out, src, false); | ||
| 272 | +} | ||
| 273 | + | ||
| 274 | +// ============================================================================= | ||
| 275 | +// Value + index extern "C" entry points -- float32 + int32_t index | ||
| 276 | +// ============================================================================= | ||
| 277 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase001(__gm__ float *outVal, __gm__ int32_t *outIdx, | ||
| 278 | + __gm__ float *src) | ||
| 279 | +{ | ||
| 280 | + runTColIdxValMax<float, int32_t, 1, 1, 1, 256, 255>(outVal, outIdx, src); | ||
| 281 | +} | ||
| 282 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase002(__gm__ float *outVal, __gm__ int32_t *outIdx, | ||
| 283 | + __gm__ float *src) | ||
| 284 | +{ | ||
| 285 | + runTColIdxValMax<float, int32_t, 16, 16, 1, 128, 127>(outVal, outIdx, src); | ||
| 286 | +} | ||
| 287 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase003(__gm__ float *outVal, __gm__ int32_t *outIdx, | ||
| 288 | + __gm__ float *src) | ||
| 289 | +{ | ||
| 290 | + runTColIdxValMax<float, int32_t, 16, 15, 1, 256, 255>(outVal, outIdx, src); | ||
| 291 | +} | ||
| 292 | + | ||
| 293 | +// ============================================================================= | ||
| 294 | +// Value + index extern "C" entry points -- float16 + int16_t index | ||
| 295 | +// ============================================================================= | ||
| 296 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase011(__gm__ half *outVal, __gm__ int16_t *outIdx, | ||
| 297 | + __gm__ half *src) | ||
| 298 | +{ | ||
| 299 | + runTColIdxValMax<half, int16_t, 1, 1, 1, 256, 255>(outVal, outIdx, src); | ||
| 300 | +} | ||
| 301 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase012(__gm__ half *outVal, __gm__ int16_t *outIdx, | ||
| 302 | + __gm__ half *src) | ||
| 303 | +{ | ||
| 304 | + runTColIdxValMax<half, int16_t, 16, 16, 1, 128, 127>(outVal, outIdx, src); | ||
| 305 | +} | ||
| 306 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase013(__gm__ half *outVal, __gm__ int16_t *outIdx, | ||
| 307 | + __gm__ half *src) | ||
| 308 | +{ | ||
| 309 | + runTColIdxValMax<half, int16_t, 16, 15, 1, 256, 255>(outVal, outIdx, src); | ||
| 310 | +} | ||
| 311 | + | ||
| 312 | +// ============================================================================= | ||
| 313 | +// Value + index extern "C" entry points -- int16 + int16_t index | ||
| 314 | +// ============================================================================= | ||
| 315 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase041(__gm__ int16_t *outVal, __gm__ int16_t *outIdx, | ||
| 316 | + __gm__ int16_t *src) | ||
| 317 | +{ | ||
| 318 | + runTColIdxValMax<int16_t, int16_t, 1, 1, 1, 256, 255>(outVal, outIdx, src); | ||
| 319 | +} | ||
| 320 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase042(__gm__ int16_t *outVal, __gm__ int16_t *outIdx, | ||
| 321 | + __gm__ int16_t *src) | ||
| 322 | +{ | ||
| 323 | + runTColIdxValMax<int16_t, int16_t, 16, 16, 1, 128, 127>(outVal, outIdx, src); | ||
| 324 | +} | ||
| 325 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase043(__gm__ int16_t *outVal, __gm__ int16_t *outIdx, | ||
| 326 | + __gm__ int16_t *src) | ||
| 327 | +{ | ||
| 328 | + runTColIdxValMax<int16_t, int16_t, 16, 15, 1, 256, 255>(outVal, outIdx, src); | ||
| 329 | +} | ||
| 330 | + | ||
| 331 | +// ============================================================================= | ||
| 332 | +// Value + index extern "C" entry points -- uint16 + int16_t index | ||
| 333 | +// ============================================================================= | ||
| 334 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase051(__gm__ uint16_t *outVal, __gm__ int16_t *outIdx, | ||
| 335 | + __gm__ uint16_t *src) | ||
| 336 | +{ | ||
| 337 | + runTColIdxValMax<uint16_t, int16_t, 1, 1, 1, 256, 255>(outVal, outIdx, src); | ||
| 338 | +} | ||
| 339 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase052(__gm__ uint16_t *outVal, __gm__ int16_t *outIdx, | ||
| 340 | + __gm__ uint16_t *src) | ||
| 341 | +{ | ||
| 342 | + runTColIdxValMax<uint16_t, int16_t, 16, 16, 1, 128, 127>(outVal, outIdx, src); | ||
| 343 | +} | ||
| 344 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase053(__gm__ uint16_t *outVal, __gm__ int16_t *outIdx, | ||
| 345 | + __gm__ uint16_t *src) | ||
| 346 | +{ | ||
| 347 | + runTColIdxValMax<uint16_t, int16_t, 16, 15, 1, 256, 255>(outVal, outIdx, src); | ||
| 348 | +} | ||
| 349 | + | ||
| 350 | +// ============================================================================= | ||
| 351 | +// Value + index extern "C" entry points -- int32 + int32_t index | ||
| 352 | +// ============================================================================= | ||
| 353 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase061(__gm__ int32_t *outVal, __gm__ int32_t *outIdx, | ||
| 354 | + __gm__ int32_t *src) | ||
| 355 | +{ | ||
| 356 | + runTColIdxValMax<int32_t, int32_t, 1, 1, 1, 256, 255>(outVal, outIdx, src); | ||
| 357 | +} | ||
| 358 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase062(__gm__ int32_t *outVal, __gm__ int32_t *outIdx, | ||
| 359 | + __gm__ int32_t *src) | ||
| 360 | +{ | ||
| 361 | + runTColIdxValMax<int32_t, int32_t, 16, 16, 1, 128, 127>(outVal, outIdx, src); | ||
| 362 | +} | ||
| 363 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase063(__gm__ int32_t *outVal, __gm__ int32_t *outIdx, | ||
| 364 | + __gm__ int32_t *src) | ||
| 365 | +{ | ||
| 366 | + runTColIdxValMax<int32_t, int32_t, 16, 15, 1, 256, 255>(outVal, outIdx, src); | ||
| 367 | +} | ||
| 368 | + | ||
| 369 | +// ============================================================================= | ||
| 370 | +// Value + index extern "C" entry points -- uint32 + int32_t index | ||
| 371 | +// ============================================================================= | ||
| 372 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase071(__gm__ uint32_t *outVal, __gm__ int32_t *outIdx, | ||
| 373 | + __gm__ uint32_t *src) | ||
| 374 | +{ | ||
| 375 | + runTColIdxValMax<uint32_t, int32_t, 1, 1, 1, 256, 255>(outVal, outIdx, src); | ||
| 376 | +} | ||
| 377 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase072(__gm__ uint32_t *outVal, __gm__ int32_t *outIdx, | ||
| 378 | + __gm__ uint32_t *src) | ||
| 379 | +{ | ||
| 380 | + runTColIdxValMax<uint32_t, int32_t, 16, 16, 1, 128, 127>(outVal, outIdx, src); | ||
| 381 | +} | ||
| 382 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase073(__gm__ uint32_t *outVal, __gm__ int32_t *outIdx, | ||
| 383 | + __gm__ uint32_t *src) | ||
| 384 | +{ | ||
| 385 | + runTColIdxValMax<uint32_t, int32_t, 16, 15, 1, 256, 255>(outVal, outIdx, src); | ||
| 386 | +} | ||
| 387 | + | ||
| 388 | +// ============================================================================= | ||
| 389 | +// Value + index extern "C" entry points -- small dim edge cases | ||
| 390 | +// ============================================================================= | ||
| 391 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase081(__gm__ half *outVal, __gm__ int16_t *outIdx, | ||
| 392 | + __gm__ half *src) | ||
| 393 | +{ | ||
| 394 | + runTColIdxValMax<half, int16_t, 16, 16, 1, 32, 32>(outVal, outIdx, src); | ||
| 395 | +} | ||
| 396 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase082(__gm__ uint16_t *outVal, __gm__ int16_t *outIdx, | ||
| 397 | + __gm__ uint16_t *src) | ||
| 398 | +{ | ||
| 399 | + runTColIdxValMax<uint16_t, int16_t, 16, 16, 1, 32, 32>(outVal, outIdx, src); | ||
| 400 | +} | ||
| 401 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase083(__gm__ uint32_t *outVal, __gm__ int32_t *outIdx, | ||
| 402 | + __gm__ uint32_t *src) | ||
| 403 | +{ | ||
| 404 | + runTColIdxValMax<uint32_t, int32_t, 16, 16, 1, 32, 31>(outVal, outIdx, src); | ||
| 405 | +} | ||
| 406 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase084(__gm__ float *outVal, __gm__ int32_t *outIdx, | ||
| 407 | + __gm__ float *src) | ||
| 408 | +{ | ||
| 409 | + runTColIdxValMax<float, int32_t, 16, 16, 1, 32, 31>(outVal, outIdx, src); | ||
| 410 | +} | ||
| 411 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase085(__gm__ int16_t *outVal, __gm__ int16_t *outIdx, | ||
| 412 | + __gm__ int16_t *src) | ||
| 413 | +{ | ||
| 414 | + runTColIdxValMax<int16_t, int16_t, 16, 16, 1, 32, 31>(outVal, outIdx, src); | ||
| 415 | +} | ||
| 416 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase086(__gm__ int32_t *outVal, __gm__ int32_t *outIdx, | ||
| 417 | + __gm__ int32_t *src) | ||
| 418 | +{ | ||
| 419 | + runTColIdxValMax<int32_t, int32_t, 16, 16, 1, 32, 31>(outVal, outIdx, src); | ||
| 420 | +} | ||
| 421 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase091(__gm__ uint16_t *outVal, __gm__ int16_t *outIdx, | ||
| 422 | + __gm__ uint16_t *src) | ||
| 423 | +{ | ||
| 424 | + runTColIdxValMax<uint16_t, int16_t, 16, 16, 1, 128, 120>(outVal, outIdx, src); | ||
| 425 | +} | ||
| 426 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase092(__gm__ half *outVal, __gm__ int16_t *outIdx, | ||
| 427 | + __gm__ half *src) | ||
| 428 | +{ | ||
| 429 | + runTColIdxValMax<half, int16_t, 16, 16, 1, 96, 88>(outVal, outIdx, src); | ||
| 430 | +} | ||
| 431 | +extern "C" __global__ AICORE void launchTCOLIDXVALMAXCase093(__gm__ uint16_t *outVal, __gm__ int16_t *outIdx, | ||
| 432 | + __gm__ uint16_t *src) | ||
| 433 | +{ | ||
| 434 | + runTColIdxValMax<uint16_t, int16_t, 4, 4, 1, 48, 34>(outVal, outIdx, src); | ||
| 435 | +} | ||
| 436 | + | ||
| 437 | +// ============================================================================= | ||
| 438 | +// Pure index dispatcher | ||
| 439 | +// ============================================================================= | ||
| 147 | template <uint32_t caseId> | 440 | template <uint32_t caseId> |
| 148 | void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream) | 441 | void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream) |
| 149 | { | 442 | { |
| @@ -172,6 +465,42 @@ void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream) | |||
| 172 | launchTCOLCMAXCase13<<<1, nullptr, stream>>>((uint32_t *)out, (half *)src); | 465 | launchTCOLCMAXCase13<<<1, nullptr, stream>>>((uint32_t *)out, (half *)src); |
| 173 | break; | 466 | break; |
| 174 | } | 467 | } |
| 468 | + case 21: { | ||
| 469 | + launchTCOLCMAXCase21<<<1, nullptr, stream>>>((uint32_t *)out, (int8_t *)src); | ||
| 470 | + break; | ||
| 471 | + } | ||
| 472 | + case 22: { | ||
| 473 | + launchTCOLCMAXCase22<<<1, nullptr, stream>>>((uint32_t *)out, (int8_t *)src); | ||
| 474 | + break; | ||
| 475 | + } | ||
| 476 | + case 23: { | ||
| 477 | + launchTCOLCMAXCase23<<<1, nullptr, stream>>>((uint32_t *)out, (int8_t *)src); | ||
| 478 | + break; | ||
| 479 | + } | ||
| 480 | + case 31: { | ||
| 481 | + launchTCOLCMAXCase31<<<1, nullptr, stream>>>((uint32_t *)out, (uint8_t *)src); | ||
| 482 | + break; | ||
| 483 | + } | ||
| 484 | + case 32: { | ||
| 485 | + launchTCOLCMAXCase32<<<1, nullptr, stream>>>((uint32_t *)out, (uint8_t *)src); | ||
| 486 | + break; | ||
| 487 | + } | ||
| 488 | + case 33: { | ||
| 489 | + launchTCOLCMAXCase33<<<1, nullptr, stream>>>((uint32_t *)out, (uint8_t *)src); | ||
| 490 | + break; | ||
| 491 | + } | ||
| 492 | + case 41: { | ||
| 493 | + launchTCOLCMAXCase41<<<1, nullptr, stream>>>((uint32_t *)out, (int16_t *)src); | ||
| 494 | + break; | ||
| 495 | + } | ||
| 496 | + case 42: { | ||
| 497 | + launchTCOLCMAXCase42<<<1, nullptr, stream>>>((uint32_t *)out, (int16_t *)src); | ||
| 498 | + break; | ||
| 499 | + } | ||
| 500 | + case 43: { | ||
| 501 | + launchTCOLCMAXCase43<<<1, nullptr, stream>>>((uint32_t *)out, (int16_t *)src); | ||
| 502 | + break; | ||
| 503 | + } | ||
| 175 | case 51: { | 504 | case 51: { |
| 176 | launchTCOLCMAXCase51<<<1, nullptr, stream>>>((uint32_t *)out, (uint16_t *)src); | 505 | launchTCOLCMAXCase51<<<1, nullptr, stream>>>((uint32_t *)out, (uint16_t *)src); |
| 177 | break; | 506 | break; |
| @@ -184,6 +513,18 @@ void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream) | |||
| 184 | launchTCOLCMAXCase53<<<1, nullptr, stream>>>((uint32_t *)out, (uint16_t *)src); | 513 | launchTCOLCMAXCase53<<<1, nullptr, stream>>>((uint32_t *)out, (uint16_t *)src); |
| 185 | break; | 514 | break; |
| 186 | } | 515 | } |
| 516 | + case 61: { | ||
| 517 | + launchTCOLCMAXCase61<<<1, nullptr, stream>>>((uint32_t *)out, (int32_t *)src); | ||
| 518 | + break; | ||
| 519 | + } | ||
| 520 | + case 62: { | ||
| 521 | + launchTCOLCMAXCase62<<<1, nullptr, stream>>>((uint32_t *)out, (int32_t *)src); | ||
| 522 | + break; | ||
| 523 | + } | ||
| 524 | + case 63: { | ||
| 525 | + launchTCOLCMAXCase63<<<1, nullptr, stream>>>((uint32_t *)out, (int32_t *)src); | ||
| 526 | + break; | ||
| 527 | + } | ||
| 187 | case 71: { | 528 | case 71: { |
| 188 | launchTCOLCMAXCase71<<<1, nullptr, stream>>>((uint32_t *)out, (uint32_t *)src); | 529 | launchTCOLCMAXCase71<<<1, nullptr, stream>>>((uint32_t *)out, (uint32_t *)src); |
| 189 | break; | 530 | break; |
| @@ -196,20 +537,241 @@ void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream) | |||
| 196 | launchTCOLCMAXCase73<<<1, nullptr, stream>>>((uint32_t *)out, (uint32_t *)src); | 537 | launchTCOLCMAXCase73<<<1, nullptr, stream>>>((uint32_t *)out, (uint32_t *)src); |
| 197 | break; | 538 | break; |
| 198 | } | 539 | } |
| 540 | + case 81: { | ||
| 541 | + launchTCOLCMAXCase81<<<1, nullptr, stream>>>((uint32_t *)out, (half *)src); | ||
| 542 | + break; | ||
| 543 | + } | ||
| 544 | + case 82: { | ||
| 545 | + launchTCOLCMAXCase82<<<1, nullptr, stream>>>((uint32_t *)out, (uint16_t *)src); | ||
| 546 | + break; | ||
| 547 | + } | ||
| 548 | + case 83: { | ||
| 549 | + launchTCOLCMAXCase83<<<1, nullptr, stream>>>((uint32_t *)out, (uint32_t *)src); | ||
| 550 | + break; | ||
| 551 | + } | ||
| 552 | + case 84: { | ||
| 553 | + launchTCOLCMAXCase84<<<1, nullptr, stream>>>((uint32_t *)out, (float *)src); | ||
| 554 | + break; | ||
| 555 | + } | ||
| 556 | + case 85: { | ||
| 557 | + launchTCOLCMAXCase85<<<1, nullptr, stream>>>((uint32_t *)out, (int8_t *)src); | ||
| 558 | + break; | ||
| 559 | + } | ||
| 560 | + case 86: { | ||
| 561 | + launchTCOLCMAXCase86<<<1, nullptr, stream>>>((uint32_t *)out, (uint8_t *)src); | ||
| 562 | + break; | ||
| 563 | + } | ||
| 564 | + case 87: { | ||
| 565 | + launchTCOLCMAXCase87<<<1, nullptr, stream>>>((uint32_t *)out, (int16_t *)src); | ||
| 566 | + break; | ||
| 567 | + } | ||
| 568 | + case 88: { | ||
| 569 | + launchTCOLCMAXCase88<<<1, nullptr, stream>>>((uint32_t *)out, (int32_t *)src); | ||
| 570 | + break; | ||
| 571 | + } | ||
| 572 | + case 91: { | ||
| 573 | + launchTCOLCMAXCase91<<<1, nullptr, stream>>>((uint32_t *)out, (uint16_t *)src); | ||
| 574 | + break; | ||
| 575 | + } | ||
| 576 | + case 92: { | ||
| 577 | + launchTCOLCMAXCase92<<<1, nullptr, stream>>>((uint32_t *)out, (half *)src); | ||
| 578 | + break; | ||
| 579 | + } | ||
| 580 | + case 93: { | ||
| 581 | + launchTCOLCMAXCase93<<<1, nullptr, stream>>>((uint32_t *)out, (uint16_t *)src); | ||
| 582 | + break; | ||
| 583 | + } | ||
| 199 | default: { | 584 | default: { |
| 200 | } | 585 | } |
| 201 | } | 586 | } |
| 202 | } | 587 | } |
| 203 | 588 | ||
| 589 | +// ============================================================================= | ||
| 590 | +// Value + index dispatcher | ||
| 591 | +// ============================================================================= | ||
| 592 | +template <uint32_t caseId> | ||
| 593 | +void launchTCOLIDXVALMAXCase(void *outVal, void *outIdx, void *src, aclrtStream stream) | ||
| 594 | +{ | ||
| 595 | + switch (caseId) { | ||
| 596 | + case 1: { | ||
| 597 | + launchTCOLIDXVALMAXCase001<<<1, nullptr, stream>>>((float *)outVal, (int32_t *)outIdx, (float *)src); | ||
| 598 | + break; | ||
| 599 | + } | ||
| 600 | + case 2: { | ||
| 601 | + launchTCOLIDXVALMAXCase002<<<1, nullptr, stream>>>((float *)outVal, (int32_t *)outIdx, (float *)src); | ||
| 602 | + break; | ||
| 603 | + } | ||
| 604 | + case 3: { | ||
| 605 | + launchTCOLIDXVALMAXCase003<<<1, nullptr, stream>>>((float *)outVal, (int32_t *)outIdx, (float *)src); | ||
| 606 | + break; | ||
| 607 | + } | ||
| 608 | + case 11: { | ||
| 609 | + launchTCOLIDXVALMAXCase011<<<1, nullptr, stream>>>((half *)outVal, (int16_t *)outIdx, (half *)src); | ||
| 610 | + break; | ||
| 611 | + } | ||
| 612 | + case 12: { | ||
| 613 | + launchTCOLIDXVALMAXCase012<<<1, nullptr, stream>>>((half *)outVal, (int16_t *)outIdx, (half *)src); | ||
| 614 | + break; | ||
| 615 | + } | ||
| 616 | + case 13: { | ||
| 617 | + launchTCOLIDXVALMAXCase013<<<1, nullptr, stream>>>((half *)outVal, (int16_t *)outIdx, (half *)src); | ||
| 618 | + break; | ||
| 619 | + } | ||
| 620 | + case 41: { | ||
| 621 | + launchTCOLIDXVALMAXCase041<<<1, nullptr, stream>>>((int16_t *)outVal, (int16_t *)outIdx, (int16_t *)src); | ||
| 622 | + break; | ||
| 623 | + } | ||
| 624 | + case 42: { | ||
| 625 | + launchTCOLIDXVALMAXCase042<<<1, nullptr, stream>>>((int16_t *)outVal, (int16_t *)outIdx, (int16_t *)src); | ||
| 626 | + break; | ||
| 627 | + } | ||
| 628 | + case 43: { | ||
| 629 | + launchTCOLIDXVALMAXCase043<<<1, nullptr, stream>>>((int16_t *)outVal, (int16_t *)outIdx, (int16_t *)src); | ||
| 630 | + break; | ||
| 631 | + } | ||
| 632 | + case 51: { | ||
| 633 | + launchTCOLIDXVALMAXCase051<<<1, nullptr, stream>>>((uint16_t *)outVal, (int16_t *)outIdx, (uint16_t *)src); | ||
| 634 | + break; | ||
| 635 | + } | ||
| 636 | + case 52: { | ||
| 637 | + launchTCOLIDXVALMAXCase052<<<1, nullptr, stream>>>((uint16_t *)outVal, (int16_t *)outIdx, (uint16_t *)src); | ||
| 638 | + break; | ||
| 639 | + } | ||
| 640 | + case 53: { | ||
| 641 | + launchTCOLIDXVALMAXCase053<<<1, nullptr, stream>>>((uint16_t *)outVal, (int16_t *)outIdx, (uint16_t *)src); | ||
| 642 | + break; | ||
| 643 | + } | ||
| 644 | + case 61: { | ||
| 645 | + launchTCOLIDXVALMAXCase061<<<1, nullptr, stream>>>((int32_t *)outVal, (int32_t *)outIdx, (int32_t *)src); | ||
| 646 | + break; | ||
| 647 | + } | ||
| 648 | + case 62: { | ||
| 649 | + launchTCOLIDXVALMAXCase062<<<1, nullptr, stream>>>((int32_t *)outVal, (int32_t *)outIdx, (int32_t *)src); | ||
| 650 | + break; | ||
| 651 | + } | ||
| 652 | + case 63: { | ||
| 653 | + launchTCOLIDXVALMAXCase063<<<1, nullptr, stream>>>((int32_t *)outVal, (int32_t *)outIdx, (int32_t *)src); | ||
| 654 | + break; | ||
| 655 | + } | ||
| 656 | + case 71: { | ||
| 657 | + launchTCOLIDXVALMAXCase071<<<1, nullptr, stream>>>((uint32_t *)outVal, (int32_t *)outIdx, (uint32_t *)src); | ||
| 658 | + break; | ||
| 659 | + } | ||
| 660 | + case 72: { | ||
| 661 | + launchTCOLIDXVALMAXCase072<<<1, nullptr, stream>>>((uint32_t *)outVal, (int32_t *)outIdx, (uint32_t *)src); | ||
| 662 | + break; | ||
| 663 | + } | ||
| 664 | + case 73: { | ||
| 665 | + launchTCOLIDXVALMAXCase073<<<1, nullptr, stream>>>((uint32_t *)outVal, (int32_t *)outIdx, (uint32_t *)src); | ||
| 666 | + break; | ||
| 667 | + } | ||
| 668 | + case 81: { | ||
| 669 | + launchTCOLIDXVALMAXCase081<<<1, nullptr, stream>>>((half *)outVal, (int16_t *)outIdx, (half *)src); | ||
| 670 | + break; | ||
| 671 | + } | ||
| 672 | + case 82: { | ||
| 673 | + launchTCOLIDXVALMAXCase082<<<1, nullptr, stream>>>((uint16_t *)outVal, (int16_t *)outIdx, (uint16_t *)src); | ||
| 674 | + break; | ||
| 675 | + } | ||
| 676 | + case 83: { | ||
| 677 | + launchTCOLIDXVALMAXCase083<<<1, nullptr, stream>>>((uint32_t *)outVal, (int32_t *)outIdx, (uint32_t *)src); | ||
| 678 | + break; | ||
| 679 | + } | ||
| 680 | + case 84: { | ||
| 681 | + launchTCOLIDXVALMAXCase084<<<1, nullptr, stream>>>((float *)outVal, (int32_t *)outIdx, (float *)src); | ||
| 682 | + break; | ||
| 683 | + } | ||
| 684 | + case 85: { | ||
| 685 | + launchTCOLIDXVALMAXCase085<<<1, nullptr, stream>>>((int16_t *)outVal, (int16_t *)outIdx, (int16_t *)src); | ||
| 686 | + break; | ||
| 687 | + } | ||
| 688 | + case 86: { | ||
| 689 | + launchTCOLIDXVALMAXCase086<<<1, nullptr, stream>>>((int32_t *)outVal, (int32_t *)outIdx, (int32_t *)src); | ||
| 690 | + break; | ||
| 691 | + } | ||
| 692 | + case 91: { | ||
| 693 | + launchTCOLIDXVALMAXCase091<<<1, nullptr, stream>>>((uint16_t *)outVal, (int16_t *)outIdx, (uint16_t *)src); | ||
| 694 | + break; | ||
| 695 | + } | ||
| 696 | + case 92: { | ||
| 697 | + launchTCOLIDXVALMAXCase092<<<1, nullptr, stream>>>((half *)outVal, (int16_t *)outIdx, (half *)src); | ||
| 698 | + break; | ||
| 699 | + } | ||
| 700 | + case 93: { | ||
| 701 | + launchTCOLIDXVALMAXCase093<<<1, nullptr, stream>>>((uint16_t *)outVal, (int16_t *)outIdx, (uint16_t *)src); | ||
| 702 | + break; | ||
| 703 | + } | ||
| 704 | + default: { | ||
| 705 | + } | ||
| 706 | + } | ||
| 707 | +} | ||
| 708 | + | ||
| 709 | +// ============================================================================= | ||
| 710 | +// Pure index template instantiations | ||
| 711 | +// ============================================================================= | ||
| 204 | template void launchTCOLCMAXTestCase<1>(void *out, void *src, aclrtStream stream); | 712 | template void launchTCOLCMAXTestCase<1>(void *out, void *src, aclrtStream stream); |
| 205 | template void launchTCOLCMAXTestCase<2>(void *out, void *src, aclrtStream stream); | 713 | template void launchTCOLCMAXTestCase<2>(void *out, void *src, aclrtStream stream); |
| 206 | template void launchTCOLCMAXTestCase<3>(void *out, void *src, aclrtStream stream); | 714 | template void launchTCOLCMAXTestCase<3>(void *out, void *src, aclrtStream stream); |
| 207 | template void launchTCOLCMAXTestCase<11>(void *out, void *src, aclrtStream stream); | 715 | template void launchTCOLCMAXTestCase<11>(void *out, void *src, aclrtStream stream); |
| 208 | template void launchTCOLCMAXTestCase<12>(void *out, void *src, aclrtStream stream); | 716 | template void launchTCOLCMAXTestCase<12>(void *out, void *src, aclrtStream stream); |
| 209 | template void launchTCOLCMAXTestCase<13>(void *out, void *src, aclrtStream stream); | 717 | template void launchTCOLCMAXTestCase<13>(void *out, void *src, aclrtStream stream); |
| 718 | +template void launchTCOLCMAXTestCase<21>(void *out, void *src, aclrtStream stream); | ||
| 719 | +template void launchTCOLCMAXTestCase<22>(void *out, void *src, aclrtStream stream); | ||
| 720 | +template void launchTCOLCMAXTestCase<23>(void *out, void *src, aclrtStream stream); | ||
| 721 | +template void launchTCOLCMAXTestCase<31>(void *out, void *src, aclrtStream stream); | ||
| 722 | +template void launchTCOLCMAXTestCase<32>(void *out, void *src, aclrtStream stream); | ||
| 723 | +template void launchTCOLCMAXTestCase<33>(void *out, void *src, aclrtStream stream); | ||
| 724 | +template void launchTCOLCMAXTestCase<41>(void *out, void *src, aclrtStream stream); | ||
| 725 | +template void launchTCOLCMAXTestCase<42>(void *out, void *src, aclrtStream stream); | ||
| 726 | +template void launchTCOLCMAXTestCase<43>(void *out, void *src, aclrtStream stream); | ||
| 210 | template void launchTCOLCMAXTestCase<51>(void *out, void *src, aclrtStream stream); | 727 | template void launchTCOLCMAXTestCase<51>(void *out, void *src, aclrtStream stream); |
| 211 | template void launchTCOLCMAXTestCase<52>(void *out, void *src, aclrtStream stream); | 728 | template void launchTCOLCMAXTestCase<52>(void *out, void *src, aclrtStream stream); |
| 212 | template void launchTCOLCMAXTestCase<53>(void *out, void *src, aclrtStream stream); | 729 | template void launchTCOLCMAXTestCase<53>(void *out, void *src, aclrtStream stream); |
| 730 | +template void launchTCOLCMAXTestCase<61>(void *out, void *src, aclrtStream stream); | ||
| 731 | +template void launchTCOLCMAXTestCase<62>(void *out, void *src, aclrtStream stream); | ||
| 732 | +template void launchTCOLCMAXTestCase<63>(void *out, void *src, aclrtStream stream); | ||
| 213 | template void launchTCOLCMAXTestCase<71>(void *out, void *src, aclrtStream stream); | 733 | template void launchTCOLCMAXTestCase<71>(void *out, void *src, aclrtStream stream); |
| 214 | template void launchTCOLCMAXTestCase<72>(void *out, void *src, aclrtStream stream); | 734 | template void launchTCOLCMAXTestCase<72>(void *out, void *src, aclrtStream stream); |
| 215 | -template void launchTCOLCMAXTestCase<73>(void *out, void *src, aclrtStream stream); | 735 | +template void launchTCOLCMAXTestCase<73>(void *out, void *src, aclrtStream stream); |
| 736 | +template void launchTCOLCMAXTestCase<81>(void *out, void *src, aclrtStream stream); | ||
| 737 | +template void launchTCOLCMAXTestCase<82>(void *out, void *src, aclrtStream stream); | ||
| 738 | +template void launchTCOLCMAXTestCase<83>(void *out, void *src, aclrtStream stream); | ||
| 739 | +template void launchTCOLCMAXTestCase<84>(void *out, void *src, aclrtStream stream); | ||
| 740 | +template void launchTCOLCMAXTestCase<85>(void *out, void *src, aclrtStream stream); | ||
| 741 | +template void launchTCOLCMAXTestCase<86>(void *out, void *src, aclrtStream stream); | ||
| 742 | +template void launchTCOLCMAXTestCase<87>(void *out, void *src, aclrtStream stream); | ||
| 743 | +template void launchTCOLCMAXTestCase<88>(void *out, void *src, aclrtStream stream); | ||
| 744 | +template void launchTCOLCMAXTestCase<91>(void *out, void *src, aclrtStream stream); | ||
| 745 | +template void launchTCOLCMAXTestCase<92>(void *out, void *src, aclrtStream stream); | ||
| 746 | +template void launchTCOLCMAXTestCase<93>(void *out, void *src, aclrtStream stream); | ||
| 747 | + | ||
| 748 | +// ============================================================================= | ||
| 749 | +// Value + index template instantiations | ||
| 750 | +// ============================================================================= | ||
| 751 | +template void launchTCOLIDXVALMAXCase<1>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 752 | +template void launchTCOLIDXVALMAXCase<2>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 753 | +template void launchTCOLIDXVALMAXCase<3>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 754 | +template void launchTCOLIDXVALMAXCase<11>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 755 | +template void launchTCOLIDXVALMAXCase<12>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 756 | +template void launchTCOLIDXVALMAXCase<13>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 757 | +template void launchTCOLIDXVALMAXCase<41>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 758 | +template void launchTCOLIDXVALMAXCase<42>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 759 | +template void launchTCOLIDXVALMAXCase<43>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 760 | +template void launchTCOLIDXVALMAXCase<51>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 761 | +template void launchTCOLIDXVALMAXCase<52>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 762 | +template void launchTCOLIDXVALMAXCase<53>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 763 | +template void launchTCOLIDXVALMAXCase<61>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 764 | +template void launchTCOLIDXVALMAXCase<62>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 765 | +template void launchTCOLIDXVALMAXCase<63>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 766 | +template void launchTCOLIDXVALMAXCase<71>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 767 | +template void launchTCOLIDXVALMAXCase<72>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 768 | +template void launchTCOLIDXVALMAXCase<73>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 769 | +template void launchTCOLIDXVALMAXCase<81>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 770 | +template void launchTCOLIDXVALMAXCase<82>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 771 | +template void launchTCOLIDXVALMAXCase<83>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 772 | +template void launchTCOLIDXVALMAXCase<84>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 773 | +template void launchTCOLIDXVALMAXCase<85>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 774 | +template void launchTCOLIDXVALMAXCase<86>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 775 | +template void launchTCOLIDXVALMAXCase<91>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 776 | +template void launchTCOLIDXVALMAXCase<92>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||
| 777 | +template void launchTCOLIDXVALMAXCase<93>(void *outVal, void *outIdx, void *src, aclrtStream stream); | ||