已合并
update TColReduceIdx to support synchronous output of value and index #928
سقط 落创建于 5月15日
update TColReduceIdx to support synchronous output of value and index #928
已合并
سقط 落创建于 5月15日
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+ 
1204template <typename TileDataOut, typename TileDataIn, typename TileDataTmp, typename... WaitEvents>1224template <typename TileDataOut, typename TileDataIn, typename TileDataTmp, typename... WaitEvents>
1205PTO_INST RecordEvent TROWMAX(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp, WaitEvents &...events)1225PTO_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
8See LICENSE in the root of the software repository for the full text of the License.8See LICENSE in the root of the software repository for the full text of the License.
9*/9*/
10 10 
11-#ifndef TCOLREDUCEIDX_HPP11+#ifndef T_COL_REDUCE_IDX_OPS_HPP
12-#define TCOLREDUCEIDX_HPP12+#define T_COL_REDUCE_IDX_OPS_HPP
13 13 
14#include <pto/common/utils.hpp>14#include <pto/common/utils.hpp>
15#include <pto/common/type.hpp>15#include <pto/common/type.hpp>
16 16 
17namespace pto {17namespace pto {
18-template <typename TileDataOut, typename TileDataIn, typename TileDataTmp>18+ 
19+template <typename TileDataOutVal, typename TileDataOutIdx, typename TileDataIn, typename TileDataTmp,
20+ bool WithVal = false>
19PTO_INTERNAL void TColReduceIdxCheck(unsigned srcValidRow, unsigned srcValidCol, unsigned dstValidRow,21PTO_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); // Min297+ 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); // Max303+ 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 pto320} // namespace pto
249#endif321#endif
@@ -17,248 +17,313 @@ See LICENSE in the root of the software repository for the full text of the Lice
17#include "utils.hpp"17#include "utils.hpp"
18 18 
19namespace pto {19namespace 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>
21PTO_INTERNAL void TColReduceIdxCheck(unsigned srcValidRow, unsigned srcValidCol, unsigned dstValidRow,60PTO_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+// ----------------------------------------------------------------------------
41template <typename TileDataOut, typename TileDataIn, bool IsArgMax>118template <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-}
238template <typename TileDataOut, typename TileDataIn, bool IsArgMax>280template <typename TileDataOut, typename TileDataIn, bool IsArgMax>
239PTO_INTERNAL void TCOLARG_DISPATCH(TileDataOut &dst, TileDataIn &src)281PTO_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+// ==========================================================================================
253template <typename TileDataOut, typename TileDataIn, typename TileDataTmp>304template <typename TileDataOut, typename TileDataIn, typename TileDataTmp>
254PTO_INTERNAL void TCOLARGMIN_IMPL(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp)305PTO_INTERNAL void TCOLARGMIN_IMPL(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp)
255{306{
256- TCOLARG_DISPATCH<TileDataOut, TileDataIn, false>(dst, src); // Min307+ TCOLARG_DISPATCH<TileDataOut, TileDataIn, false>(dst, src);
257}308}
309+ 
258template <typename TileDataOut, typename TileDataIn, typename TileDataTmp>310template <typename TileDataOut, typename TileDataIn, typename TileDataTmp>
259PTO_INTERNAL void TCOLARGMAX_IMPL(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp)311PTO_INTERNAL void TCOLARGMAX_IMPL(TileDataOut &dst, TileDataIn &src, TileDataTmp &tmp)
260{312{
261- TCOLARG_DISPATCH<TileDataOut, TileDataIn, true>(dst, src); // Max313+ 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 pto328} // namespace pto
264#endif329#endif
@@ -13,6 +13,7 @@
13import os13import os
14import numpy as np14import numpy as np
15import math15import math
16+ 
16np.random.seed(19)17np.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:] = 036 output_arr[valid_col:] = 0
36- dst_col = math.ceil(valid_col / 8) * 837+ 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 
44class TColCMaxParams:57class 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 = name59 self.name = name
47 self.data_type = data_type60 self.data_type = data_type
48 self.row = row61 self.row = row
49 self.valid_row = valid_row62 self.valid_row = valid_row
50 self.col = col63 self.col = col
51 self.valid_col = valid_col64 self.valid_col = valid_col
65+ self.idx = idx
66+ 
52 67 
53if __name__ == "__main__":68if __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;
18template <uint32_t caseId>18template <uint32_t caseId>
19void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream);19void 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+ 
21std::string GetGoldenDir()24std::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+ 
38protected:46protected:
39 void SetUp() override47 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 
96TEST_F(TCOLCMAXTest, case01)162TEST_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+ 
51extern "C" __global__ AICORE void launchTCOLCMAXCase01(__gm__ uint32_t *out, __gm__ float *src)96extern "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+ 
128template <uint32_t caseId>269template <uint32_t caseId>
129void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream)270void 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+ 
213template void launchTCOLCMAXTestCase<1>(void *out, void *src, aclrtStream stream);439template void launchTCOLCMAXTestCase<1>(void *out, void *src, aclrtStream stream);
214template void launchTCOLCMAXTestCase<2>(void *out, void *src, aclrtStream stream);440template void launchTCOLCMAXTestCase<2>(void *out, void *src, aclrtStream stream);
215template void launchTCOLCMAXTestCase<3>(void *out, void *src, aclrtStream stream);441template void launchTCOLCMAXTestCase<3>(void *out, void *src, aclrtStream stream);
@@ -228,4 +454,24 @@ template void launchTCOLCMAXTestCase<83>(void *out, void *src, aclrtStream strea
228template void launchTCOLCMAXTestCase<84>(void *out, void *src, aclrtStream stream);454template void launchTCOLCMAXTestCase<84>(void *out, void *src, aclrtStream stream);
229template void launchTCOLCMAXTestCase<91>(void *out, void *src, aclrtStream stream);455template void launchTCOLCMAXTestCase<91>(void *out, void *src, aclrtStream stream);
230template void launchTCOLCMAXTestCase<92>(void *out, void *src, aclrtStream stream);456template 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/python31#!/usr/bin/python3
2# coding=utf-82# 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 of5+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 
13import os13import os
14+import math
14import numpy as np15import numpy as np
16+ 
15np.random.seed(19)17np.random.seed(19)
16 18 
17 19 
@@ -23,44 +25,146 @@ def gen_golden_data(param):
23 valid_col = param.valid_col25 valid_col = param.valid_col
24 value_max = 10026 value_max = 100
25 value_min = -10027 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 = 20029 value_max = 200
28 value_min = 030 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 = 1035 value_max = 10
31 value_min = 036 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:] = 039+ 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 
41class TColCMaxParams:66class 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 = name68 self.name = name
44 self.data_type = data_type69 self.data_type = data_type
45 self.row = row70 self.row = row
46 self.valid_row = valid_row71 self.valid_row = valid_row
47 self.col = col72 self.col = col
48 self.valid_col = valid_col73 self.valid_col = valid_col
74+ self.idx = idx
75+ 
49 76 
50if __name__ == "__main__":77if __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
15using namespace std;15using namespace std;
16using namespace PtoTestCommon;16using namespace PtoTestCommon;
17 17 
18+// =============================================================================
19+// Pure index dispatcher (3-arg TCOLARGMAX)
20+// =============================================================================
18template <uint32_t caseId>21template <uint32_t caseId>
19void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream);22void 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+// =============================================================================
21std::string GetGoldenDir()31std::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+// =============================================================================
30class TCOLCMAXTest : public testing::Test {43class TCOLCMAXTest : public testing::Test {
31public:44public:
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+ 
38protected:56protected:
39 void SetUp() override57 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+// =============================================================================
95TEST_F(TCOLCMAXTest, case01)180TEST_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+ 
110TEST_F(TCOLCMAXTest, case11)196TEST_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+ 
125TEST_F(TCOLCMAXTest, case51)260TEST_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+ 
140TEST_F(TCOLCMAXTest, case71)292TEST_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
15using namespace std;15using namespace std;
16using namespace pto;16using namespace pto;
17 17 
18+// =============================================================================
19+// Pure index mode kernel (3-arg TCOLARGMAX)
20+// =============================================================================
18template <typename T, int srcRow, int srcValidRow, int dstRow, int col, int validCol>21template <typename T, int srcRow, int srcValidRow, int dstRow, int col, int validCol>
19PTO_INTERNAL void runTColCMax(__gm__ uint32_t __out__ *out, __gm__ T __in__ *src, bool isBinary)22PTO_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+// =============================================================================
50extern "C" __global__ AICORE void launchTCOLCMAXCase01(__gm__ uint32_t *out, __gm__ float *src)101extern "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+// =============================================================================
62extern "C" __global__ AICORE void launchTCOLCMAXCase11(__gm__ uint32_t *out, __gm__ half *src)117extern "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+// =============================================================================
74extern "C" __global__ AICORE void launchTCOLCMAXCase21(__gm__ uint32_t *out, __gm__ int8_t *src)133extern "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+// =============================================================================
86extern "C" __global__ AICORE void launchTCOLCMAXCase31(__gm__ uint32_t *out, __gm__ uint8_t *src)149extern "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+// =============================================================================
98extern "C" __global__ AICORE void launchTCOLCMAXCase41(__gm__ uint32_t *out, __gm__ int16_t *src)165extern "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+// =============================================================================
110extern "C" __global__ AICORE void launchTCOLCMAXCase51(__gm__ uint32_t *out, __gm__ uint16_t *src)181extern "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+// =============================================================================
122extern "C" __global__ AICORE void launchTCOLCMAXCase61(__gm__ uint32_t *out, __gm__ int32_t *src)197extern "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+// =============================================================================
134extern "C" __global__ AICORE void launchTCOLCMAXCase71(__gm__ uint32_t *out, __gm__ uint32_t *src)213extern "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+// =============================================================================
147template <uint32_t caseId>440template <uint32_t caseId>
148void launchTCOLCMAXTestCase(void *out, void *src, aclrtStream stream)441void 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+// =============================================================================
204template void launchTCOLCMAXTestCase<1>(void *out, void *src, aclrtStream stream);712template void launchTCOLCMAXTestCase<1>(void *out, void *src, aclrtStream stream);
205template void launchTCOLCMAXTestCase<2>(void *out, void *src, aclrtStream stream);713template void launchTCOLCMAXTestCase<2>(void *out, void *src, aclrtStream stream);
206template void launchTCOLCMAXTestCase<3>(void *out, void *src, aclrtStream stream);714template void launchTCOLCMAXTestCase<3>(void *out, void *src, aclrtStream stream);
207template void launchTCOLCMAXTestCase<11>(void *out, void *src, aclrtStream stream);715template void launchTCOLCMAXTestCase<11>(void *out, void *src, aclrtStream stream);
208template void launchTCOLCMAXTestCase<12>(void *out, void *src, aclrtStream stream);716template void launchTCOLCMAXTestCase<12>(void *out, void *src, aclrtStream stream);
209template void launchTCOLCMAXTestCase<13>(void *out, void *src, aclrtStream stream);717template 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);
210template void launchTCOLCMAXTestCase<51>(void *out, void *src, aclrtStream stream);727template void launchTCOLCMAXTestCase<51>(void *out, void *src, aclrtStream stream);
211template void launchTCOLCMAXTestCase<52>(void *out, void *src, aclrtStream stream);728template void launchTCOLCMAXTestCase<52>(void *out, void *src, aclrtStream stream);
212template void launchTCOLCMAXTestCase<53>(void *out, void *src, aclrtStream stream);729template 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);
213template void launchTCOLCMAXTestCase<71>(void *out, void *src, aclrtStream stream);733template void launchTCOLCMAXTestCase<71>(void *out, void *src, aclrtStream stream);
214template void launchTCOLCMAXTestCase<72>(void *out, void *src, aclrtStream stream);734template 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);