已合并
replace MicroAPI namespace to Reg #834
li-xingyue-lxy创建于 3月11日
replace MicroAPI namespace to Reg #834
已合并
共 249 个文件变更+20824-20819
| @@ -37,6 +37,11 @@ asc-devkit: | |||
| 37 | - impl/basic_api/reg_compute/kernel_reg_compute_common_intf_impl.h | 37 | - impl/basic_api/reg_compute/kernel_reg_compute_common_intf_impl.h |
| 38 | - impl/basic_api/reg_compute/kernel_reg_compute_intf_impl.h | 38 | - impl/basic_api/reg_compute/kernel_reg_compute_intf_impl.h |
| 39 | - impl/basic_api/reg_compute/dav_c310/kernel_reg_compute_datatype_impl.h | 39 | - impl/basic_api/reg_compute/dav_c310/kernel_reg_compute_datatype_impl.h |
| 40 | + - impl/adv_api/detail/math/**/*.h | ||
| 41 | + - impl/adv_api/detail/quantization/**/*.h | ||
| 42 | + - impl/adv_api/detail/activation/**/*.h | ||
| 43 | + - impl/adv_api/detail/normalization/**/*.h | ||
| 44 | + - impl/adv_api/detail/common/common.h | ||
| 40 | llt: | 45 | llt: |
| 41 | ut_check: true | 46 | ut_check: true |
| 42 | st_check: false | 47 | st_check: false |
| @@ -25,16 +25,16 @@ template <typename T> | |||
| 25 | __simd_vf__ inline void AddVF(__ubuf__ T* dstAddr, __ubuf__ T* src0Addr, __ubuf__ T* src1Addr, uint32_t count, | 25 | __simd_vf__ inline void AddVF(__ubuf__ T* dstAddr, __ubuf__ T* src0Addr, __ubuf__ T* src1Addr, uint32_t count, |
| 26 | uint32_t oneRepeatSize, uint16_t repeatTimes) | 26 | uint32_t oneRepeatSize, uint16_t repeatTimes) |
| 27 | { | 27 | { |
| 28 | - AscendC::MicroAPI::RegTensor<T> srcReg0; | 28 | + AscendC::Reg::RegTensor<T> srcReg0; |
| 29 | - AscendC::MicroAPI::RegTensor<T> srcReg1; | 29 | + AscendC::Reg::RegTensor<T> srcReg1; |
| 30 | - AscendC::MicroAPI::RegTensor<T> dstReg; | 30 | + AscendC::Reg::RegTensor<T> dstReg; |
| 31 | - AscendC::MicroAPI::MaskReg mask; | 31 | + AscendC::Reg::MaskReg mask; |
| 32 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 32 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 33 | - mask = AscendC::MicroAPI::UpdateMask<T>(count); | 33 | + mask = AscendC::Reg::UpdateMask<T>(count); |
| 34 | - AscendC::MicroAPI::LoadAlign(srcReg0, src0Addr + i * oneRepeatSize); | 34 | + AscendC::Reg::LoadAlign(srcReg0, src0Addr + i * oneRepeatSize); |
| 35 | - AscendC::MicroAPI::LoadAlign(srcReg1, src1Addr + i * oneRepeatSize); | 35 | + AscendC::Reg::LoadAlign(srcReg1, src1Addr + i * oneRepeatSize); |
| 36 | - AscendC::MicroAPI::Add(dstReg, srcReg0, srcReg1, mask); | 36 | + AscendC::Reg::Add(dstReg, srcReg0, srcReg1, mask); |
| 37 | - AscendC::MicroAPI::StoreAlign(dstAddr + i * oneRepeatSize, dstReg, mask); | 37 | + AscendC::Reg::StoreAlign(dstAddr + i * oneRepeatSize, dstReg, mask); |
| 38 | } | 38 | } |
| 39 | } | 39 | } |
| 40 | 40 | ||
| @@ -63,12 +63,12 @@ | |||
| 63 | gather_output[threadIdx.x] = input[gather_idx]; | 63 | gather_output[threadIdx.x] = input[gather_idx]; |
| 64 | ``` | 64 | ``` |
| 65 | 65 | ||
| 66 | - simd_adds负责将Local Memory中数据做加1操作。使用MicroAPI::LoadAlign将数据从Local Memory搬运到寄存器上,调用MicroAPI::Adds完成加1运算并输出到目标寄存器,最后调用MicroAPI::StoreAlign将数据从寄存器搬运到Local Memory。重复上述操作即可完成1024个数据的加1运算。 | 66 | + simd_adds负责将Local Memory中数据做加1操作。使用Reg::LoadAlign将数据从Local Memory搬运到寄存器上,调用Reg::Adds完成加1运算并输出到目标寄存器,最后调用Reg::StoreAlign将数据从寄存器搬运到Local Memory。重复上述操作即可完成1024个数据的加1运算。 |
| 67 | ``` | 67 | ``` |
| 68 | for (uint16_t i = 0; i < repeat_times; i++) { | 68 | for (uint16_t i = 0; i < repeat_times; i++) { |
| 69 | - AscendC::MicroAPI::LoadAlign(src_reg0, input + i * one_repeat_size); | 69 | + AscendC::Reg::LoadAlign(src_reg0, input + i * one_repeat_size); |
| 70 | - AscendC::MicroAPI::Adds(dst_reg0, src_reg0, ADDS_ADDEND, mask_reg); | 70 | + AscendC::Reg::Adds(dst_reg0, src_reg0, ADDS_ADDEND, mask_reg); |
| 71 | - AscendC::MicroAPI::StoreAlign(output + i * one_repeat_size, dst_reg0, mask_reg); | 71 | + AscendC::Reg::StoreAlign(output + i * one_repeat_size, dst_reg0, mask_reg); |
| 72 | } | 72 | } |
| 73 | ``` | 73 | ``` |
| 74 | 74 | ||
| @@ -61,20 +61,20 @@ __simt_vf__ __launch_bounds__(THREAD_COUNT) inline void simt_gather( | |||
| 61 | __simd_vf__ inline void simd_adds(__ubuf__ float* output, __ubuf__ float* input, | 61 | __simd_vf__ inline void simd_adds(__ubuf__ float* output, __ubuf__ float* input, |
| 62 | uint32_t count, uint32_t one_repeat_size, uint16_t repeat_times) | 62 | uint32_t count, uint32_t one_repeat_size, uint16_t repeat_times) |
| 63 | { | 63 | { |
| 64 | - AscendC::MicroAPI::RegTensor<float> src_reg0; | 64 | + AscendC::Reg::RegTensor<float> src_reg0; |
| 65 | - AscendC::MicroAPI::RegTensor<float> dst_reg0; | 65 | + AscendC::Reg::RegTensor<float> dst_reg0; |
| 66 | // asc_update_mask() will be supported later. | 66 | // asc_update_mask() will be supported later. |
| 67 | // init MaskReg with the count of all numbers. | 67 | // init MaskReg with the count of all numbers. |
| 68 | - AscendC::MicroAPI::MaskReg mask_reg; | 68 | + AscendC::Reg::MaskReg mask_reg; |
| 69 | 69 | ||
| 70 | for (uint16_t i = 0; i < repeat_times; i++) { | 70 | for (uint16_t i = 0; i < repeat_times; i++) { |
| 71 | - mask_reg = AscendC::MicroAPI::UpdateMask<float>(count); | 71 | + mask_reg = AscendC::Reg::UpdateMask<float>(count); |
| 72 | // asc_load, asc_adds and asc_store will be supported later. | 72 | // asc_load, asc_adds and asc_store will be supported later. |
| 73 | // load data from UB to RegTensor. | 73 | // load data from UB to RegTensor. |
| 74 | - AscendC::MicroAPI::LoadAlign(src_reg0, input + i * one_repeat_size); | 74 | + AscendC::Reg::LoadAlign(src_reg0, input + i * one_repeat_size); |
| 75 | - AscendC::MicroAPI::Adds(dst_reg0, src_reg0, ADDS_ADDEND, mask_reg); | 75 | + AscendC::Reg::Adds(dst_reg0, src_reg0, ADDS_ADDEND, mask_reg); |
| 76 | // store data from RegTensor to UB. | 76 | // store data from RegTensor to UB. |
| 77 | - AscendC::MicroAPI::StoreAlign(output + i * one_repeat_size, dst_reg0, mask_reg); | 77 | + AscendC::Reg::StoreAlign(output + i * one_repeat_size, dst_reg0, mask_reg); |
| 78 | } | 78 | } |
| 79 | } | 79 | } |
| 80 | 80 | ||
| @@ -16,14 +16,14 @@ | |||
| 16 | template <typename T> | 16 | template <typename T> |
| 17 | __simd_vf__ inline void CopyVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize) | 17 | __simd_vf__ inline void CopyVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize) |
| 18 | { | 18 | { |
| 19 | - AscendC::MicroAPI::RegTensor<T> reg; | 19 | + AscendC::Reg::RegTensor<T> reg; |
| 20 | - AscendC::MicroAPI::MaskReg mask; | 20 | + AscendC::Reg::MaskReg mask; |
| 21 | - AscendC::MicroAPI::LoadAlign(mask, srcAddr); | 21 | + AscendC::Reg::LoadAlign(mask, srcAddr); |
| 22 | - AscendC::MicroAPI::Duplicate<T>(reg, 1, mask); | 22 | + AscendC::Reg::Duplicate<T>(reg, 1, mask); |
| 23 | - AscendC::MicroAPI::StoreAlign(dstAddr, reg, mask); | 23 | + AscendC::Reg::StoreAlign(dstAddr, reg, mask); |
| 24 | // save mask | 24 | // save mask |
| 25 | - mask = AscendC::MicroAPI::CreateMask<T, AscendC::MicroAPI::MaskPattern::ALL>(); | 25 | + mask = AscendC::Reg::CreateMask<T, AscendC::Reg::MaskPattern::ALL>(); |
| 26 | - AscendC::MicroAPI::StoreAlign(dstAddr + oneRepeatSize, mask); | 26 | + AscendC::Reg::StoreAlign(dstAddr + oneRepeatSize, mask); |
| 27 | } | 27 | } |
| 28 | 28 | ||
| 29 | template <typename T> | 29 | template <typename T> |
| @@ -19,14 +19,14 @@ template <typename T> | |||
| 19 | __simd_vf__ inline void CopyVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize, uint32_t repeatTimes, | 19 | __simd_vf__ inline void CopyVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize, uint32_t repeatTimes, |
| 20 | uint32_t count) | 20 | uint32_t count) |
| 21 | { | 21 | { |
| 22 | - AscendC::MicroAPI::RegTensor<T> reg; | 22 | + AscendC::Reg::RegTensor<T> reg; |
| 23 | - AscendC::MicroAPI::MaskReg mask; | 23 | + AscendC::Reg::MaskReg mask; |
| 24 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 24 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 25 | - mask = AscendC::MicroAPI::UpdateMask<T>(count); | 25 | + mask = AscendC::Reg::UpdateMask<T>(count); |
| 26 | // read | 26 | // read |
| 27 | - AscendC::MicroAPI::LoadAlign(reg, srcAddr + i * oneRepeatSize); | 27 | + AscendC::Reg::LoadAlign(reg, srcAddr + i * oneRepeatSize); |
| 28 | // write | 28 | // write |
| 29 | - AscendC::MicroAPI::StoreAlign(dstAddr + i * oneRepeatSize, reg, mask); | 29 | + AscendC::Reg::StoreAlign(dstAddr + i * oneRepeatSize, reg, mask); |
| 30 | } | 30 | } |
| 31 | } | 31 | } |
| 32 | 32 | ||
| @@ -34,14 +34,14 @@ template <typename T> | |||
| 34 | __simd_vf__ inline void CopyWithPostModeVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize, | 34 | __simd_vf__ inline void CopyWithPostModeVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize, |
| 35 | uint32_t repeatTimes, uint32_t count) | 35 | uint32_t repeatTimes, uint32_t count) |
| 36 | { | 36 | { |
| 37 | - AscendC::MicroAPI::RegTensor<T> reg; | 37 | + AscendC::Reg::RegTensor<T> reg; |
| 38 | - AscendC::MicroAPI::MaskReg mask; | 38 | + AscendC::Reg::MaskReg mask; |
| 39 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 39 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 40 | - mask = AscendC::MicroAPI::UpdateMask<T>(count); | 40 | + mask = AscendC::Reg::UpdateMask<T>(count); |
| 41 | // read | 41 | // read |
| 42 | - AscendC::MicroAPI::LoadAlign<T, AscendC::MicroAPI::PostLiteral::POST_MODE_UPDATE>(reg, srcAddr, oneRepeatSize); | 42 | + AscendC::Reg::LoadAlign<T, AscendC::Reg::PostLiteral::POST_MODE_UPDATE>(reg, srcAddr, oneRepeatSize); |
| 43 | // write | 43 | // write |
| 44 | - AscendC::MicroAPI::StoreAlign<T, AscendC::MicroAPI::PostLiteral::POST_MODE_UPDATE>(dstAddr, reg, oneRepeatSize, | 44 | + AscendC::Reg::StoreAlign<T, AscendC::Reg::PostLiteral::POST_MODE_UPDATE>(dstAddr, reg, oneRepeatSize, |
| 45 | mask); | 45 | mask); |
| 46 | } | 46 | } |
| 47 | } | 47 | } |
| @@ -18,16 +18,16 @@ | |||
| 18 | template <typename T> | 18 | template <typename T> |
| 19 | __simd_vf__ inline void CopyVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize) | 19 | __simd_vf__ inline void CopyVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize) |
| 20 | { | 20 | { |
| 21 | - AscendC::MicroAPI::RegTensor<T> srcReg, dstReg; | 21 | + AscendC::Reg::RegTensor<T> srcReg, dstReg; |
| 22 | - AscendC::MicroAPI::UnalignRegForLoad ureg0; | 22 | + AscendC::Reg::UnalignRegForLoad ureg0; |
| 23 | - AscendC::MicroAPI::UnalignRegForStore ureg1; | 23 | + AscendC::Reg::UnalignRegForStore ureg1; |
| 24 | - AscendC::MicroAPI::LoadUnAlignPre(ureg0, srcAddr); | 24 | + AscendC::Reg::LoadUnAlignPre(ureg0, srcAddr); |
| 25 | for (uint16_t i = 0; i < 2; ++i) { | 25 | for (uint16_t i = 0; i < 2; ++i) { |
| 26 | - AscendC::MicroAPI::LoadUnAlign(srcReg, ureg0, srcAddr + i * oneRepeatSize); | 26 | + AscendC::Reg::LoadUnAlign(srcReg, ureg0, srcAddr + i * oneRepeatSize); |
| 27 | // store data with post mode, so don't need to change dst operator's addrss | 27 | // store data with post mode, so don't need to change dst operator's addrss |
| 28 | - AscendC::MicroAPI::StoreUnAlign(dstAddr, srcReg, ureg1, oneRepeatSize); | 28 | + AscendC::Reg::StoreUnAlign(dstAddr, srcReg, ureg1, oneRepeatSize); |
| 29 | } | 29 | } |
| 30 | - AscendC::MicroAPI::StoreUnAlignPost(dstAddr, ureg1, 0); | 30 | + AscendC::Reg::StoreUnAlignPost(dstAddr, ureg1, 0); |
| 31 | } | 31 | } |
| 32 | 32 | ||
| 33 | template<typename T> | 33 | template<typename T> |
| @@ -19,15 +19,15 @@ template <typename T> | |||
| 19 | __simd_vf__ inline void CopyVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize, uint32_t repeatTimes, | 19 | __simd_vf__ inline void CopyVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize, uint32_t repeatTimes, |
| 20 | uint32_t count) | 20 | uint32_t count) |
| 21 | { | 21 | { |
| 22 | - AscendC::MicroAPI::RegTensor<T> reg; | 22 | + AscendC::Reg::RegTensor<T> reg; |
| 23 | - AscendC::MicroAPI::MaskReg mask; | 23 | + AscendC::Reg::MaskReg mask; |
| 24 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 24 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 25 | - mask = AscendC::MicroAPI::UpdateMask<T>(count); | 25 | + mask = AscendC::Reg::UpdateMask<T>(count); |
| 26 | // read | 26 | // read |
| 27 | - AscendC::MicroAPI::LoadAlign<T, AscendC::MicroAPI::DataCopyMode::DATA_BLOCK_COPY>( | 27 | + AscendC::Reg::LoadAlign<T, AscendC::Reg::DataCopyMode::DATA_BLOCK_COPY>( |
| 28 | reg, srcAddr + i * oneRepeatSize, 1, mask); | 28 | reg, srcAddr + i * oneRepeatSize, 1, mask); |
| 29 | // write | 29 | // write |
| 30 | - AscendC::MicroAPI::StoreAlign<T, AscendC::MicroAPI::DataCopyMode::DATA_BLOCK_COPY>(dstAddr + i * oneRepeatSize, | 30 | + AscendC::Reg::StoreAlign<T, AscendC::Reg::DataCopyMode::DATA_BLOCK_COPY>(dstAddr + i * oneRepeatSize, |
| 31 | reg, 1, mask); | 31 | reg, 1, mask); |
| 32 | } | 32 | } |
| 33 | } | 33 | } |
| @@ -36,16 +36,16 @@ template <typename T> | |||
| 36 | __simd_vf__ inline void CopyWithPostModeVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize, | 36 | __simd_vf__ inline void CopyWithPostModeVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize, |
| 37 | uint32_t repeatTimes, uint32_t count) | 37 | uint32_t repeatTimes, uint32_t count) |
| 38 | { | 38 | { |
| 39 | - AscendC::MicroAPI::RegTensor<T> reg; | 39 | + AscendC::Reg::RegTensor<T> reg; |
| 40 | - AscendC::MicroAPI::MaskReg mask; | 40 | + AscendC::Reg::MaskReg mask; |
| 41 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 41 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 42 | - mask = AscendC::MicroAPI::UpdateMask<T>(count); | 42 | + mask = AscendC::Reg::UpdateMask<T>(count); |
| 43 | // read | 43 | // read |
| 44 | - AscendC::MicroAPI::LoadAlign<T, AscendC::MicroAPI::DataCopyMode::DATA_BLOCK_COPY, | 44 | + AscendC::Reg::LoadAlign<T, AscendC::Reg::DataCopyMode::DATA_BLOCK_COPY, |
| 45 | - AscendC::MicroAPI::PostLiteral::POST_MODE_UPDATE>(reg, srcAddr, 1, 0, mask); | 45 | + AscendC::Reg::PostLiteral::POST_MODE_UPDATE>(reg, srcAddr, 1, 0, mask); |
| 46 | // write | 46 | // write |
| 47 | - AscendC::MicroAPI::StoreAlign<T, AscendC::MicroAPI::DataCopyMode::DATA_BLOCK_COPY, | 47 | + AscendC::Reg::StoreAlign<T, AscendC::Reg::DataCopyMode::DATA_BLOCK_COPY, |
| 48 | - AscendC::MicroAPI::PostLiteral::POST_MODE_UPDATE>(dstAddr, reg, 1, 0, mask); | 48 | + AscendC::Reg::PostLiteral::POST_MODE_UPDATE>(dstAddr, reg, 1, 0, mask); |
| 49 | } | 49 | } |
| 50 | } | 50 | } |
| 51 | 51 | ||
| @@ -17,14 +17,14 @@ template <typename T> | |||
| 17 | __simd_vf__ inline void AddVF(__ubuf__ T* dstAddr, __ubuf__ T* src0Addr, __ubuf__ T* src1Addr, uint16_t oneRepeatSize, | 17 | __simd_vf__ inline void AddVF(__ubuf__ T* dstAddr, __ubuf__ T* src0Addr, __ubuf__ T* src1Addr, uint16_t oneRepeatSize, |
| 18 | uint32_t repeatTimes, uint32_t count) | 18 | uint32_t repeatTimes, uint32_t count) |
| 19 | { | 19 | { |
| 20 | - AscendC::MicroAPI::RegTensor<T> reg0, reg1; | 20 | + AscendC::Reg::RegTensor<T> reg0, reg1; |
| 21 | - AscendC::MicroAPI::MaskReg mask; | 21 | + AscendC::Reg::MaskReg mask; |
| 22 | for (uint16_t i = 0; i < repeatTimes; i++) { | 22 | for (uint16_t i = 0; i < repeatTimes; i++) { |
| 23 | - mask = AscendC::MicroAPI::UpdateMask<T>(count); | 23 | + mask = AscendC::Reg::UpdateMask<T>(count); |
| 24 | - AscendC::MicroAPI::LoadAlign(reg0, src0Addr + i * oneRepeatSize); | 24 | + AscendC::Reg::LoadAlign(reg0, src0Addr + i * oneRepeatSize); |
| 25 | - AscendC::MicroAPI::LoadAlign(reg1, src1Addr + i * oneRepeatSize); | 25 | + AscendC::Reg::LoadAlign(reg1, src1Addr + i * oneRepeatSize); |
| 26 | - AscendC::MicroAPI::Add(reg0, reg0, reg1, mask); | 26 | + AscendC::Reg::Add(reg0, reg0, reg1, mask); |
| 27 | - AscendC::MicroAPI::StoreAlign(dstAddr + i * oneRepeatSize, reg0, mask); | 27 | + AscendC::Reg::StoreAlign(dstAddr + i * oneRepeatSize, reg0, mask); |
| 28 | } | 28 | } |
| 29 | } | 29 | } |
| 30 | 30 | ||
| @@ -17,17 +17,17 @@ template <typename T> | |||
| 17 | __simd_vf__ inline void AddVF(__ubuf__ T* dstAddr, uint32_t inner, uint32_t outter, uint32_t repeatTimes, | 17 | __simd_vf__ inline void AddVF(__ubuf__ T* dstAddr, uint32_t inner, uint32_t outter, uint32_t repeatTimes, |
| 18 | uint32_t oneRepeatSize) | 18 | uint32_t oneRepeatSize) |
| 19 | { | 19 | { |
| 20 | - AscendC::MicroAPI::MaskReg mask = AscendC::MicroAPI::CreateMask<T, AscendC::MicroAPI::MaskPattern::ALL>(); | 20 | + AscendC::Reg::MaskReg mask = AscendC::Reg::CreateMask<T, AscendC::Reg::MaskPattern::ALL>(); |
| 21 | - AscendC::MicroAPI::RegTensor<T> srcReg; | 21 | + AscendC::Reg::RegTensor<T> srcReg; |
| 22 | - AscendC::MicroAPI::RegTensor<T> dstReg; | 22 | + AscendC::Reg::RegTensor<T> dstReg; |
| 23 | for (uint16_t i = 0; i < outter - 1; ++i) { | 23 | for (uint16_t i = 0; i < outter - 1; ++i) { |
| 24 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 24 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 25 | - AscendC::MicroAPI::LoadAlign(srcReg, dstAddr + i * inner + j * oneRepeatSize); | 25 | + AscendC::Reg::LoadAlign(srcReg, dstAddr + i * inner + j * oneRepeatSize); |
| 26 | - AscendC::MicroAPI::LoadAlign(dstReg, dstAddr + (i + 1) * inner + j * oneRepeatSize); | 26 | + AscendC::Reg::LoadAlign(dstReg, dstAddr + (i + 1) * inner + j * oneRepeatSize); |
| 27 | - AscendC::MicroAPI::Add(dstReg, dstReg, srcReg, mask); | 27 | + AscendC::Reg::Add(dstReg, dstReg, srcReg, mask); |
| 28 | - AscendC::MicroAPI::StoreAlign(dstAddr + (i + 1) * inner + j * oneRepeatSize, dstReg, mask); | 28 | + AscendC::Reg::StoreAlign(dstAddr + (i + 1) * inner + j * oneRepeatSize, dstReg, mask); |
| 29 | - AscendC::MicroAPI::LocalMemBar<AscendC::MicroAPI::MemType::VEC_STORE, | 29 | + AscendC::Reg::LocalMemBar<AscendC::Reg::MemType::VEC_STORE, |
| 30 | - AscendC::MicroAPI::MemType::VEC_LOAD>(); | 30 | + AscendC::Reg::MemType::VEC_LOAD>(); |
| 31 | } | 31 | } |
| 32 | } | 32 | } |
| 33 | } | 33 | } |
| @@ -16,22 +16,22 @@ | |||
| 16 | template <typename T> | 16 | template <typename T> |
| 17 | __simd_vf__ inline void ComputeVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize) | 17 | __simd_vf__ inline void ComputeVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize) |
| 18 | { | 18 | { |
| 19 | - AscendC::MicroAPI::RegTensor<T> reg0, reg1; | 19 | + AscendC::Reg::RegTensor<T> reg0, reg1; |
| 20 | - AscendC::MicroAPI::MaskReg mask = AscendC::MicroAPI::CreateMask<T, AscendC::MicroAPI::MaskPattern::ALL>(); | 20 | + AscendC::Reg::MaskReg mask = AscendC::Reg::CreateMask<T, AscendC::Reg::MaskPattern::ALL>(); |
| 21 | // read 1 | 21 | // read 1 |
| 22 | - AscendC::MicroAPI::LoadAlign(reg0, srcAddr + oneRepeatSize); | 22 | + AscendC::Reg::LoadAlign(reg0, srcAddr + oneRepeatSize); |
| 23 | // read 2 | 23 | // read 2 |
| 24 | - AscendC::MicroAPI::LoadAlign(reg1, srcAddr + oneRepeatSize); | 24 | + AscendC::Reg::LoadAlign(reg1, srcAddr + oneRepeatSize); |
| 25 | // duplicate 1 | 25 | // duplicate 1 |
| 26 | - AscendC::MicroAPI::Duplicate<T>(reg0, 1.0f, mask); | 26 | + AscendC::Reg::Duplicate<T>(reg0, 1.0f, mask); |
| 27 | // duplicate 2 | 27 | // duplicate 2 |
| 28 | - AscendC::MicroAPI::Duplicate<T>(reg1, 2.0f, mask); | 28 | + AscendC::Reg::Duplicate<T>(reg1, 2.0f, mask); |
| 29 | // write 1 | 29 | // write 1 |
| 30 | - AscendC::MicroAPI::StoreAlign(dstAddr, reg0, mask); | 30 | + AscendC::Reg::StoreAlign(dstAddr, reg0, mask); |
| 31 | // sync | 31 | // sync |
| 32 | - AscendC::MicroAPI::LocalMemBar<AscendC::MicroAPI::MemType::VEC_STORE, AscendC::MicroAPI::MemType::VEC_STORE>(); | 32 | + AscendC::Reg::LocalMemBar<AscendC::Reg::MemType::VEC_STORE, AscendC::Reg::MemType::VEC_STORE>(); |
| 33 | // write 2 | 33 | // write 2 |
| 34 | - AscendC::MicroAPI::StoreAlign(dstAddr + oneRepeatSize, reg1, mask); | 34 | + AscendC::Reg::StoreAlign(dstAddr + oneRepeatSize, reg1, mask); |
| 35 | } | 35 | } |
| 36 | 36 | ||
| 37 | template <typename T> | 37 | template <typename T> |
| @@ -19,15 +19,15 @@ template <typename T> | |||
| 19 | __simd_vf__ inline void LoadVF(__ubuf__ T* srcAddr, __ubuf__ T* dstAddr, uint16_t postUpdateStride, uint32_t repeatTimes, | 19 | __simd_vf__ inline void LoadVF(__ubuf__ T* srcAddr, __ubuf__ T* dstAddr, uint16_t postUpdateStride, uint32_t repeatTimes, |
| 20 | uint16_t count) | 20 | uint16_t count) |
| 21 | { | 21 | { |
| 22 | - AscendC::MicroAPI::RegTensor<T> srcReg, dstReg; | 22 | + AscendC::Reg::RegTensor<T> srcReg, dstReg; |
| 23 | - AscendC::MicroAPI::UnalignRegForLoad ureg0; | 23 | + AscendC::Reg::UnalignRegForLoad ureg0; |
| 24 | - AscendC::MicroAPI::UnalignRegForStore ureg1; | 24 | + AscendC::Reg::UnalignRegForStore ureg1; |
| 25 | - AscendC::MicroAPI::LoadUnAlignPre(ureg0, srcAddr); | 25 | + AscendC::Reg::LoadUnAlignPre(ureg0, srcAddr); |
| 26 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 26 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 27 | - AscendC::MicroAPI::LoadUnAlign(srcReg, ureg0, srcAddr + i * postUpdateStride); | 27 | + AscendC::Reg::LoadUnAlign(srcReg, ureg0, srcAddr + i * postUpdateStride); |
| 28 | - AscendC::MicroAPI::StoreUnAlign(dstAddr, srcReg, ureg1, postUpdateStride); | 28 | + AscendC::Reg::StoreUnAlign(dstAddr, srcReg, ureg1, postUpdateStride); |
| 29 | } | 29 | } |
| 30 | - AscendC::MicroAPI::StoreUnAlignPost(dstAddr, ureg1, 0); | 30 | + AscendC::Reg::StoreUnAlignPost(dstAddr, ureg1, 0); |
| 31 | } | 31 | } |
| 32 | 32 | ||
| 33 | template <typename T> | 33 | template <typename T> |
Mexamples/04_best_practices/12_high_performance_vf/optimize_vf_dual_instr/optimize_vf_dual_instr.asc+19-19
| @@ -18,28 +18,28 @@ template <typename T> | |||
| 18 | __simd_vf__ inline void AddVF(__ubuf__ T* src0Addr, __ubuf__ T* src1Addr, __ubuf__ T* dstAddr, uint16_t oneRepeatSize, | 18 | __simd_vf__ inline void AddVF(__ubuf__ T* src0Addr, __ubuf__ T* src1Addr, __ubuf__ T* dstAddr, uint16_t oneRepeatSize, |
| 19 | uint16_t repeatTimes, uint32_t count) | 19 | uint16_t repeatTimes, uint32_t count) |
| 20 | { | 20 | { |
| 21 | - AscendC::MicroAPI::RegTensor<T> srcReg0; | 21 | + AscendC::Reg::RegTensor<T> srcReg0; |
| 22 | - AscendC::MicroAPI::RegTensor<T> srcReg1; | 22 | + AscendC::Reg::RegTensor<T> srcReg1; |
| 23 | - AscendC::MicroAPI::RegTensor<T> dstReg0; | 23 | + AscendC::Reg::RegTensor<T> dstReg0; |
| 24 | - AscendC::MicroAPI::RegTensor<T> dstReg1; | 24 | + AscendC::Reg::RegTensor<T> dstReg1; |
| 25 | - AscendC::MicroAPI::MaskReg mask; | 25 | + AscendC::Reg::MaskReg mask; |
| 26 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 26 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 27 | - mask = AscendC::MicroAPI::UpdateMask<T>(count); | 27 | + mask = AscendC::Reg::UpdateMask<T>(count); |
| 28 | - AscendC::MicroAPI::LoadAlign(srcReg0, src0Addr + i * oneRepeatSize); | 28 | + AscendC::Reg::LoadAlign(srcReg0, src0Addr + i * oneRepeatSize); |
| 29 | - AscendC::MicroAPI::LoadAlign(srcReg1, src1Addr + i * oneRepeatSize); | 29 | + AscendC::Reg::LoadAlign(srcReg1, src1Addr + i * oneRepeatSize); |
| 30 | - AscendC::MicroAPI::Add(dstReg0, srcReg0, srcReg1, mask); | 30 | + AscendC::Reg::Add(dstReg0, srcReg0, srcReg1, mask); |
| 31 | - AscendC::MicroAPI::Adds(dstReg0, dstReg0, 10, mask); | 31 | + AscendC::Reg::Adds(dstReg0, dstReg0, 10, mask); |
| 32 | - AscendC::MicroAPI::Adds(dstReg0, dstReg0, 10, mask); | 32 | + AscendC::Reg::Adds(dstReg0, dstReg0, 10, mask); |
| 33 | - AscendC::MicroAPI::Adds(dstReg0, dstReg0, 10, mask); | 33 | + AscendC::Reg::Adds(dstReg0, dstReg0, 10, mask); |
| 34 | - AscendC::MicroAPI::StoreAlign(dstAddr + i * oneRepeatSize, dstReg0, mask); | 34 | + AscendC::Reg::StoreAlign(dstAddr + i * oneRepeatSize, dstReg0, mask); |
| 35 | } | 35 | } |
| 36 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 36 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 37 | - mask = AscendC::MicroAPI::UpdateMask<T>(count); | 37 | + mask = AscendC::Reg::UpdateMask<T>(count); |
| 38 | - AscendC::MicroAPI::LoadAlign(dstReg0, dstAddr + i * oneRepeatSize); | 38 | + AscendC::Reg::LoadAlign(dstReg0, dstAddr + i * oneRepeatSize); |
| 39 | - AscendC::MicroAPI::Adds(dstReg0, dstReg0, 10, mask); | 39 | + AscendC::Reg::Adds(dstReg0, dstReg0, 10, mask); |
| 40 | - AscendC::MicroAPI::Adds(dstReg0, dstReg0, 10, mask); | 40 | + AscendC::Reg::Adds(dstReg0, dstReg0, 10, mask); |
| 41 | - AscendC::MicroAPI::Adds(dstReg0, dstReg0, 10, mask); | 41 | + AscendC::Reg::Adds(dstReg0, dstReg0, 10, mask); |
| 42 | - AscendC::MicroAPI::StoreAlign(dstAddr + i * oneRepeatSize, dstReg0, mask); | 42 | + AscendC::Reg::StoreAlign(dstAddr + i * oneRepeatSize, dstReg0, mask); |
| 43 | } | 43 | } |
| 44 | } | 44 | } |
| 45 | 45 | ||
| @@ -17,15 +17,15 @@ template <typename T> | |||
| 17 | __simd_vf__ inline void DivVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize, uint32_t repeatTimes, | 17 | __simd_vf__ inline void DivVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize, uint32_t repeatTimes, |
| 18 | uint32_t count) | 18 | uint32_t count) |
| 19 | { | 19 | { |
| 20 | - AscendC::MicroAPI::MaskReg mask; | 20 | + AscendC::Reg::MaskReg mask; |
| 21 | - AscendC::MicroAPI::RegTensor<T> srcReg0, srcReg1, dstReg; | 21 | + AscendC::Reg::RegTensor<T> srcReg0, srcReg1, dstReg; |
| 22 | constexpr float num = 1.0f; | 22 | constexpr float num = 1.0f; |
| 23 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 23 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 24 | - mask = AscendC::MicroAPI::UpdateMask<T>(count); | 24 | + mask = AscendC::Reg::UpdateMask<T>(count); |
| 25 | - AscendC::MicroAPI::LoadAlign(srcReg0, srcAddr + i * oneRepeatSize); | 25 | + AscendC::Reg::LoadAlign(srcReg0, srcAddr + i * oneRepeatSize); |
| 26 | - AscendC::MicroAPI::Duplicate(srcReg1, num, mask); | 26 | + AscendC::Reg::Duplicate(srcReg1, num, mask); |
| 27 | - AscendC::MicroAPI::Div(dstReg, srcReg1, srcReg0, mask); | 27 | + AscendC::Reg::Div(dstReg, srcReg1, srcReg0, mask); |
| 28 | - AscendC::MicroAPI::StoreAlign(dstAddr + i * oneRepeatSize, dstReg, mask); | 28 | + AscendC::Reg::StoreAlign(dstAddr + i * oneRepeatSize, dstReg, mask); |
| 29 | } | 29 | } |
| 30 | } | 30 | } |
| 31 | 31 | ||
| @@ -33,14 +33,14 @@ template <typename T> | |||
| 33 | __simd_vf__ inline void AddVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize, uint32_t repeatTimes, | 33 | __simd_vf__ inline void AddVF(__ubuf__ T* dstAddr, __ubuf__ T* srcAddr, uint16_t oneRepeatSize, uint32_t repeatTimes, |
| 34 | uint32_t count) | 34 | uint32_t count) |
| 35 | { | 35 | { |
| 36 | - AscendC::MicroAPI::MaskReg mask; | 36 | + AscendC::Reg::MaskReg mask; |
| 37 | - AscendC::MicroAPI::RegTensor<T> srcReg, dstReg; | 37 | + AscendC::Reg::RegTensor<T> srcReg, dstReg; |
| 38 | constexpr float num = 1.0f; | 38 | constexpr float num = 1.0f; |
| 39 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 39 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 40 | - mask = AscendC::MicroAPI::UpdateMask<T>(count); | 40 | + mask = AscendC::Reg::UpdateMask<T>(count); |
| 41 | - AscendC::MicroAPI::LoadAlign(srcReg, srcAddr + i * oneRepeatSize); | 41 | + AscendC::Reg::LoadAlign(srcReg, srcAddr + i * oneRepeatSize); |
| 42 | - AscendC::MicroAPI::Adds(dstReg, srcReg, num, mask); | 42 | + AscendC::Reg::Adds(dstReg, srcReg, num, mask); |
| 43 | - AscendC::MicroAPI::StoreAlign(dstAddr + i * oneRepeatSize, dstReg, mask); | 43 | + AscendC::Reg::StoreAlign(dstAddr + i * oneRepeatSize, dstReg, mask); |
| 44 | } | 44 | } |
| 45 | } | 45 | } |
| 46 | 46 | ||
| @@ -17,17 +17,17 @@ template <typename T> | |||
| 17 | __simd_vf__ inline void AddVF(__ubuf__ T* srcAddr, __ubuf__ T* dstAddr, uint16_t stride1, uint16_t stride2, | 17 | __simd_vf__ inline void AddVF(__ubuf__ T* srcAddr, __ubuf__ T* dstAddr, uint16_t stride1, uint16_t stride2, |
| 18 | uint16_t stride3, uint16_t stride4, uint16_t loop) | 18 | uint16_t stride3, uint16_t stride4, uint16_t loop) |
| 19 | { | 19 | { |
| 20 | - AscendC::MicroAPI::RegTensor<T> srcReg; | 20 | + AscendC::Reg::RegTensor<T> srcReg; |
| 21 | - AscendC::MicroAPI::RegTensor<T> dstReg; | 21 | + AscendC::Reg::RegTensor<T> dstReg; |
| 22 | - AscendC::MicroAPI::MaskReg mask = AscendC::MicroAPI::CreateMask<T, AscendC::MicroAPI::MaskPattern::VL8>(); | 22 | + AscendC::Reg::MaskReg mask = AscendC::Reg::CreateMask<T, AscendC::Reg::MaskPattern::VL8>(); |
| 23 | for (uint16_t i = 0; i < loop; i++) { | 23 | for (uint16_t i = 0; i < loop; i++) { |
| 24 | for (uint16_t j = 0; j < loop; j++) { | 24 | for (uint16_t j = 0; j < loop; j++) { |
| 25 | for (uint16_t k = 0; k < loop; k++) { | 25 | for (uint16_t k = 0; k < loop; k++) { |
| 26 | for (uint16_t m = 0; m < loop; m++) { | 26 | for (uint16_t m = 0; m < loop; m++) { |
| 27 | - AscendC::MicroAPI::LoadAlign(srcReg, | 27 | + AscendC::Reg::LoadAlign(srcReg, |
| 28 | srcAddr + i * stride1 + j * stride2 + k * stride3 + m * stride4); | 28 | srcAddr + i * stride1 + j * stride2 + k * stride3 + m * stride4); |
| 29 | - AscendC::MicroAPI::Adds(dstReg, srcReg, 1.0f, mask); | 29 | + AscendC::Reg::Adds(dstReg, srcReg, 1.0f, mask); |
| 30 | - AscendC::MicroAPI::StoreAlign(dstAddr + i * stride1 + j * stride2 + k * stride3 + m * stride4, | 30 | + AscendC::Reg::StoreAlign(dstAddr + i * stride1 + j * stride2 + k * stride3 + m * stride4, |
| 31 | dstReg, mask); | 31 | dstReg, mask); |
| 32 | } | 32 | } |
| 33 | } | 33 | } |
| @@ -72,19 +72,19 @@ | |||
| 72 | floor_mod_simd负责float场景的计算。 | 72 | floor_mod_simd负责float场景的计算。 |
| 73 | ```cpp | 73 | ```cpp |
| 74 | for (uint16_t j = 0; j < loopTimes; j++) { | 74 | for (uint16_t j = 0; j < loopTimes; j++) { |
| 75 | - preg = AscendC::MicroAPI::UpdateMask<T>(sregMask); | 75 | + preg = AscendC::Reg::UpdateMask<T>(sregMask); |
| 76 | - AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_NORM>(fmodResValue, fmodResAddr + VL_T * j); | 76 | + AscendC::Reg::DataCopy<T, AscendC::Reg::LoadDist::DIST_NORM>(fmodResValue, fmodResAddr + VL_T * j); |
| 77 | - AscendC::MicroAPI::Compare<T, AscendC::CMPMODE::NE>(negValue, fmodResValue, zeroValue, preg); | 77 | + AscendC::Reg::Compare<T, AscendC::CMPMODE::NE>(negValue, fmodResValue, zeroValue, preg); |
| 78 | 78 | ||
| 79 | - AscendC::MicroAPI::And(fmodSignValue, (AscendC::MicroAPI::RegTensor<uint32_t>&)fmodResValue, signValue, preg); | 79 | + AscendC::Reg::And(fmodSignValue, (AscendC::Reg::RegTensor<uint32_t>&)fmodResValue, signValue, preg); |
| 80 | - AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_NORM>(inputX2Value, otherAddr + VL_T * j); | 80 | + AscendC::Reg::DataCopy<T, AscendC::Reg::LoadDist::DIST_NORM>(inputX2Value, otherAddr + VL_T * j); |
| 81 | - AscendC::MicroAPI::Add(addValue, fmodResValue, inputX2Value, preg); | 81 | + AscendC::Reg::Add(addValue, fmodResValue, inputX2Value, preg); |
| 82 | - AscendC::MicroAPI::And(inputX2signValue, (AscendC::MicroAPI::RegTensor<uint32_t>&)inputX2Value, signValue, preg); | 82 | + AscendC::Reg::And(inputX2signValue, (AscendC::Reg::RegTensor<uint32_t>&)inputX2Value, signValue, preg); |
| 83 | - AscendC::MicroAPI::Compare<uint32_t, AscendC::CMPMODE::NE>(signNegValue, fmodSignValue, inputX2signValue, preg); | 83 | + AscendC::Reg::Compare<uint32_t, AscendC::CMPMODE::NE>(signNegValue, fmodSignValue, inputX2signValue, preg); |
| 84 | 84 | ||
| 85 | - AscendC::MicroAPI::MaskAnd(resMaskValue, signNegValue, negValue, preg); | 85 | + AscendC::Reg::MaskAnd(resMaskValue, signNegValue, negValue, preg); |
| 86 | - AscendC::MicroAPI::Select(resValue, addValue, fmodResValue, resMaskValue); | 86 | + AscendC::Reg::Select(resValue, addValue, fmodResValue, resMaskValue); |
| 87 | - AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::StoreDist::DIST_NORM>(dstAddr + VL_T * j, resValue, preg); | 87 | + AscendC::Reg::DataCopy<T, AscendC::Reg::StoreDist::DIST_NORM>(dstAddr + VL_T * j, resValue, preg); |
| 88 | } | 88 | } |
| 89 | ``` | 89 | ``` |
| 90 | 90 | ||
| @@ -37,39 +37,39 @@ __simd_vf__ inline void floor_mod_float_simd(__ubuf__ T* fmodResAddr, __ubuf__ T | |||
| 37 | uint32_t vecLen = VECTOR_LENGTH / sizeof(T); | 37 | uint32_t vecLen = VECTOR_LENGTH / sizeof(T); |
| 38 | uint16_t loopTimes = (count + vecLen - 1) / vecLen; | 38 | uint16_t loopTimes = (count + vecLen - 1) / vecLen; |
| 39 | 39 | ||
| 40 | - AscendC::MicroAPI::RegTensor<T> zeroValue; | 40 | + AscendC::Reg::RegTensor<T> zeroValue; |
| 41 | - AscendC::MicroAPI::RegTensor<T> fmodResValue; | 41 | + AscendC::Reg::RegTensor<T> fmodResValue; |
| 42 | - AscendC::MicroAPI::RegTensor<T> inputX2Value; | 42 | + AscendC::Reg::RegTensor<T> inputX2Value; |
| 43 | - AscendC::MicroAPI::RegTensor<T> addValue; | 43 | + AscendC::Reg::RegTensor<T> addValue; |
| 44 | - AscendC::MicroAPI::RegTensor<T> resValue; | 44 | + AscendC::Reg::RegTensor<T> resValue; |
| 45 | 45 | ||
| 46 | - AscendC::MicroAPI::RegTensor<uint32_t> signValue; | 46 | + AscendC::Reg::RegTensor<uint32_t> signValue; |
| 47 | - AscendC::MicroAPI::RegTensor<uint32_t> fmodSignValue; | 47 | + AscendC::Reg::RegTensor<uint32_t> fmodSignValue; |
| 48 | - AscendC::MicroAPI::RegTensor<uint32_t> inputX2signValue; | 48 | + AscendC::Reg::RegTensor<uint32_t> inputX2signValue; |
| 49 | 49 | ||
| 50 | - AscendC::MicroAPI::MaskReg preg; | 50 | + AscendC::Reg::MaskReg preg; |
| 51 | - AscendC::MicroAPI::MaskReg negValue; | 51 | + AscendC::Reg::MaskReg negValue; |
| 52 | - AscendC::MicroAPI::MaskReg signNegValue; | 52 | + AscendC::Reg::MaskReg signNegValue; |
| 53 | - AscendC::MicroAPI::MaskReg resMaskValue; | 53 | + AscendC::Reg::MaskReg resMaskValue; |
| 54 | uint32_t sregMask = count; | 54 | uint32_t sregMask = count; |
| 55 | 55 | ||
| 56 | - AscendC::MicroAPI::Duplicate(zeroValue, T(0)); | 56 | + AscendC::Reg::Duplicate(zeroValue, T(0)); |
| 57 | - AscendC::MicroAPI::Duplicate(signValue, FMOD_B32_SIGN); | 57 | + AscendC::Reg::Duplicate(signValue, FMOD_B32_SIGN); |
| 58 | 58 | ||
| 59 | for (uint16_t j = 0; j < loopTimes; j++) { | 59 | for (uint16_t j = 0; j < loopTimes; j++) { |
| 60 | - preg = AscendC::MicroAPI::UpdateMask<T>(sregMask); | 60 | + preg = AscendC::Reg::UpdateMask<T>(sregMask); |
| 61 | - AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_NORM>(fmodResValue, fmodResAddr + vecLen * j); | 61 | + AscendC::Reg::DataCopy<T, AscendC::Reg::LoadDist::DIST_NORM>(fmodResValue, fmodResAddr + vecLen * j); |
| 62 | - AscendC::MicroAPI::Compare<T, AscendC::CMPMODE::NE>(negValue, fmodResValue, zeroValue, preg); | 62 | + AscendC::Reg::Compare<T, AscendC::CMPMODE::NE>(negValue, fmodResValue, zeroValue, preg); |
| 63 | 63 | ||
| 64 | - AscendC::MicroAPI::And(fmodSignValue, (AscendC::MicroAPI::RegTensor<uint32_t>&)fmodResValue, signValue, preg); | 64 | + AscendC::Reg::And(fmodSignValue, (AscendC::Reg::RegTensor<uint32_t>&)fmodResValue, signValue, preg); |
| 65 | - AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_NORM>(inputX2Value, otherAddr + vecLen * j); | 65 | + AscendC::Reg::DataCopy<T, AscendC::Reg::LoadDist::DIST_NORM>(inputX2Value, otherAddr + vecLen * j); |
| 66 | - AscendC::MicroAPI::Add(addValue, fmodResValue, inputX2Value, preg); | 66 | + AscendC::Reg::Add(addValue, fmodResValue, inputX2Value, preg); |
| 67 | - AscendC::MicroAPI::And(inputX2signValue, (AscendC::MicroAPI::RegTensor<uint32_t>&)inputX2Value, signValue, preg); | 67 | + AscendC::Reg::And(inputX2signValue, (AscendC::Reg::RegTensor<uint32_t>&)inputX2Value, signValue, preg); |
| 68 | - AscendC::MicroAPI::Compare<uint32_t, AscendC::CMPMODE::NE>(signNegValue, fmodSignValue, inputX2signValue, preg); | 68 | + AscendC::Reg::Compare<uint32_t, AscendC::CMPMODE::NE>(signNegValue, fmodSignValue, inputX2signValue, preg); |
| 69 | 69 | ||
| 70 | - AscendC::MicroAPI::MaskAnd(resMaskValue, signNegValue, negValue, preg); | 70 | + AscendC::Reg::MaskAnd(resMaskValue, signNegValue, negValue, preg); |
| 71 | - AscendC::MicroAPI::Select(resValue, addValue, fmodResValue, resMaskValue); | 71 | + AscendC::Reg::Select(resValue, addValue, fmodResValue, resMaskValue); |
| 72 | - AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::StoreDist::DIST_NORM>(dstAddr + vecLen * j, resValue, preg); | 72 | + AscendC::Reg::DataCopy<T, AscendC::Reg::StoreDist::DIST_NORM>(dstAddr + vecLen * j, resValue, preg); |
| 73 | } | 73 | } |
| 74 | } | 74 | } |
| 75 | 75 | ||
| @@ -80,51 +80,51 @@ __simd_vf__ inline void floor_mod_int_simd(__ubuf__ T* dstAddr, __ubuf__ T* inpu | |||
| 80 | uint32_t vecLen = VECTOR_LENGTH / sizeof(T); | 80 | uint32_t vecLen = VECTOR_LENGTH / sizeof(T); |
| 81 | uint16_t loopTimes = (count + vecLen - 1) / vecLen; | 81 | uint16_t loopTimes = (count + vecLen - 1) / vecLen; |
| 82 | 82 | ||
| 83 | - AscendC::MicroAPI::RegTensor<T> zeroValue; | 83 | + AscendC::Reg::RegTensor<T> zeroValue; |
| 84 | - AscendC::MicroAPI::RegTensor<T> defaultValue; | 84 | + AscendC::Reg::RegTensor<T> defaultValue; |
| 85 | - AscendC::MicroAPI::RegTensor<T> signValue; | 85 | + AscendC::Reg::RegTensor<T> signValue; |
| 86 | - AscendC::MicroAPI::RegTensor<T> input1Value; | 86 | + AscendC::Reg::RegTensor<T> input1Value; |
| 87 | - AscendC::MicroAPI::RegTensor<T> input2Value; | 87 | + AscendC::Reg::RegTensor<T> input2Value; |
| 88 | - AscendC::MicroAPI::RegTensor<T> divValue; | 88 | + AscendC::Reg::RegTensor<T> divValue; |
| 89 | - AscendC::MicroAPI::RegTensor<T> mulValue; | 89 | + AscendC::Reg::RegTensor<T> mulValue; |
| 90 | - AscendC::MicroAPI::RegTensor<T> subValue; | 90 | + AscendC::Reg::RegTensor<T> subValue; |
| 91 | - AscendC::MicroAPI::RegTensor<T> modValue; | 91 | + AscendC::Reg::RegTensor<T> modValue; |
| 92 | - AscendC::MicroAPI::RegTensor<T> modSignValue; | 92 | + AscendC::Reg::RegTensor<T> modSignValue; |
| 93 | - AscendC::MicroAPI::RegTensor<T> addValue; | 93 | + AscendC::Reg::RegTensor<T> addValue; |
| 94 | - AscendC::MicroAPI::RegTensor<T> input2SignValue; | 94 | + AscendC::Reg::RegTensor<T> input2SignValue; |
| 95 | - AscendC::MicroAPI::RegTensor<T> resValue; | 95 | + AscendC::Reg::RegTensor<T> resValue; |
| 96 | 96 | ||
| 97 | - AscendC::MicroAPI::MaskReg preg; | 97 | + AscendC::Reg::MaskReg preg; |
| 98 | - AscendC::MicroAPI::MaskReg cmpValue; | 98 | + AscendC::Reg::MaskReg cmpValue; |
| 99 | - AscendC::MicroAPI::MaskReg negValue; | 99 | + AscendC::Reg::MaskReg negValue; |
| 100 | - AscendC::MicroAPI::MaskReg signNegValue; | 100 | + AscendC::Reg::MaskReg signNegValue; |
| 101 | - AscendC::MicroAPI::MaskReg resMaskValue; | 101 | + AscendC::Reg::MaskReg resMaskValue; |
| 102 | uint32_t sregMask = count; | 102 | uint32_t sregMask = count; |
| 103 | 103 | ||
| 104 | - AscendC::MicroAPI::Duplicate(zeroValue, T(0)); | 104 | + AscendC::Reg::Duplicate(zeroValue, T(0)); |
| 105 | - AscendC::MicroAPI::Duplicate(defaultValue, T(-1)); | 105 | + AscendC::Reg::Duplicate(defaultValue, T(-1)); |
| 106 | - AscendC::MicroAPI::Duplicate(signValue, FMOD_B32_SIGN); | 106 | + AscendC::Reg::Duplicate(signValue, FMOD_B32_SIGN); |
| 107 | 107 | ||
| 108 | for (uint16_t j = 0; j < loopTimes; j++) { | 108 | for (uint16_t j = 0; j < loopTimes; j++) { |
| 109 | // handel -1 | 109 | // handel -1 |
| 110 | - preg = AscendC::MicroAPI::UpdateMask<T>(sregMask); | 110 | + preg = AscendC::Reg::UpdateMask<T>(sregMask); |
| 111 | - AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_NORM>(input2Value, input2Addr + vecLen * j); | 111 | + AscendC::Reg::DataCopy<T, AscendC::Reg::LoadDist::DIST_NORM>(input2Value, input2Addr + vecLen * j); |
| 112 | - AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_NORM>(divValue, divAddr + vecLen * j); | 112 | + AscendC::Reg::DataCopy<T, AscendC::Reg::LoadDist::DIST_NORM>(divValue, divAddr + vecLen * j); |
| 113 | - AscendC::MicroAPI::Mul(mulValue, input2Value, divValue, preg); | 113 | + AscendC::Reg::Mul(mulValue, input2Value, divValue, preg); |
| 114 | - AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_NORM>(input1Value, input1Addr + vecLen * j); | 114 | + AscendC::Reg::DataCopy<T, AscendC::Reg::LoadDist::DIST_NORM>(input1Value, input1Addr + vecLen * j); |
| 115 | - AscendC::MicroAPI::Sub(subValue, input1Value, mulValue, preg); | 115 | + AscendC::Reg::Sub(subValue, input1Value, mulValue, preg); |
| 116 | - AscendC::MicroAPI::Compare<T, AscendC::CMPMODE::NE>(cmpValue, input2Value, zeroValue, preg); | 116 | + AscendC::Reg::Compare<T, AscendC::CMPMODE::NE>(cmpValue, input2Value, zeroValue, preg); |
| 117 | - AscendC::MicroAPI::Select(modValue, subValue, defaultValue, cmpValue); | 117 | + AscendC::Reg::Select(modValue, subValue, defaultValue, cmpValue); |
| 118 | 118 | ||
| 119 | // post handel | 119 | // post handel |
| 120 | - AscendC::MicroAPI::Add(addValue, modValue, input2Value, preg); | 120 | + AscendC::Reg::Add(addValue, modValue, input2Value, preg); |
| 121 | - AscendC::MicroAPI::Compare<T, AscendC::CMPMODE::NE>(negValue, modValue, zeroValue, preg); | 121 | + AscendC::Reg::Compare<T, AscendC::CMPMODE::NE>(negValue, modValue, zeroValue, preg); |
| 122 | - AscendC::MicroAPI::And(input2SignValue, input2Value, signValue, preg); | 122 | + AscendC::Reg::And(input2SignValue, input2Value, signValue, preg); |
| 123 | - AscendC::MicroAPI::And(modSignValue, modValue, signValue, preg); | 123 | + AscendC::Reg::And(modSignValue, modValue, signValue, preg); |
| 124 | - AscendC::MicroAPI::Compare<T, AscendC::CMPMODE::NE>(signNegValue, modSignValue, input2SignValue, preg); | 124 | + AscendC::Reg::Compare<T, AscendC::CMPMODE::NE>(signNegValue, modSignValue, input2SignValue, preg); |
| 125 | - AscendC::MicroAPI::MaskAnd(resMaskValue, signNegValue, negValue, preg); | 125 | + AscendC::Reg::MaskAnd(resMaskValue, signNegValue, negValue, preg); |
| 126 | - AscendC::MicroAPI::Select(resValue, addValue, modValue, resMaskValue); | 126 | + AscendC::Reg::Select(resValue, addValue, modValue, resMaskValue); |
| 127 | - AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::StoreDist::DIST_NORM>(dstAddr + vecLen * j, resValue, preg); | 127 | + AscendC::Reg::DataCopy<T, AscendC::Reg::StoreDist::DIST_NORM>(dstAddr + vecLen * j, resValue, preg); |
| 128 | } | 128 | } |
| 129 | } | 129 | } |
| 130 | 130 | ||
| @@ -35,40 +35,40 @@ template <typename T> | |||
| 35 | __simd_vf__ inline void GeGLUImplVF( | 35 | __simd_vf__ inline void GeGLUImplVF( |
| 36 | __ubuf__ T* dst, __ubuf__ T* src0, __ubuf__ T* src1, uint32_t count, const uint16_t repeatTimes) | 36 | __ubuf__ T* dst, __ubuf__ T* src0, __ubuf__ T* src1, uint32_t count, const uint16_t repeatTimes) |
| 37 | { | 37 | { |
| 38 | - MicroAPI::RegTensor<half> srcOrigin0; | 38 | + Reg::RegTensor<half> srcOrigin0; |
| 39 | - MicroAPI::RegTensor<half> srcOrigin1; | 39 | + Reg::RegTensor<half> srcOrigin1; |
| 40 | - MicroAPI::RegTensor<float> srcVreg0; | 40 | + Reg::RegTensor<float> srcVreg0; |
| 41 | - MicroAPI::RegTensor<float> srcVreg1; | 41 | + Reg::RegTensor<float> srcVreg1; |
| 42 | - MicroAPI::RegTensor<float> tmpReg0; | 42 | + Reg::RegTensor<float> tmpReg0; |
| 43 | - MicroAPI::RegTensor<float> tmpReg1; | 43 | + Reg::RegTensor<float> tmpReg1; |
| 44 | - MicroAPI::RegTensor<float> dstVreg; | 44 | + Reg::RegTensor<float> dstVreg; |
| 45 | - MicroAPI::MaskReg mask; | 45 | + Reg::MaskReg mask; |
| 46 | constexpr uint32_t oneRepElm = static_cast<uint32_t>(GetVecLen() / sizeof(float)); | 46 | constexpr uint32_t oneRepElm = static_cast<uint32_t>(GetVecLen() / sizeof(float)); |
| 47 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 47 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 48 | - mask = MicroAPI::UpdateMask<float>(count); | 48 | + mask = Reg::UpdateMask<float>(count); |
| 49 | if constexpr (sizeof(T) == sizeof(half)) { | 49 | if constexpr (sizeof(T) == sizeof(half)) { |
| 50 | - MicroAPI::LoadAlign<half, MicroAPI::LoadDist::DIST_UNPACK_B16>(srcOrigin0, src0 + i * oneRepElm); | 50 | + Reg::LoadAlign<half, Reg::LoadDist::DIST_UNPACK_B16>(srcOrigin0, src0 + i * oneRepElm); |
| 51 | - MicroAPI::Cast<float, half, castTraitB16ToB32>(srcVreg0, srcOrigin0, mask); | 51 | + Reg::Cast<float, half, castTraitB16ToB32>(srcVreg0, srcOrigin0, mask); |
| 52 | - MicroAPI::LoadAlign<half, MicroAPI::LoadDist::DIST_UNPACK_B16>(srcOrigin1, src1 + i * oneRepElm); | 52 | + Reg::LoadAlign<half, Reg::LoadDist::DIST_UNPACK_B16>(srcOrigin1, src1 + i * oneRepElm); |
| 53 | - MicroAPI::Cast<float, half, castTraitB16ToB32>(srcVreg1, srcOrigin1, mask); | 53 | + Reg::Cast<float, half, castTraitB16ToB32>(srcVreg1, srcOrigin1, mask); |
| 54 | } else { | 54 | } else { |
| 55 | - MicroAPI::LoadAlign(srcVreg0, src0 + i * oneRepElm); | 55 | + Reg::LoadAlign(srcVreg0, src0 + i * oneRepElm); |
| 56 | - MicroAPI::LoadAlign(srcVreg1, src1 + i * oneRepElm); | 56 | + Reg::LoadAlign(srcVreg1, src1 + i * oneRepElm); |
| 57 | } | 57 | } |
| 58 | - MicroAPI::Mul(tmpReg0, srcVreg1, srcVreg1, mask); | 58 | + Reg::Mul(tmpReg0, srcVreg1, srcVreg1, mask); |
| 59 | - MicroAPI::Adds(tmpReg0, tmpReg0, gegluConstantA, mask); | 59 | + Reg::Adds(tmpReg0, tmpReg0, gegluConstantA, mask); |
| 60 | - MicroAPI::Mul(tmpReg0, tmpReg0, srcVreg1, mask); | 60 | + Reg::Mul(tmpReg0, tmpReg0, srcVreg1, mask); |
| 61 | - MicroAPI::Muls(tmpReg0, tmpReg0, gegluConstantB, mask); | 61 | + Reg::Muls(tmpReg0, tmpReg0, gegluConstantB, mask); |
| 62 | - MicroAPI::Exp(tmpReg1, tmpReg0, mask); | 62 | + Reg::Exp(tmpReg1, tmpReg0, mask); |
| 63 | - MicroAPI::Adds(tmpReg1, tmpReg1, 1.0f, mask); | 63 | + Reg::Adds(tmpReg1, tmpReg1, 1.0f, mask); |
| 64 | - MicroAPI::Div(tmpReg1, srcVreg1, tmpReg1, mask); | 64 | + Reg::Div(tmpReg1, srcVreg1, tmpReg1, mask); |
| 65 | - MicroAPI::Mul(dstVreg, srcVreg0, tmpReg1, mask); | 65 | + Reg::Mul(dstVreg, srcVreg0, tmpReg1, mask); |
| 66 | if constexpr (sizeof(T) == sizeof(half)) { | 66 | if constexpr (sizeof(T) == sizeof(half)) { |
| 67 | - MicroAPI::Cast<half, float, castTraitB32ToB16>((MicroAPI::RegTensor<half>&)dstVreg, dstVreg, mask); | 67 | + Reg::Cast<half, float, castTraitB32ToB16>((Reg::RegTensor<half>&)dstVreg, dstVreg, mask); |
| 68 | - MicroAPI::StoreAlign<half, MicroAPI::StoreDist::DIST_PACK_B32>( | 68 | + Reg::StoreAlign<half, Reg::StoreDist::DIST_PACK_B32>( |
| 69 | - dst + i * oneRepElm, (MicroAPI::RegTensor<half>&)dstVreg, mask); | 69 | + dst + i * oneRepElm, (Reg::RegTensor<half>&)dstVreg, mask); |
| 70 | } else { | 70 | } else { |
| 71 | - MicroAPI::StoreAlign(dst + i * oneRepElm, dstVreg, mask); | 71 | + Reg::StoreAlign(dst + i * oneRepElm, dstVreg, mask); |
| 72 | } | 72 | } |
| 73 | } | 73 | } |
| 74 | } | 74 | } |
| @@ -36,87 +36,87 @@ __simd_vf__ inline void GeluImplVF(__ubuf__ T* dst, __ubuf__ T* src, uint32_t co | |||
| 36 | constexpr uint32_t oneRepElm = static_cast<uint32_t>(GetVecLen() / sizeof(T)); | 36 | constexpr uint32_t oneRepElm = static_cast<uint32_t>(GetVecLen() / sizeof(T)); |
| 37 | constexpr float coefficientsA = 0.044715; | 37 | constexpr float coefficientsA = 0.044715; |
| 38 | constexpr float coefficientsB = 1.5957691216057308; | 38 | constexpr float coefficientsB = 1.5957691216057308; |
| 39 | - MicroAPI::RegTensor<T> srcVreg; | 39 | + Reg::RegTensor<T> srcVreg; |
| 40 | - MicroAPI::RegTensor<T> dstVreg; | 40 | + Reg::RegTensor<T> dstVreg; |
| 41 | - MicroAPI::RegTensor<T> tmpReg0; | 41 | + Reg::RegTensor<T> tmpReg0; |
| 42 | - MicroAPI::RegTensor<T> tmpReg1; | 42 | + Reg::RegTensor<T> tmpReg1; |
| 43 | - MicroAPI::RegTensor<T> tmpReg2; | 43 | + Reg::RegTensor<T> tmpReg2; |
| 44 | - MicroAPI::RegTensor<T> tmpReg3; | 44 | + Reg::RegTensor<T> tmpReg3; |
| 45 | - MicroAPI::MaskReg mask; | 45 | + Reg::MaskReg mask; |
| 46 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 46 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 47 | - mask = MicroAPI::UpdateMask<T>(count); | 47 | + mask = Reg::UpdateMask<T>(count); |
| 48 | if constexpr (highPrecision) { | 48 | if constexpr (highPrecision) { |
| 49 | - MicroAPI::LoadAlign<half, MicroAPI::LoadDist::DIST_UNPACK_B16>( | 49 | + Reg::LoadAlign<half, Reg::LoadDist::DIST_UNPACK_B16>( |
| 50 | - (MicroAPI::RegTensor<half>&)srcVreg, (__ubuf__ half*)src + i * oneRepElm); | 50 | + (Reg::RegTensor<half>&)srcVreg, (__ubuf__ half*)src + i * oneRepElm); |
| 51 | - MicroAPI::Cast<float, half, castTraitB16ToB32>(srcVreg, (MicroAPI::RegTensor<half>&)srcVreg, mask); | 51 | + Reg::Cast<float, half, castTraitB16ToB32>(srcVreg, (Reg::RegTensor<half>&)srcVreg, mask); |
| 52 | } else { | 52 | } else { |
| 53 | - MicroAPI::LoadAlign(srcVreg, src + i * oneRepElm); | 53 | + Reg::LoadAlign(srcVreg, src + i * oneRepElm); |
| 54 | } | 54 | } |
| 55 | // y = (input_x + 0.044715 * input_x ^ 3) * 1.5957691 | 55 | // y = (input_x + 0.044715 * input_x ^ 3) * 1.5957691 |
| 56 | - MicroAPI::Mul(tmpReg0, srcVreg, srcVreg, mask); | 56 | + Reg::Mul(tmpReg0, srcVreg, srcVreg, mask); |
| 57 | - MicroAPI::Mul(tmpReg0, tmpReg0, srcVreg, mask); | 57 | + Reg::Mul(tmpReg0, tmpReg0, srcVreg, mask); |
| 58 | - MicroAPI::Muls(tmpReg0, tmpReg0, coefficientsA, mask); | 58 | + Reg::Muls(tmpReg0, tmpReg0, coefficientsA, mask); |
| 59 | - MicroAPI::Add(tmpReg0, tmpReg0, srcVreg, mask); | 59 | + Reg::Add(tmpReg0, tmpReg0, srcVreg, mask); |
| 60 | - MicroAPI::Muls(tmpReg0, tmpReg0, coefficientsB, mask); | 60 | + Reg::Muls(tmpReg0, tmpReg0, coefficientsB, mask); |
| 61 | // exp(min(y, 0)) | 61 | // exp(min(y, 0)) |
| 62 | - MicroAPI::Mins(tmpReg1, tmpReg0, 0.0f, mask); | 62 | + Reg::Mins(tmpReg1, tmpReg0, 0.0f, mask); |
| 63 | - MicroAPI::Exp(tmpReg1, tmpReg1, mask); | 63 | + Reg::Exp(tmpReg1, tmpReg1, mask); |
| 64 | // x / (exp^(-abs(y)) + 1) | 64 | // x / (exp^(-abs(y)) + 1) |
| 65 | - MicroAPI::Abs(tmpReg2, tmpReg0, mask); | 65 | + Reg::Abs(tmpReg2, tmpReg0, mask); |
| 66 | - MicroAPI::Muls(tmpReg2, tmpReg2, -1.0f, mask); | 66 | + Reg::Muls(tmpReg2, tmpReg2, -1.0f, mask); |
| 67 | - MicroAPI::Exp(tmpReg3, tmpReg2, mask); | 67 | + Reg::Exp(tmpReg3, tmpReg2, mask); |
| 68 | - MicroAPI::Adds(tmpReg3, tmpReg3, 1.0f, mask); | 68 | + Reg::Adds(tmpReg3, tmpReg3, 1.0f, mask); |
| 69 | - MicroAPI::Div(tmpReg3, srcVreg, tmpReg3, mask); | 69 | + Reg::Div(tmpReg3, srcVreg, tmpReg3, mask); |
| 70 | // x / (exp^(-abs(y)) + 1) * exp(min(y, 0)) | 70 | // x / (exp^(-abs(y)) + 1) * exp(min(y, 0)) |
| 71 | - MicroAPI::Mul(dstVreg, tmpReg1, tmpReg3, mask); | 71 | + Reg::Mul(dstVreg, tmpReg1, tmpReg3, mask); |
| 72 | if constexpr (highPrecision) { | 72 | if constexpr (highPrecision) { |
| 73 | - MicroAPI::Cast<half, float, castTraitB32ToB16>((MicroAPI::RegTensor<half>&)dstVreg, dstVreg, mask); | 73 | + Reg::Cast<half, float, castTraitB32ToB16>((Reg::RegTensor<half>&)dstVreg, dstVreg, mask); |
| 74 | - MicroAPI::StoreAlign<half, MicroAPI::StoreDist::DIST_PACK_B32>( | 74 | + Reg::StoreAlign<half, Reg::StoreDist::DIST_PACK_B32>( |
| 75 | - (__ubuf__ half*)dst + i * oneRepElm, (MicroAPI::RegTensor<half>&)dstVreg, mask); | 75 | + (__ubuf__ half*)dst + i * oneRepElm, (Reg::RegTensor<half>&)dstVreg, mask); |
| 76 | } else { | 76 | } else { |
| 77 | - MicroAPI::StoreAlign(dst + i * oneRepElm, dstVreg, mask); | 77 | + Reg::StoreAlign(dst + i * oneRepElm, dstVreg, mask); |
| 78 | } | 78 | } |
| 79 | } | 79 | } |
| 80 | } | 80 | } |
| 81 | 81 | ||
| 82 | template <typename T> | 82 | template <typename T> |
| 83 | -__simd_callee__ inline void FastGeluCoreAlg(MicroAPI::RegTensor<T>& dstVreg, | 83 | +__simd_callee__ inline void FastGeluCoreAlg(Reg::RegTensor<T>& dstVreg, |
| 84 | - MicroAPI::RegTensor<T>& srcVreg, MicroAPI::MaskReg& mask, MicroAPI::RegTensor<T>& stackVreg) | 84 | + Reg::RegTensor<T>& srcVreg, Reg::MaskReg& mask, Reg::RegTensor<T>& stackVreg) |
| 85 | { | 85 | { |
| 86 | constexpr float coefficients = -1.702f; | 86 | constexpr float coefficients = -1.702f; |
| 87 | constexpr float oneFloatScalar = 1.0f; | 87 | constexpr float oneFloatScalar = 1.0f; |
| 88 | - MicroAPI::Muls(stackVreg, srcVreg, coefficients, mask); | 88 | + Reg::Muls(stackVreg, srcVreg, coefficients, mask); |
| 89 | - MicroAPI::Exp(stackVreg, stackVreg, mask); | 89 | + Reg::Exp(stackVreg, stackVreg, mask); |
| 90 | - MicroAPI::Adds(stackVreg, stackVreg, oneFloatScalar, mask); | 90 | + Reg::Adds(stackVreg, stackVreg, oneFloatScalar, mask); |
| 91 | - MicroAPI::Div(dstVreg, srcVreg, stackVreg, mask); | 91 | + Reg::Div(dstVreg, srcVreg, stackVreg, mask); |
| 92 | } | 92 | } |
| 93 | 93 | ||
| 94 | template <typename T = half> | 94 | template <typename T = half> |
| 95 | __simd_vf__ inline void FastGeluHighPrecisionAlgVF(__ubuf__ T* dst, __ubuf__ T* src, | 95 | __simd_vf__ inline void FastGeluHighPrecisionAlgVF(__ubuf__ T* dst, __ubuf__ T* src, |
| 96 | const uint32_t dataSize) | 96 | const uint32_t dataSize) |
| 97 | { | 97 | { |
| 98 | - MicroAPI::RegTensor<T> srcVreg; | 98 | + Reg::RegTensor<T> srcVreg; |
| 99 | - MicroAPI::RegTensor<float> srcVregFloat; | 99 | + Reg::RegTensor<float> srcVregFloat; |
| 100 | - MicroAPI::RegTensor<T> dstVreg; | 100 | + Reg::RegTensor<T> dstVreg; |
| 101 | - MicroAPI::RegTensor<float> dstVregFloat; | 101 | + Reg::RegTensor<float> dstVregFloat; |
| 102 | 102 | ||
| 103 | constexpr uint32_t stackSize = GetVecLen() / sizeof(float); | 103 | constexpr uint32_t stackSize = GetVecLen() / sizeof(float); |
| 104 | uint32_t sreg = dataSize; | 104 | uint32_t sreg = dataSize; |
| 105 | 105 | ||
| 106 | - MicroAPI::RegTensor<float> stackVregFloat; | 106 | + Reg::RegTensor<float> stackVregFloat; |
| 107 | 107 | ||
| 108 | - MicroAPI::MaskReg mask; | 108 | + Reg::MaskReg mask; |
| 109 | 109 | ||
| 110 | uint16_t repeatTimes = static_cast<uint16_t>(CeilDivision(dataSize, stackSize)); | 110 | uint16_t repeatTimes = static_cast<uint16_t>(CeilDivision(dataSize, stackSize)); |
| 111 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 111 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 112 | - mask = MicroAPI::UpdateMask<float>(sreg); | 112 | + mask = Reg::UpdateMask<float>(sreg); |
| 113 | - MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_UNPACK_B16>(srcVreg, src + i * stackSize); | 113 | + Reg::LoadAlign<T, Reg::LoadDist::DIST_UNPACK_B16>(srcVreg, src + i * stackSize); |
| 114 | - MicroAPI::Cast<float, half, castTraitB16ToB32>(srcVregFloat, srcVreg, mask); | 114 | + Reg::Cast<float, half, castTraitB16ToB32>(srcVregFloat, srcVreg, mask); |
| 115 | 115 | ||
| 116 | FastGeluCoreAlg<float>(dstVregFloat, srcVregFloat, mask, stackVregFloat); | 116 | FastGeluCoreAlg<float>(dstVregFloat, srcVregFloat, mask, stackVregFloat); |
| 117 | 117 | ||
| 118 | - MicroAPI::Cast<half, float, castTraitB32ToB16>(dstVreg, dstVregFloat, mask); | 118 | + Reg::Cast<half, float, castTraitB32ToB16>(dstVreg, dstVregFloat, mask); |
| 119 | - MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_PACK_B32>(dst + i * stackSize, dstVreg, mask); | 119 | + Reg::StoreAlign<T, Reg::StoreDist::DIST_PACK_B32>(dst + i * stackSize, dstVreg, mask); |
| 120 | } | 120 | } |
| 121 | } | 121 | } |
| 122 | 122 | ||
| @@ -134,18 +134,18 @@ template <typename T> | |||
| 134 | __simd_vf__ inline void FastGeluAlgVF(__ubuf__ T* dst, __ubuf__ T* src, | 134 | __simd_vf__ inline void FastGeluAlgVF(__ubuf__ T* dst, __ubuf__ T* src, |
| 135 | const uint32_t dataSize) | 135 | const uint32_t dataSize) |
| 136 | { | 136 | { |
| 137 | - MicroAPI::RegTensor<T> srcVreg; | 137 | + Reg::RegTensor<T> srcVreg; |
| 138 | - MicroAPI::RegTensor<T> dstVreg; | 138 | + Reg::RegTensor<T> dstVreg; |
| 139 | constexpr uint32_t stackSize = GetVecLen() / sizeof(T); | 139 | constexpr uint32_t stackSize = GetVecLen() / sizeof(T); |
| 140 | uint32_t sreg = dataSize; | 140 | uint32_t sreg = dataSize; |
| 141 | - MicroAPI::RegTensor<T> stackVreg; | 141 | + Reg::RegTensor<T> stackVreg; |
| 142 | - MicroAPI::MaskReg mask; | 142 | + Reg::MaskReg mask; |
| 143 | uint16_t repeatTimes = static_cast<uint16_t>(CeilDivision(dataSize, stackSize)); | 143 | uint16_t repeatTimes = static_cast<uint16_t>(CeilDivision(dataSize, stackSize)); |
| 144 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 144 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 145 | - mask = MicroAPI::UpdateMask<T>(sreg); | 145 | + mask = Reg::UpdateMask<T>(sreg); |
| 146 | - MicroAPI::LoadAlign<T>(srcVreg, src + i * stackSize); | 146 | + Reg::LoadAlign<T>(srcVreg, src + i * stackSize); |
| 147 | FastGeluCoreAlg<T>(dstVreg, srcVreg, mask, stackVreg); | 147 | FastGeluCoreAlg<T>(dstVreg, srcVreg, mask, stackVreg); |
| 148 | - MicroAPI::StoreAlign<T>(dst + i * stackSize, dstVreg, mask); | 148 | + Reg::StoreAlign<T>(dst + i * stackSize, dstVreg, mask); |
| 149 | } | 149 | } |
| 150 | } | 150 | } |
| 151 | 151 | ||
| @@ -160,9 +160,9 @@ __aicore__ inline void FastGeluAlg(const LocalTensor<T>& dstLocal, const LocalTe | |||
| 160 | } | 160 | } |
| 161 | 161 | ||
| 162 | template <typename T> | 162 | template <typename T> |
| 163 | -__simd_callee__ inline void FastGeluV2CoreAlg(MicroAPI::RegTensor<T>& dstVreg, | 163 | +__simd_callee__ inline void FastGeluV2CoreAlg(Reg::RegTensor<T>& dstVreg, |
| 164 | - MicroAPI::RegTensor<T>& srcVreg, MicroAPI::MaskReg& mask, MicroAPI::RegTensor<T>& stackVregA, | 164 | + Reg::RegTensor<T>& srcVreg, Reg::MaskReg& mask, Reg::RegTensor<T>& stackVregA, |
| 165 | - MicroAPI::RegTensor<T>& stackVregB, MicroAPI::RegTensor<T>& stackVregC) | 165 | + Reg::RegTensor<T>& stackVregB, Reg::RegTensor<T>& stackVregC) |
| 166 | { | 166 | { |
| 167 | constexpr float coefficients = 0.000000000001; | 167 | constexpr float coefficients = 0.000000000001; |
| 168 | constexpr float coefficientsHalf = 0.5; | 168 | constexpr float coefficientsHalf = 0.5; |
| @@ -171,54 +171,54 @@ __simd_callee__ inline void FastGeluV2CoreAlg(MicroAPI::RegTensor<T>& dstVreg, | |||
| 171 | constexpr float coefficientsBInv = 1.769; | 171 | constexpr float coefficientsBInv = 1.769; |
| 172 | constexpr float coefficientsC = 0.7071; | 172 | constexpr float coefficientsC = 0.7071; |
| 173 | constexpr float coefficientsD = 0.5; | 173 | constexpr float coefficientsD = 0.5; |
| 174 | - MicroAPI::Muls(stackVregA, srcVreg, coefficientsC, mask); | 174 | + Reg::Muls(stackVregA, srcVreg, coefficientsC, mask); |
| 175 | - MicroAPI::Abs(stackVregA, stackVregA, mask); | 175 | + Reg::Abs(stackVregA, stackVregA, mask); |
| 176 | - MicroAPI::Mins(stackVregA, stackVregA, coefficientsBInv, mask); | 176 | + Reg::Mins(stackVregA, stackVregA, coefficientsBInv, mask); |
| 177 | - MicroAPI::Adds(stackVregA, stackVregA, coefficientsB, mask); | 177 | + Reg::Adds(stackVregA, stackVregA, coefficientsB, mask); |
| 178 | - MicroAPI::Mul(stackVregA, stackVregA, stackVregA, mask); | 178 | + Reg::Mul(stackVregA, stackVregA, stackVregA, mask); |
| 179 | - MicroAPI::Muls(stackVregA, stackVregA, coefficientsA, mask); | 179 | + Reg::Muls(stackVregA, stackVregA, coefficientsA, mask); |
| 180 | - MicroAPI::Adds(stackVregA, stackVregA, coefficientsD, mask); | 180 | + Reg::Adds(stackVregA, stackVregA, coefficientsD, mask); |
| 181 | 181 | ||
| 182 | - MicroAPI::Adds(stackVregB, srcVreg, coefficients, mask); | 182 | + Reg::Adds(stackVregB, srcVreg, coefficients, mask); |
| 183 | - MicroAPI::Abs(stackVregC, stackVregB, mask); | 183 | + Reg::Abs(stackVregC, stackVregB, mask); |
| 184 | - MicroAPI::Div(stackVregB, stackVregB, stackVregC, mask); | 184 | + Reg::Div(stackVregB, stackVregB, stackVregC, mask); |
| 185 | 185 | ||
| 186 | - MicroAPI::Mul(stackVregA, stackVregA, stackVregB, mask); | 186 | + Reg::Mul(stackVregA, stackVregA, stackVregB, mask); |
| 187 | - MicroAPI::Adds(stackVregA, stackVregA, coefficientsHalf, mask); | 187 | + Reg::Adds(stackVregA, stackVregA, coefficientsHalf, mask); |
| 188 | 188 | ||
| 189 | - MicroAPI::Mul(dstVreg, srcVreg, stackVregA, mask); | 189 | + Reg::Mul(dstVreg, srcVreg, stackVregA, mask); |
| 190 | } | 190 | } |
| 191 | 191 | ||
| 192 | template <typename T = half> | 192 | template <typename T = half> |
| 193 | __simd_vf__ inline void FastGeluV2HighPrecisionAlgVF(__ubuf__ T* dst, __ubuf__ T* src, | 193 | __simd_vf__ inline void FastGeluV2HighPrecisionAlgVF(__ubuf__ T* dst, __ubuf__ T* src, |
| 194 | const uint32_t dataSize) | 194 | const uint32_t dataSize) |
| 195 | { | 195 | { |
| 196 | - MicroAPI::RegTensor<T> srcVreg; | 196 | + Reg::RegTensor<T> srcVreg; |
| 197 | - MicroAPI::RegTensor<float> srcVregFloat; | 197 | + Reg::RegTensor<float> srcVregFloat; |
| 198 | - MicroAPI::RegTensor<T> dstVreg; | 198 | + Reg::RegTensor<T> dstVreg; |
| 199 | - MicroAPI::RegTensor<float> dstVregFloat; | 199 | + Reg::RegTensor<float> dstVregFloat; |
| 200 | 200 | ||
| 201 | constexpr uint32_t stackSize = GetVecLen() / sizeof(float); | 201 | constexpr uint32_t stackSize = GetVecLen() / sizeof(float); |
| 202 | uint32_t sreg = dataSize; | 202 | uint32_t sreg = dataSize; |
| 203 | 203 | ||
| 204 | - MicroAPI::RegTensor<float> stackVregFloat; | 204 | + Reg::RegTensor<float> stackVregFloat; |
| 205 | 205 | ||
| 206 | - MicroAPI::MaskReg mask; | 206 | + Reg::MaskReg mask; |
| 207 | 207 | ||
| 208 | - MicroAPI::RegTensor<float> stackVregA; | 208 | + Reg::RegTensor<float> stackVregA; |
| 209 | - MicroAPI::RegTensor<float> stackVregB; | 209 | + Reg::RegTensor<float> stackVregB; |
| 210 | - MicroAPI::RegTensor<float> stackVregC; | 210 | + Reg::RegTensor<float> stackVregC; |
| 211 | 211 | ||
| 212 | uint16_t repeatTimes = static_cast<uint16_t>(CeilDivision(dataSize, stackSize)); | 212 | uint16_t repeatTimes = static_cast<uint16_t>(CeilDivision(dataSize, stackSize)); |
| 213 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 213 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 214 | - mask = MicroAPI::UpdateMask<float>(sreg); | 214 | + mask = Reg::UpdateMask<float>(sreg); |
| 215 | - MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_UNPACK_B16>(srcVreg, src + i * stackSize); | 215 | + Reg::LoadAlign<T, Reg::LoadDist::DIST_UNPACK_B16>(srcVreg, src + i * stackSize); |
| 216 | - MicroAPI::Cast<float, half, castTraitB16ToB32>(srcVregFloat, srcVreg, mask); | 216 | + Reg::Cast<float, half, castTraitB16ToB32>(srcVregFloat, srcVreg, mask); |
| 217 | 217 | ||
| 218 | FastGeluV2CoreAlg<float>(dstVregFloat, srcVregFloat, mask, stackVregA, stackVregB, stackVregC); | 218 | FastGeluV2CoreAlg<float>(dstVregFloat, srcVregFloat, mask, stackVregA, stackVregB, stackVregC); |
| 219 | 219 | ||
| 220 | - MicroAPI::Cast<half, float, castTraitB32ToB16>(dstVreg, dstVregFloat, mask); | 220 | + Reg::Cast<half, float, castTraitB32ToB16>(dstVreg, dstVregFloat, mask); |
| 221 | - MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_PACK_B32>(dst + i * stackSize, dstVreg, mask); | 221 | + Reg::StoreAlign<T, Reg::StoreDist::DIST_PACK_B32>(dst + i * stackSize, dstVreg, mask); |
| 222 | } | 222 | } |
| 223 | } | 223 | } |
| 224 | 224 | ||
| @@ -236,21 +236,21 @@ template <typename T> | |||
| 236 | __simd_vf__ inline void FastGeluV2AlgVF(__ubuf__ T* dst, __ubuf__ T* src, | 236 | __simd_vf__ inline void FastGeluV2AlgVF(__ubuf__ T* dst, __ubuf__ T* src, |
| 237 | const uint32_t dataSize) | 237 | const uint32_t dataSize) |
| 238 | { | 238 | { |
| 239 | - MicroAPI::RegTensor<T> srcVreg; | 239 | + Reg::RegTensor<T> srcVreg; |
| 240 | - MicroAPI::RegTensor<T> dstVreg; | 240 | + Reg::RegTensor<T> dstVreg; |
| 241 | constexpr uint32_t stackSize = GetVecLen() / sizeof(T); | 241 | constexpr uint32_t stackSize = GetVecLen() / sizeof(T); |
| 242 | uint32_t sreg = dataSize; | 242 | uint32_t sreg = dataSize; |
| 243 | 243 | ||
| 244 | - MicroAPI::RegTensor<T> stackVregA; | 244 | + Reg::RegTensor<T> stackVregA; |
| 245 | - MicroAPI::RegTensor<T> stackVregB; | 245 | + Reg::RegTensor<T> stackVregB; |
| 246 | - MicroAPI::RegTensor<T> stackVregC; | 246 | + Reg::RegTensor<T> stackVregC; |
| 247 | - MicroAPI::MaskReg mask; | 247 | + Reg::MaskReg mask; |
| 248 | uint16_t repeatTimes = static_cast<uint16_t>(CeilDivision(dataSize, stackSize)); | 248 | uint16_t repeatTimes = static_cast<uint16_t>(CeilDivision(dataSize, stackSize)); |
| 249 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 249 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 250 | - mask = MicroAPI::UpdateMask<T>(sreg); | 250 | + mask = Reg::UpdateMask<T>(sreg); |
| 251 | - MicroAPI::LoadAlign<T>(srcVreg, src + i * stackSize); | 251 | + Reg::LoadAlign<T>(srcVreg, src + i * stackSize); |
| 252 | FastGeluV2CoreAlg<T>(dstVreg, srcVreg, mask, stackVregA, stackVregB, stackVregC); | 252 | FastGeluV2CoreAlg<T>(dstVreg, srcVreg, mask, stackVregA, stackVregB, stackVregC); |
| 253 | - MicroAPI::StoreAlign<T>(dst + i * stackSize, dstVreg, mask); | 253 | + Reg::StoreAlign<T>(dst + i * stackSize, dstVreg, mask); |
| 254 | } | 254 | } |
| 255 | } | 255 | } |
| 256 | 256 | ||
| @@ -33,33 +33,33 @@ template <typename T> | |||
| 33 | __simd_vf__ inline void ReGluImplVF( | 33 | __simd_vf__ inline void ReGluImplVF( |
| 34 | __ubuf__ T* dst, __ubuf__ T* src0, __ubuf__ T* src1, uint32_t count, const uint16_t repeatTimes) | 34 | __ubuf__ T* dst, __ubuf__ T* src0, __ubuf__ T* src1, uint32_t count, const uint16_t repeatTimes) |
| 35 | { | 35 | { |
| 36 | - MicroAPI::RegTensor<T> srcOrigin0; | 36 | + Reg::RegTensor<T> srcOrigin0; |
| 37 | - MicroAPI::RegTensor<T> srcOrigin1; | 37 | + Reg::RegTensor<T> srcOrigin1; |
| 38 | - MicroAPI::RegTensor<float> srcVreg0; | 38 | + Reg::RegTensor<float> srcVreg0; |
| 39 | - MicroAPI::RegTensor<float> srcVreg1; | 39 | + Reg::RegTensor<float> srcVreg1; |
| 40 | - MicroAPI::RegTensor<float> tmpReg0; | 40 | + Reg::RegTensor<float> tmpReg0; |
| 41 | - MicroAPI::RegTensor<float> dstVreg; | 41 | + Reg::RegTensor<float> dstVreg; |
| 42 | - MicroAPI::MaskReg mask; | 42 | + Reg::MaskReg mask; |
| 43 | constexpr uint32_t oneRepElm = static_cast<uint32_t>(GetVecLen() / sizeof(float)); | 43 | constexpr uint32_t oneRepElm = static_cast<uint32_t>(GetVecLen() / sizeof(float)); |
| 44 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 44 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 45 | - mask = MicroAPI::UpdateMask<float>(count); | 45 | + mask = Reg::UpdateMask<float>(count); |
| 46 | if constexpr (sizeof(T) == sizeof(half)) { | 46 | if constexpr (sizeof(T) == sizeof(half)) { |
| 47 | - MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_UNPACK_B16>(srcOrigin0, src0 + i * oneRepElm); | 47 | + Reg::LoadAlign<T, Reg::LoadDist::DIST_UNPACK_B16>(srcOrigin0, src0 + i * oneRepElm); |
| 48 | - MicroAPI::Cast<float, T, castTraitB16ToB32>(srcVreg0, srcOrigin0, mask); | 48 | + Reg::Cast<float, T, castTraitB16ToB32>(srcVreg0, srcOrigin0, mask); |
| 49 | - MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_UNPACK_B16>(srcOrigin1, src1 + i * oneRepElm); | 49 | + Reg::LoadAlign<T, Reg::LoadDist::DIST_UNPACK_B16>(srcOrigin1, src1 + i * oneRepElm); |
| 50 | - MicroAPI::Cast<float, T, castTraitB16ToB32>(srcVreg1, srcOrigin1, mask); | 50 | + Reg::Cast<float, T, castTraitB16ToB32>(srcVreg1, srcOrigin1, mask); |
| 51 | } else { | 51 | } else { |
| 52 | - MicroAPI::LoadAlign(srcVreg0, src0 + i * oneRepElm); | 52 | + Reg::LoadAlign(srcVreg0, src0 + i * oneRepElm); |
| 53 | - MicroAPI::LoadAlign(srcVreg1, src1 + i * oneRepElm); | 53 | + Reg::LoadAlign(srcVreg1, src1 + i * oneRepElm); |
| 54 | } | 54 | } |
| 55 | - MicroAPI::Maxs(tmpReg0, srcVreg1, 0.0f, mask); | 55 | + Reg::Maxs(tmpReg0, srcVreg1, 0.0f, mask); |
| 56 | - MicroAPI::Mul(dstVreg, srcVreg0, tmpReg0, mask); | 56 | + Reg::Mul(dstVreg, srcVreg0, tmpReg0, mask); |
| 57 | if constexpr (sizeof(T) == sizeof(half)) { | 57 | if constexpr (sizeof(T) == sizeof(half)) { |
| 58 | - MicroAPI::Cast<T, float, castTraitB32ToB16>((MicroAPI::RegTensor<T>&)dstVreg, dstVreg, mask); | 58 | + Reg::Cast<T, float, castTraitB32ToB16>((Reg::RegTensor<T>&)dstVreg, dstVreg, mask); |
| 59 | - MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_PACK_B32>( | 59 | + Reg::StoreAlign<T, Reg::StoreDist::DIST_PACK_B32>( |
| 60 | - dst + i * oneRepElm, (MicroAPI::RegTensor<T>&)dstVreg, mask); | 60 | + dst + i * oneRepElm, (Reg::RegTensor<T>&)dstVreg, mask); |
| 61 | } else { | 61 | } else { |
| 62 | - MicroAPI::StoreAlign(dst + i * oneRepElm, dstVreg, mask); | 62 | + Reg::StoreAlign(dst + i * oneRepElm, dstVreg, mask); |
| 63 | } | 63 | } |
| 64 | } | 64 | } |
| 65 | } | 65 | } |
| @@ -35,31 +35,31 @@ template<typename T> | |||
| 35 | __simd_vf__ inline void SigmoidImplVF(__ubuf__ T* dstUb, __ubuf__ T* srcUb, uint32_t count, const uint16_t repeatTimes) | 35 | __simd_vf__ inline void SigmoidImplVF(__ubuf__ T* dstUb, __ubuf__ T* srcUb, uint32_t count, const uint16_t repeatTimes) |
| 36 | { | 36 | { |
| 37 | uint32_t sreg = count; | 37 | uint32_t sreg = count; |
| 38 | - MicroAPI::MaskReg preg; | 38 | + Reg::MaskReg preg; |
| 39 | - MicroAPI::RegTensor<T> srcReg; | 39 | + Reg::RegTensor<T> srcReg; |
| 40 | - MicroAPI::RegTensor<float> castReg; | 40 | + Reg::RegTensor<float> castReg; |
| 41 | - MicroAPI::RegTensor<float> tmpReg; | 41 | + Reg::RegTensor<float> tmpReg; |
| 42 | - MicroAPI::RegTensor<float> dstReg; | 42 | + Reg::RegTensor<float> dstReg; |
| 43 | 43 | ||
| 44 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 44 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 45 | - preg = MicroAPI::UpdateMask<float>(sreg); | 45 | + preg = Reg::UpdateMask<float>(sreg); |
| 46 | if constexpr (sizeof(T) == sizeof(half)) { | 46 | if constexpr (sizeof(T) == sizeof(half)) { |
| 47 | - MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_UNPACK_B16>(srcReg, srcUb + i * B32_DATA_NUM_PER_REPEAT); | 47 | + Reg::LoadAlign<T, Reg::LoadDist::DIST_UNPACK_B16>(srcReg, srcUb + i * B32_DATA_NUM_PER_REPEAT); |
| 48 | - MicroAPI::Cast<float, T, castTraitB16ToB32>(castReg, srcReg, preg); | 48 | + Reg::Cast<float, T, castTraitB16ToB32>(castReg, srcReg, preg); |
| 49 | } else { | 49 | } else { |
| 50 | - MicroAPI::LoadAlign(castReg, srcUb + i * B32_DATA_NUM_PER_REPEAT); | 50 | + Reg::LoadAlign(castReg, srcUb + i * B32_DATA_NUM_PER_REPEAT); |
| 51 | } | 51 | } |
| 52 | - MicroAPI::Muls(tmpReg, castReg, -1.0f, preg); | 52 | + Reg::Muls(tmpReg, castReg, -1.0f, preg); |
| 53 | - MicroAPI::Exp(tmpReg, tmpReg, preg); | 53 | + Reg::Exp(tmpReg, tmpReg, preg); |
| 54 | 54 | ||
| 55 | - MicroAPI::Adds(tmpReg, tmpReg, 1.0f, preg); | 55 | + Reg::Adds(tmpReg, tmpReg, 1.0f, preg); |
| 56 | - MicroAPI::Duplicate(dstReg, 1.0f, preg); | 56 | + Reg::Duplicate(dstReg, 1.0f, preg); |
| 57 | - MicroAPI::Div(dstReg, dstReg, tmpReg, preg); | 57 | + Reg::Div(dstReg, dstReg, tmpReg, preg); |
| 58 | if constexpr (sizeof(T) == sizeof(half)) { | 58 | if constexpr (sizeof(T) == sizeof(half)) { |
| 59 | - MicroAPI::Cast<T, float, castTraitB32ToB16>(srcReg, dstReg, preg); | 59 | + Reg::Cast<T, float, castTraitB32ToB16>(srcReg, dstReg, preg); |
| 60 | - MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_PACK_B32>(dstUb + i * B32_DATA_NUM_PER_REPEAT, srcReg, preg); | 60 | + Reg::StoreAlign<T, Reg::StoreDist::DIST_PACK_B32>(dstUb + i * B32_DATA_NUM_PER_REPEAT, srcReg, preg); |
| 61 | } else { | 61 | } else { |
| 62 | - MicroAPI::StoreAlign(dstUb + i * B32_DATA_NUM_PER_REPEAT, dstReg, preg); | 62 | + Reg::StoreAlign(dstUb + i * B32_DATA_NUM_PER_REPEAT, dstReg, preg); |
| 63 | } | 63 | } |
| 64 | } | 64 | } |
| 65 | } | 65 | } |
| @@ -31,18 +31,18 @@ template<typename T> | |||
| 31 | __simd_vf__ inline void SiluComputeVF(__ubuf__ T* dst, __ubuf__ T* src, uint32_t count, const uint16_t repeatTimes) | 31 | __simd_vf__ inline void SiluComputeVF(__ubuf__ T* dst, __ubuf__ T* src, uint32_t count, const uint16_t repeatTimes) |
| 32 | { | 32 | { |
| 33 | constexpr uint32_t oneRepElm = static_cast<uint32_t>(GetVecLen() / sizeof(T)); | 33 | constexpr uint32_t oneRepElm = static_cast<uint32_t>(GetVecLen() / sizeof(T)); |
| 34 | - MicroAPI::RegTensor<T> srcVreg; | 34 | + Reg::RegTensor<T> srcVreg; |
| 35 | - MicroAPI::RegTensor<T> tmpReg0; | 35 | + Reg::RegTensor<T> tmpReg0; |
| 36 | - MicroAPI::RegTensor<T> dstVreg; | 36 | + Reg::RegTensor<T> dstVreg; |
| 37 | - MicroAPI::MaskReg mask; | 37 | + Reg::MaskReg mask; |
| 38 | for (uint16_t i = 0; i < repeatTimes; ++i) { | 38 | for (uint16_t i = 0; i < repeatTimes; ++i) { |
| 39 | - mask = MicroAPI::UpdateMask<T>(count); | 39 | + mask = Reg::UpdateMask<T>(count); |
| 40 | - MicroAPI::LoadAlign(srcVreg, src + i * oneRepElm); | 40 | + Reg::LoadAlign(srcVreg, src + i * oneRepElm); |
| 41 | - MicroAPI::Muls(tmpReg0, srcVreg, -1.0f, mask); | 41 | + Reg::Muls(tmpReg0, srcVreg, -1.0f, mask); |
| 42 | - MicroAPI::Exp(tmpReg0, tmpReg0, mask); | 42 | + Reg::Exp(tmpReg0, tmpReg0, mask); |
| 43 | - MicroAPI::Adds(tmpReg0, tmpReg0, 1.0f, mask); | 43 | + Reg::Adds(tmpReg0, tmpReg0, 1.0f, mask); |
| 44 | - MicroAPI::Div(dstVreg, srcVreg, tmpReg0, mask); | 44 | + Reg::Div(dstVreg, srcVreg, tmpReg0, mask); |
| 45 | - MicroAPI::StoreAlign(dst + i * oneRepElm, dstVreg, mask); | 45 | + Reg::StoreAlign(dst + i * oneRepElm, dstVreg, mask); |
| 46 | } | 46 | } |
| 47 | } | 47 | } |
| 48 | } // namespace Internal | 48 | } // namespace Internal |
| @@ -31,43 +31,43 @@ __simd_vf__ inline void SimpleSoftMaxGenericNZImpl(__ubuf__ T1* dstUb, __ubuf__ | |||
| 31 | __ubuf__ T2* maxUb, __ubuf__ T1* srcUb, const uint16_t mRepeatTimes, | 31 | __ubuf__ T2* maxUb, __ubuf__ T1* srcUb, const uint16_t mRepeatTimes, |
| 32 | const uint16_t kRepeatTimes, const uint16_t outNum, const uint16_t dataBlock) | 32 | const uint16_t kRepeatTimes, const uint16_t outNum, const uint16_t dataBlock) |
| 33 | { | 33 | { |
| 34 | - MicroAPI::MaskReg maskCnt; | 34 | + Reg::MaskReg maskCnt; |
| 35 | - MicroAPI::MaskReg maskFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 35 | + Reg::MaskReg maskFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 36 | - MicroAPI::RegTensor<float> srcVreg; | 36 | + Reg::RegTensor<float> srcVreg; |
| 37 | - MicroAPI::RegTensor<float> maxVreg; | 37 | + Reg::RegTensor<float> maxVreg; |
| 38 | - MicroAPI::RegTensor<float> sumVreg; | 38 | + Reg::RegTensor<float> sumVreg; |
| 39 | - MicroAPI::RegTensor<float> tmpVreg; | 39 | + Reg::RegTensor<float> tmpVreg; |
| 40 | - MicroAPI::RegTensor<float> dstVreg; | 40 | + Reg::RegTensor<float> dstVreg; |
| 41 | - MicroAPI::RegTensor<float> maxVreg1; | 41 | + Reg::RegTensor<float> maxVreg1; |
| 42 | - MicroAPI::RegTensor<float> maxVreg2; | 42 | + Reg::RegTensor<float> maxVreg2; |
| 43 | - MicroAPI::RegTensor<float> sumVreg1; | 43 | + Reg::RegTensor<float> sumVreg1; |
| 44 | - MicroAPI::RegTensor<float> sumVreg2; | 44 | + Reg::RegTensor<float> sumVreg2; |
| 45 | 45 | ||
| 46 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 46 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 47 | uint32_t sreg = outNum; | 47 | uint32_t sreg = outNum; |
| 48 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 48 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 49 | - maskCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 49 | + maskCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 50 | LoadIfNeedCast<T2>(maxVreg, maxUb + i * FLOAT_REPEAT_SIZE, maskFull); | 50 | LoadIfNeedCast<T2>(maxVreg, maxUb + i * FLOAT_REPEAT_SIZE, maskFull); |
| 51 | LoadIfNeedCast<T2>(sumVreg, sumUb + i * FLOAT_REPEAT_SIZE, maskFull); | 51 | LoadIfNeedCast<T2>(sumVreg, sumUb + i * FLOAT_REPEAT_SIZE, maskFull); |
| 52 | if constexpr (SupportType<T2, float>()) { | 52 | if constexpr (SupportType<T2, float>()) { |
| 53 | - MicroAPI::Interleave(maxVreg1, maxVreg2, maxVreg, maxVreg); | 53 | + Reg::Interleave(maxVreg1, maxVreg2, maxVreg, maxVreg); |
| 54 | - MicroAPI::Interleave(sumVreg1, sumVreg2, sumVreg, sumVreg); | 54 | + Reg::Interleave(sumVreg1, sumVreg2, sumVreg, sumVreg); |
| 55 | LoadIfNeedCast<T1>(srcVreg, srcUb + 2 * i * FLOAT_REPEAT_SIZE + j * dataBlock, maskFull); | 55 | LoadIfNeedCast<T1>(srcVreg, srcUb + 2 * i * FLOAT_REPEAT_SIZE + j * dataBlock, maskFull); |
| 56 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg1, maskCnt); | 56 | + Reg::Sub(dstVreg, srcVreg, maxVreg1, maskCnt); |
| 57 | - MicroAPI::Exp(tmpVreg, dstVreg, maskCnt); | 57 | + Reg::Exp(tmpVreg, dstVreg, maskCnt); |
| 58 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg1, maskCnt); | 58 | + Reg::Div(dstVreg, tmpVreg, sumVreg1, maskCnt); |
| 59 | StoreIfNeedCast<T1>(dstUb + 2 * i * FLOAT_REPEAT_SIZE + j * dataBlock, dstVreg, maskCnt); | 59 | StoreIfNeedCast<T1>(dstUb + 2 * i * FLOAT_REPEAT_SIZE + j * dataBlock, dstVreg, maskCnt); |
| 60 | - maskCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 60 | + maskCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 61 | LoadIfNeedCast<T1>(srcVreg, srcUb + (2 * i + 1) * FLOAT_REPEAT_SIZE + j * dataBlock, maskFull); | 61 | LoadIfNeedCast<T1>(srcVreg, srcUb + (2 * i + 1) * FLOAT_REPEAT_SIZE + j * dataBlock, maskFull); |
| 62 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg2, maskCnt); | 62 | + Reg::Sub(dstVreg, srcVreg, maxVreg2, maskCnt); |
| 63 | - MicroAPI::Exp(tmpVreg, dstVreg, maskCnt); | 63 | + Reg::Exp(tmpVreg, dstVreg, maskCnt); |
| 64 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg2, maskCnt); | 64 | + Reg::Div(dstVreg, tmpVreg, sumVreg2, maskCnt); |
| 65 | StoreIfNeedCast<T1>(dstUb + (2 * i + 1) * FLOAT_REPEAT_SIZE + j * dataBlock, dstVreg, maskCnt); | 65 | StoreIfNeedCast<T1>(dstUb + (2 * i + 1) * FLOAT_REPEAT_SIZE + j * dataBlock, dstVreg, maskCnt); |
| 66 | } else { | 66 | } else { |
| 67 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, maskFull); | 67 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, maskFull); |
| 68 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, maskCnt); | 68 | + Reg::Sub(dstVreg, srcVreg, maxVreg, maskCnt); |
| 69 | - MicroAPI::Exp(tmpVreg, dstVreg, maskCnt); | 69 | + Reg::Exp(tmpVreg, dstVreg, maskCnt); |
| 70 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, maskCnt); | 70 | + Reg::Div(dstVreg, tmpVreg, sumVreg, maskCnt); |
| 71 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, dstVreg, maskCnt); | 71 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, dstVreg, maskCnt); |
| 72 | } | 72 | } |
| 73 | } | 73 | } |
| @@ -79,23 +79,23 @@ __simd_vf__ inline void SimpleSoftMaxGenericNDImpl(__ubuf__ T1* dstUb, __ubuf__ | |||
| 79 | __ubuf__ T2* maxUb, __ubuf__ T1* srcUb, const uint16_t srcM, const uint16_t srcK, | 79 | __ubuf__ T2* maxUb, __ubuf__ T1* srcUb, const uint16_t srcM, const uint16_t srcK, |
| 80 | const uint16_t repeatTimes, const uint16_t blockStride) | 80 | const uint16_t repeatTimes, const uint16_t blockStride) |
| 81 | { | 81 | { |
| 82 | - MicroAPI::MaskReg maskFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 82 | + Reg::MaskReg maskFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 83 | - MicroAPI::RegTensor<float> srcVreg; | 83 | + Reg::RegTensor<float> srcVreg; |
| 84 | - MicroAPI::RegTensor<float> maxVreg; | 84 | + Reg::RegTensor<float> maxVreg; |
| 85 | - MicroAPI::RegTensor<float> sumVreg; | 85 | + Reg::RegTensor<float> sumVreg; |
| 86 | - MicroAPI::RegTensor<float> tmpVreg; | 86 | + Reg::RegTensor<float> tmpVreg; |
| 87 | - MicroAPI::RegTensor<float> dstVreg; | 87 | + Reg::RegTensor<float> dstVreg; |
| 88 | 88 | ||
| 89 | for (uint16_t i = 0; i < srcM; ++i) { | 89 | for (uint16_t i = 0; i < srcM; ++i) { |
| 90 | LoadIfNeedCast<T2>(maxVreg, maxUb + i * blockStride, maskFull); | 90 | LoadIfNeedCast<T2>(maxVreg, maxUb + i * blockStride, maskFull); |
| 91 | LoadIfNeedCast<T2>(sumVreg, sumUb + i * blockStride, maskFull); | 91 | LoadIfNeedCast<T2>(sumVreg, sumUb + i * blockStride, maskFull); |
| 92 | - MicroAPI::Duplicate(maxVreg, maxVreg, maskFull); | 92 | + Reg::Duplicate(maxVreg, maxVreg, maskFull); |
| 93 | - MicroAPI::Duplicate(sumVreg, sumVreg, maskFull); | 93 | + Reg::Duplicate(sumVreg, sumVreg, maskFull); |
| 94 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 94 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 95 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, maskFull); | 95 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, maskFull); |
| 96 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, maskFull); | 96 | + Reg::Sub(dstVreg, srcVreg, maxVreg, maskFull); |
| 97 | - MicroAPI::Exp(tmpVreg, dstVreg, maskFull); | 97 | + Reg::Exp(tmpVreg, dstVreg, maskFull); |
| 98 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, maskFull); | 98 | + Reg::Div(dstVreg, tmpVreg, sumVreg, maskFull); |
| 99 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, dstVreg, maskFull); | 99 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, dstVreg, maskFull); |
| 100 | } | 100 | } |
| 101 | } | 101 | } |
| @@ -106,26 +106,26 @@ __simd_vf__ inline void SimpleSoftMaxGenericNDWithTailImpl(__ubuf__ T1* dstUb, _ | |||
| 106 | __ubuf__ T2* maxUb, __ubuf__ T1* srcUb, const uint16_t srcM, const uint16_t srcK, | 106 | __ubuf__ T2* maxUb, __ubuf__ T1* srcUb, const uint16_t srcM, const uint16_t srcK, |
| 107 | const uint16_t repeatTimes, const uint16_t blockStride) | 107 | const uint16_t repeatTimes, const uint16_t blockStride) |
| 108 | { | 108 | { |
| 109 | - MicroAPI::MaskReg maskCnt; | 109 | + Reg::MaskReg maskCnt; |
| 110 | - MicroAPI::MaskReg maskFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 110 | + Reg::MaskReg maskFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 111 | - MicroAPI::RegTensor<float> srcVreg; | 111 | + Reg::RegTensor<float> srcVreg; |
| 112 | - MicroAPI::RegTensor<float> maxVreg; | 112 | + Reg::RegTensor<float> maxVreg; |
| 113 | - MicroAPI::RegTensor<float> sumVreg; | 113 | + Reg::RegTensor<float> sumVreg; |
| 114 | - MicroAPI::RegTensor<float> tmpVreg; | 114 | + Reg::RegTensor<float> tmpVreg; |
| 115 | - MicroAPI::RegTensor<float> dstVreg; | 115 | + Reg::RegTensor<float> dstVreg; |
| 116 | 116 | ||
| 117 | for (uint16_t i = 0; i < srcM; ++i) { | 117 | for (uint16_t i = 0; i < srcM; ++i) { |
| 118 | LoadIfNeedCast<T2>(maxVreg, maxUb + i * blockStride, maskFull); | 118 | LoadIfNeedCast<T2>(maxVreg, maxUb + i * blockStride, maskFull); |
| 119 | LoadIfNeedCast<T2>(sumVreg, sumUb + i * blockStride, maskFull); | 119 | LoadIfNeedCast<T2>(sumVreg, sumUb + i * blockStride, maskFull); |
| 120 | - MicroAPI::Duplicate(maxVreg, maxVreg, maskFull); | 120 | + Reg::Duplicate(maxVreg, maxVreg, maskFull); |
| 121 | - MicroAPI::Duplicate(sumVreg, sumVreg, maskFull); | 121 | + Reg::Duplicate(sumVreg, sumVreg, maskFull); |
| 122 | uint32_t sreg = srcK; | 122 | uint32_t sreg = srcK; |
| 123 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 123 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 124 | - maskCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 124 | + maskCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 125 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, maskFull); | 125 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, maskFull); |
| 126 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, maskCnt); | 126 | + Reg::Sub(dstVreg, srcVreg, maxVreg, maskCnt); |
| 127 | - MicroAPI::Exp(tmpVreg, dstVreg, maskCnt); | 127 | + Reg::Exp(tmpVreg, dstVreg, maskCnt); |
| 128 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, maskCnt); | 128 | + Reg::Div(dstVreg, tmpVreg, sumVreg, maskCnt); |
| 129 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, dstVreg, maskCnt); | 129 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, dstVreg, maskCnt); |
| 130 | } | 130 | } |
| 131 | } | 131 | } |
| @@ -27,67 +27,67 @@ | |||
| 27 | namespace AscendC { | 27 | namespace AscendC { |
| 28 | template <typename T> | 28 | template <typename T> |
| 29 | __simd_callee__ inline void LoadIfNeedCast( | 29 | __simd_callee__ inline void LoadIfNeedCast( |
| 30 | - MicroAPI::RegTensor<float>& dstReg, __ubuf__ T* srcUb, MicroAPI::MaskReg& preg) | 30 | + Reg::RegTensor<float>& dstReg, __ubuf__ T* srcUb, Reg::MaskReg& preg) |
| 31 | { | 31 | { |
| 32 | if constexpr (sizeof(T) == 2) { | 32 | if constexpr (sizeof(T) == 2) { |
| 33 | - MicroAPI::RegTensor<T> tmpReg; | 33 | + Reg::RegTensor<T> tmpReg; |
| 34 | - MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_UNPACK_B16>(tmpReg, srcUb); | 34 | + Reg::LoadAlign<T, Reg::LoadDist::DIST_UNPACK_B16>(tmpReg, srcUb); |
| 35 | - MicroAPI::Cast<float, T, Internal::castTraitB16ToB32>(dstReg, tmpReg, preg); | 35 | + Reg::Cast<float, T, Internal::castTraitB16ToB32>(dstReg, tmpReg, preg); |
| 36 | } else { | 36 | } else { |
| 37 | - MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(dstReg, srcUb); | 37 | + Reg::LoadAlign<T, Reg::LoadDist::DIST_NORM>(dstReg, srcUb); |
| 38 | } | 38 | } |
| 39 | } | 39 | } |
| 40 | 40 | ||
| 41 | template <typename T> | 41 | template <typename T> |
| 42 | __simd_callee__ inline void LoadIfNeedCastM1( | 42 | __simd_callee__ inline void LoadIfNeedCastM1( |
| 43 | - MicroAPI::RegTensor<float>& dstReg, __ubuf__ T* srcUb, MicroAPI::MaskReg& preg) | 43 | + Reg::RegTensor<float>& dstReg, __ubuf__ T* srcUb, Reg::MaskReg& preg) |
| 44 | { | 44 | { |
| 45 | if constexpr (sizeof(T) == 2) { | 45 | if constexpr (sizeof(T) == 2) { |
| 46 | - MicroAPI::RegTensor<T> castVreg; | 46 | + Reg::RegTensor<T> castVreg; |
| 47 | - MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B16>(castVreg, srcUb); | 47 | + Reg::LoadAlign<T, Reg::LoadDist::DIST_BRC_B16>(castVreg, srcUb); |
| 48 | - MicroAPI::UnPack<uint32_t, uint16_t>( | 48 | + Reg::UnPack<uint32_t, uint16_t>( |
| 49 | - (MicroAPI::RegTensor<uint32_t>&)castVreg, (MicroAPI::RegTensor<uint16_t>&)castVreg); | 49 | + (Reg::RegTensor<uint32_t>&)castVreg, (Reg::RegTensor<uint16_t>&)castVreg); |
| 50 | - MicroAPI::Cast<float, T, Internal::castTraitB16ToB32>(dstReg, castVreg, preg); | 50 | + Reg::Cast<float, T, Internal::castTraitB16ToB32>(dstReg, castVreg, preg); |
| 51 | } else { | 51 | } else { |
| 52 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(dstReg, srcUb); | 52 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_BRC_B32>(dstReg, srcUb); |
| 53 | } | 53 | } |
| 54 | } | 54 | } |
| 55 | 55 | ||
| 56 | template <typename T> | 56 | template <typename T> |
| 57 | __simd_callee__ inline void StoreIfNeedCastM1( | 57 | __simd_callee__ inline void StoreIfNeedCastM1( |
| 58 | - __ubuf__ T* dstUb, MicroAPI::RegTensor<float>& srcReg, MicroAPI::MaskReg& preg) | 58 | + __ubuf__ T* dstUb, Reg::RegTensor<float>& srcReg, Reg::MaskReg& preg) |
| 59 | { | 59 | { |
| 60 | if constexpr (sizeof(T) == 2) { | 60 | if constexpr (sizeof(T) == 2) { |
| 61 | - MicroAPI::RegTensor<T> castVreg; | 61 | + Reg::RegTensor<T> castVreg; |
| 62 | - MicroAPI::Cast<T, float, Internal::castTraitB32ToB16>(castVreg, srcReg, preg); | 62 | + Reg::Cast<T, float, Internal::castTraitB32ToB16>(castVreg, srcReg, preg); |
| 63 | - MicroAPI::Pack<uint16_t, uint32_t>( | 63 | + Reg::Pack<uint16_t, uint32_t>( |
| 64 | - (MicroAPI::RegTensor<uint16_t>&)castVreg, (MicroAPI::RegTensor<uint32_t>&)castVreg); | 64 | + (Reg::RegTensor<uint16_t>&)castVreg, (Reg::RegTensor<uint32_t>&)castVreg); |
| 65 | - MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B16>(dstUb, castVreg, preg); | 65 | + Reg::StoreAlign<T, Reg::StoreDist::DIST_FIRST_ELEMENT_B16>(dstUb, castVreg, preg); |
| 66 | } else { | 66 | } else { |
| 67 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dstUb, srcReg, preg); | 67 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_FIRST_ELEMENT_B32>(dstUb, srcReg, preg); |
| 68 | } | 68 | } |
| 69 | } | 69 | } |
| 70 | 70 | ||
| 71 | template <typename T> | 71 | template <typename T> |
| 72 | __simd_callee__ inline void StoreIfNeedCast( | 72 | __simd_callee__ inline void StoreIfNeedCast( |
| 73 | - __ubuf__ T* dstUb, MicroAPI::RegTensor<float>& srcReg, MicroAPI::MaskReg& preg) | 73 | + __ubuf__ T* dstUb, Reg::RegTensor<float>& srcReg, Reg::MaskReg& preg) |
| 74 | { | 74 | { |
| 75 | if constexpr (sizeof(T) == 2) { | 75 | if constexpr (sizeof(T) == 2) { |
| 76 | - MicroAPI::RegTensor<T> tmpReg; | 76 | + Reg::RegTensor<T> tmpReg; |
| 77 | - MicroAPI::Cast<T, float, Internal::castTraitB32ToB16>(tmpReg, srcReg, preg); | 77 | + Reg::Cast<T, float, Internal::castTraitB32ToB16>(tmpReg, srcReg, preg); |
| 78 | - MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_PACK_B32>(dstUb, tmpReg, preg); | 78 | + Reg::StoreAlign<T, Reg::StoreDist::DIST_PACK_B32>(dstUb, tmpReg, preg); |
| 79 | } else { | 79 | } else { |
| 80 | - MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM>(dstUb, srcReg, preg); | 80 | + Reg::StoreAlign<T, Reg::StoreDist::DIST_NORM>(dstUb, srcReg, preg); |
| 81 | } | 81 | } |
| 82 | } | 82 | } |
| 83 | 83 | ||
| 84 | template <typename T> | 84 | template <typename T> |
| 85 | -__simd_callee__ inline void LoadE2B(MicroAPI::RegTensor<T>& dstReg, __ubuf__ T* srcUb) | 85 | +__simd_callee__ inline void LoadE2B(Reg::RegTensor<T>& dstReg, __ubuf__ T* srcUb) |
| 86 | { | 86 | { |
| 87 | if constexpr (sizeof(T) == 2) { | 87 | if constexpr (sizeof(T) == 2) { |
| 88 | - MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B16>(dstReg, srcUb); | 88 | + Reg::LoadAlign<T, Reg::LoadDist::DIST_E2B_B16>(dstReg, srcUb); |
| 89 | } else { | 89 | } else { |
| 90 | - MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(dstReg, srcUb); | 90 | + Reg::LoadAlign<T, Reg::LoadDist::DIST_E2B_B32>(dstReg, srcUb); |
| 91 | } | 91 | } |
| 92 | } | 92 | } |
| 93 | } // namespace AscendC | 93 | } // namespace AscendC |
| @@ -40,51 +40,51 @@ __simd_vf__ inline void SoftmaxFlashNDImpl(__ubuf__ T1* dstUb, __ubuf__ T2* sumU | |||
| 40 | NotNumUnion notNum; | 40 | NotNumUnion notNum; |
| 41 | notNum.i = F32_NEG_INF; | 41 | notNum.i = F32_NEG_INF; |
| 42 | 42 | ||
| 43 | - MicroAPI::MaskReg maskFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 43 | + Reg::MaskReg maskFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 44 | - MicroAPI::MaskReg maskExpMax = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL16>(); | 44 | + Reg::MaskReg maskExpMax = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL16>(); |
| 45 | - MicroAPI::MaskReg maskOneBlk; | 45 | + Reg::MaskReg maskOneBlk; |
| 46 | if constexpr (IsSameType<T2, half>::value) { | 46 | if constexpr (IsSameType<T2, half>::value) { |
| 47 | - maskOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL16>(); | 47 | + maskOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL16>(); |
| 48 | } else { | 48 | } else { |
| 49 | - maskOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 49 | + maskOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 50 | } | 50 | } |
| 51 | - MicroAPI::RegTensor<float> srcVreg; | 51 | + Reg::RegTensor<float> srcVreg; |
| 52 | - MicroAPI::RegTensor<float> maxVreg; | 52 | + Reg::RegTensor<float> maxVreg; |
| 53 | - MicroAPI::RegTensor<float> expMaxVreg; | 53 | + Reg::RegTensor<float> expMaxVreg; |
| 54 | - MicroAPI::RegTensor<float> sumVreg; | 54 | + Reg::RegTensor<float> sumVreg; |
| 55 | - MicroAPI::RegTensor<float> tmpVreg; | 55 | + Reg::RegTensor<float> tmpVreg; |
| 56 | - MicroAPI::RegTensor<float> dstVreg; | 56 | + Reg::RegTensor<float> dstVreg; |
| 57 | - MicroAPI::RegTensor<T1> t1Reg; | 57 | + Reg::RegTensor<T1> t1Reg; |
| 58 | 58 | ||
| 59 | for (uint16_t i = 0; i < srcM; ++i) { | 59 | for (uint16_t i = 0; i < srcM; ++i) { |
| 60 | - MicroAPI::Duplicate(maxVreg, notNum.f); | 60 | + Reg::Duplicate(maxVreg, notNum.f); |
| 61 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 61 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 62 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * stride, maskFull); | 62 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * stride, maskFull); |
| 63 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, maskFull); | 63 | + Reg::Max(maxVreg, maxVreg, srcVreg, maskFull); |
| 64 | } | 64 | } |
| 65 | - MicroAPI::ReduceMax(maxVreg, maxVreg, maskFull); | 65 | + Reg::ReduceMax(maxVreg, maxVreg, maskFull); |
| 66 | - MicroAPI::Duplicate(maxVreg, maxVreg, maskOneBlk); | 66 | + Reg::Duplicate(maxVreg, maxVreg, maskOneBlk); |
| 67 | LoadIfNeedCast<T2>(tmpVreg, inMaxUb + i * blockStride, maskOneBlk); | 67 | LoadIfNeedCast<T2>(tmpVreg, inMaxUb + i * blockStride, maskOneBlk); |
| 68 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, maskOneBlk); | 68 | + Reg::Max(maxVreg, maxVreg, tmpVreg, maskOneBlk); |
| 69 | StoreIfNeedCast<T2>(maxUb + i * blockStride, maxVreg, maskOneBlk); | 69 | StoreIfNeedCast<T2>(maxUb + i * blockStride, maxVreg, maskOneBlk); |
| 70 | 70 | ||
| 71 | - MicroAPI::FusedExpSub(expMaxVreg, tmpVreg, maxVreg, maskOneBlk); | 71 | + Reg::FusedExpSub(expMaxVreg, tmpVreg, maxVreg, maskOneBlk); |
| 72 | - MicroAPI::Duplicate(sumVreg, 0); | 72 | + Reg::Duplicate(sumVreg, 0); |
| 73 | - MicroAPI::Duplicate(maxVreg, maxVreg, maskFull); | 73 | + Reg::Duplicate(maxVreg, maxVreg, maskFull); |
| 74 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 74 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 75 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * stride, maskFull); | 75 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * stride, maskFull); |
| 76 | - MicroAPI::FusedExpSub(tmpVreg, srcVreg, maxVreg, maskFull); | 76 | + Reg::FusedExpSub(tmpVreg, srcVreg, maxVreg, maskFull); |
| 77 | StoreIfNeedCast<float>(workUb + i * srcK + j * stride, tmpVreg, maskFull); | 77 | StoreIfNeedCast<float>(workUb + i * srcK + j * stride, tmpVreg, maskFull); |
| 78 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, maskFull); | 78 | + Reg::Add(sumVreg, sumVreg, tmpVreg, maskFull); |
| 79 | } | 79 | } |
| 80 | - MicroAPI::ReduceSum(sumVreg, sumVreg, maskFull); | 80 | + Reg::ReduceSum(sumVreg, sumVreg, maskFull); |
| 81 | - MicroAPI::Duplicate(sumVreg, sumVreg, maskOneBlk); | 81 | + Reg::Duplicate(sumVreg, sumVreg, maskOneBlk); |
| 82 | LoadIfNeedCast<T2>(tmpVreg, inSumUb + i * blockStride, maskOneBlk); | 82 | LoadIfNeedCast<T2>(tmpVreg, inSumUb + i * blockStride, maskOneBlk); |
| 83 | - MicroAPI::MulAddDst(sumVreg, expMaxVreg, tmpVreg, maskOneBlk); | 83 | + Reg::MulAddDst(sumVreg, expMaxVreg, tmpVreg, maskOneBlk); |
| 84 | - MicroAPI::Mul(tmpVreg, expMaxVreg, tmpVreg, maskOneBlk); | 84 | + Reg::Mul(tmpVreg, expMaxVreg, tmpVreg, maskOneBlk); |
| 85 | - MicroAPI::Div(expMaxVreg, tmpVreg, sumVreg, maskOneBlk); | 85 | + Reg::Div(expMaxVreg, tmpVreg, sumVreg, maskOneBlk); |
| 86 | if constexpr (sizeof(T1) == 2 && sizeof (T2) == 4) { | 86 | if constexpr (sizeof(T1) == 2 && sizeof (T2) == 4) { |
| 87 | - MicroAPI::Interleave(expMaxVreg, tmpVreg, expMaxVreg, expMaxVreg); | 87 | + Reg::Interleave(expMaxVreg, tmpVreg, expMaxVreg, expMaxVreg); |
| 88 | StoreIfNeedCast<T1>(expMaxUb + i * blockStride * 2, expMaxVreg, maskExpMax); | 88 | StoreIfNeedCast<T1>(expMaxUb + i * blockStride * 2, expMaxVreg, maskExpMax); |
| 89 | } else { | 89 | } else { |
| 90 | StoreIfNeedCast<T1>(expMaxUb + i * blockStride, expMaxVreg, maskOneBlk); | 90 | StoreIfNeedCast<T1>(expMaxUb + i * blockStride, expMaxVreg, maskOneBlk); |
| @@ -92,22 +92,22 @@ __simd_vf__ inline void SoftmaxFlashNDImpl(__ubuf__ T1* dstUb, __ubuf__ T2* sumU | |||
| 92 | 92 | ||
| 93 | StoreIfNeedCast<T2>(sumUb + i * blockStride, sumVreg, maskOneBlk); | 93 | StoreIfNeedCast<T2>(sumUb + i * blockStride, sumVreg, maskOneBlk); |
| 94 | if constexpr (sizeof(T2) == sizeof(half)) { | 94 | if constexpr (sizeof(T2) == sizeof(half)) { |
| 95 | - MicroAPI::StoreAlign(tmpUb + i * blockStride, sumVreg, maskOneBlk); | 95 | + Reg::StoreAlign(tmpUb + i * blockStride, sumVreg, maskOneBlk); |
| 96 | } | 96 | } |
| 97 | } | 97 | } |
| 98 | 98 | ||
| 99 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 99 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 100 | 100 | ||
| 101 | for (uint16_t i = 0; i < srcM; ++i) { | 101 | for (uint16_t i = 0; i < srcM; ++i) { |
| 102 | if constexpr (sizeof(T2) == sizeof(half)) { | 102 | if constexpr (sizeof(T2) == sizeof(half)) { |
| 103 | - MicroAPI::LoadAlign(sumVreg, tmpUb + i * blockStride); | 103 | + Reg::LoadAlign(sumVreg, tmpUb + i * blockStride); |
| 104 | } else { | 104 | } else { |
| 105 | - MicroAPI::LoadAlign(sumVreg, sumUb + i * blockStride); | 105 | + Reg::LoadAlign(sumVreg, sumUb + i * blockStride); |
| 106 | } | 106 | } |
| 107 | - MicroAPI::Duplicate(sumVreg, sumVreg, maskFull); | 107 | + Reg::Duplicate(sumVreg, sumVreg, maskFull); |
| 108 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 108 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 109 | - MicroAPI::LoadAlign(tmpVreg, workUb + i * srcK + j * stride); | 109 | + Reg::LoadAlign(tmpVreg, workUb + i * srcK + j * stride); |
| 110 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, maskFull); | 110 | + Reg::Div(dstVreg, tmpVreg, sumVreg, maskFull); |
| 111 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * stride, dstVreg, maskFull); | 111 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * stride, dstVreg, maskFull); |
| 112 | } | 112 | } |
| 113 | } | 113 | } |
| @@ -124,87 +124,87 @@ __simd_vf__ inline void SoftmaxFlashNDWithTailImpl(__ubuf__ T1* dstUb, __ubuf__ | |||
| 124 | NotNumUnion notNum; | 124 | NotNumUnion notNum; |
| 125 | notNum.i = F32_NEG_INF; | 125 | notNum.i = F32_NEG_INF; |
| 126 | 126 | ||
| 127 | - MicroAPI::MaskReg maskCnt; | 127 | + Reg::MaskReg maskCnt; |
| 128 | - MicroAPI::MaskReg maskFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 128 | + Reg::MaskReg maskFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 129 | - MicroAPI::MaskReg maskExpMax = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL16>(); | 129 | + Reg::MaskReg maskExpMax = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL16>(); |
| 130 | - MicroAPI::MaskReg maskOneBlk; | 130 | + Reg::MaskReg maskOneBlk; |
| 131 | if constexpr (IsSameType<T2, half>::value) { | 131 | if constexpr (IsSameType<T2, half>::value) { |
| 132 | - maskOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL16>(); | 132 | + maskOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL16>(); |
| 133 | } else { | 133 | } else { |
| 134 | - maskOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 134 | + maskOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 135 | } | 135 | } |
| 136 | - MicroAPI::RegTensor<float> srcVreg; | 136 | + Reg::RegTensor<float> srcVreg; |
| 137 | - MicroAPI::RegTensor<float> maxVreg; | 137 | + Reg::RegTensor<float> maxVreg; |
| 138 | - MicroAPI::RegTensor<float> expMaxVreg; | 138 | + Reg::RegTensor<float> expMaxVreg; |
| 139 | - MicroAPI::RegTensor<float> sumVreg; | 139 | + Reg::RegTensor<float> sumVreg; |
| 140 | - MicroAPI::RegTensor<float> tmpVreg; | 140 | + Reg::RegTensor<float> tmpVreg; |
| 141 | - MicroAPI::RegTensor<float> minVreg; | 141 | + Reg::RegTensor<float> minVreg; |
| 142 | - MicroAPI::RegTensor<float> dstVreg; | 142 | + Reg::RegTensor<float> dstVreg; |
| 143 | - MicroAPI::RegTensor<T1> t1Reg; | 143 | + Reg::RegTensor<T1> t1Reg; |
| 144 | 144 | ||
| 145 | Duplicate(minVreg, notNum.f); | 145 | Duplicate(minVreg, notNum.f); |
| 146 | for (uint16_t i = 0; i < srcM; ++i) { | 146 | for (uint16_t i = 0; i < srcM; ++i) { |
| 147 | uint32_t sreg = originK; | 147 | uint32_t sreg = originK; |
| 148 | Duplicate(maxVreg, notNum.f); | 148 | Duplicate(maxVreg, notNum.f); |
| 149 | for (uint16_t j = 0; j < static_cast<uint16_t>(repeatTimes - 1); ++j) { | 149 | for (uint16_t j = 0; j < static_cast<uint16_t>(repeatTimes - 1); ++j) { |
| 150 | - maskCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 150 | + maskCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 151 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * stride, maskFull); | 151 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * stride, maskFull); |
| 152 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, maskCnt); | 152 | + Reg::Max(maxVreg, maxVreg, srcVreg, maskCnt); |
| 153 | } | 153 | } |
| 154 | - maskCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 154 | + maskCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 155 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + (repeatTimes - 1) * stride, maskFull); | 155 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + (repeatTimes - 1) * stride, maskFull); |
| 156 | - MicroAPI::Select(srcVreg, srcVreg, minVreg, maskCnt); | 156 | + Reg::Select(srcVreg, srcVreg, minVreg, maskCnt); |
| 157 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, maskFull); | 157 | + Reg::Max(maxVreg, maxVreg, srcVreg, maskFull); |
| 158 | 158 | ||
| 159 | - MicroAPI::ReduceMax(maxVreg, maxVreg, maskFull); | 159 | + Reg::ReduceMax(maxVreg, maxVreg, maskFull); |
| 160 | Duplicate(maxVreg, maxVreg, maskOneBlk); | 160 | Duplicate(maxVreg, maxVreg, maskOneBlk); |
| 161 | LoadIfNeedCast<T2>(tmpVreg, inMaxUb + i *blockStride, maskOneBlk); | 161 | LoadIfNeedCast<T2>(tmpVreg, inMaxUb + i *blockStride, maskOneBlk); |
| 162 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, maskOneBlk); | 162 | + Reg::Max(maxVreg, maxVreg, tmpVreg, maskOneBlk); |
| 163 | StoreIfNeedCast<T2>(maxUb + i * blockStride, maxVreg, maskOneBlk); | 163 | StoreIfNeedCast<T2>(maxUb + i * blockStride, maxVreg, maskOneBlk); |
| 164 | 164 | ||
| 165 | - MicroAPI::FusedExpSub(expMaxVreg, tmpVreg, maxVreg, maskOneBlk); | 165 | + Reg::FusedExpSub(expMaxVreg, tmpVreg, maxVreg, maskOneBlk); |
| 166 | Duplicate(sumVreg, 0); | 166 | Duplicate(sumVreg, 0); |
| 167 | Duplicate(maxVreg, maxVreg, maskFull); | 167 | Duplicate(maxVreg, maxVreg, maskFull); |
| 168 | sreg = originK; | 168 | sreg = originK; |
| 169 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 169 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 170 | - maskCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 170 | + maskCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 171 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * stride, maskFull); | 171 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * stride, maskFull); |
| 172 | - MicroAPI::FusedExpSub(tmpVreg, srcVreg, maxVreg, maskCnt); | 172 | + Reg::FusedExpSub(tmpVreg, srcVreg, maxVreg, maskCnt); |
| 173 | StoreIfNeedCast<float>(workUb + i * srcK + j * stride, tmpVreg, maskCnt); | 173 | StoreIfNeedCast<float>(workUb + i * srcK + j * stride, tmpVreg, maskCnt); |
| 174 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, maskFull); | 174 | + Reg::Add(sumVreg, sumVreg, tmpVreg, maskFull); |
| 175 | } | 175 | } |
| 176 | - MicroAPI::ReduceSum(sumVreg, sumVreg, maskFull); | 176 | + Reg::ReduceSum(sumVreg, sumVreg, maskFull); |
| 177 | Duplicate(sumVreg, sumVreg, maskOneBlk); | 177 | Duplicate(sumVreg, sumVreg, maskOneBlk); |
| 178 | LoadIfNeedCast<T2>(tmpVreg, inSumUb + i * blockStride, maskOneBlk); | 178 | LoadIfNeedCast<T2>(tmpVreg, inSumUb + i * blockStride, maskOneBlk); |
| 179 | - MicroAPI::MulAddDst(sumVreg, expMaxVreg, tmpVreg, maskOneBlk); | 179 | + Reg::MulAddDst(sumVreg, expMaxVreg, tmpVreg, maskOneBlk); |
| 180 | - MicroAPI::Mul(tmpVreg, expMaxVreg, tmpVreg, maskOneBlk); | 180 | + Reg::Mul(tmpVreg, expMaxVreg, tmpVreg, maskOneBlk); |
| 181 | - MicroAPI::Div(expMaxVreg, tmpVreg, sumVreg, maskOneBlk); | 181 | + Reg::Div(expMaxVreg, tmpVreg, sumVreg, maskOneBlk); |
| 182 | if constexpr (sizeof(T1) == 2 && sizeof (T2) == 4) { | 182 | if constexpr (sizeof(T1) == 2 && sizeof (T2) == 4) { |
| 183 | - MicroAPI::Interleave(expMaxVreg, tmpVreg, expMaxVreg, expMaxVreg); | 183 | + Reg::Interleave(expMaxVreg, tmpVreg, expMaxVreg, expMaxVreg); |
| 184 | StoreIfNeedCast<T1>(expMaxUb + i * blockStride * 2, expMaxVreg, maskExpMax); | 184 | StoreIfNeedCast<T1>(expMaxUb + i * blockStride * 2, expMaxVreg, maskExpMax); |
| 185 | } else { | 185 | } else { |
| 186 | StoreIfNeedCast<T1>(expMaxUb + i * blockStride, expMaxVreg, maskOneBlk); | 186 | StoreIfNeedCast<T1>(expMaxUb + i * blockStride, expMaxVreg, maskOneBlk); |
| 187 | } | 187 | } |
| 188 | StoreIfNeedCast<T2>(sumUb + i * blockStride, sumVreg, maskOneBlk); | 188 | StoreIfNeedCast<T2>(sumUb + i * blockStride, sumVreg, maskOneBlk); |
| 189 | if constexpr (sizeof(T2) == sizeof(half)) { | 189 | if constexpr (sizeof(T2) == sizeof(half)) { |
| 190 | - MicroAPI::StoreAlign(tmpUb + i * blockStride, sumVreg, maskOneBlk); | 190 | + Reg::StoreAlign(tmpUb + i * blockStride, sumVreg, maskOneBlk); |
| 191 | } | 191 | } |
| 192 | } | 192 | } |
| 193 | 193 | ||
| 194 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 194 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 195 | 195 | ||
| 196 | for (uint16_t i = 0; i < srcM; ++i) { | 196 | for (uint16_t i = 0; i < srcM; ++i) { |
| 197 | if constexpr (sizeof(T2) == sizeof(half)) { | 197 | if constexpr (sizeof(T2) == sizeof(half)) { |
| 198 | - MicroAPI::LoadAlign(sumVreg, tmpUb + i * blockStride); | 198 | + Reg::LoadAlign(sumVreg, tmpUb + i * blockStride); |
| 199 | } else { | 199 | } else { |
| 200 | - MicroAPI::LoadAlign(sumVreg, sumUb + i * blockStride); | 200 | + Reg::LoadAlign(sumVreg, sumUb + i * blockStride); |
| 201 | } | 201 | } |
| 202 | Duplicate(sumVreg, sumVreg, maskFull); | 202 | Duplicate(sumVreg, sumVreg, maskFull); |
| 203 | uint32_t sreg = originK; | 203 | uint32_t sreg = originK; |
| 204 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 204 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 205 | - maskCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 205 | + maskCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 206 | - MicroAPI::LoadAlign(tmpVreg, workUb + i * srcK + j * stride); | 206 | + Reg::LoadAlign(tmpVreg, workUb + i * srcK + j * stride); |
| 207 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, maskCnt); | 207 | + Reg::Div(dstVreg, tmpVreg, sumVreg, maskCnt); |
| 208 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * stride, dstVreg, maskCnt); | 208 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * stride, dstVreg, maskCnt); |
| 209 | } | 209 | } |
| 210 | } | 210 | } |
| @@ -40,44 +40,44 @@ __simd_vf__ inline void SoftmaxFlashV2M1NDUpdateVFImpl(__ubuf__ T1* dstUb, __ubu | |||
| 40 | NotNumUnion notNum; | 40 | NotNumUnion notNum; |
| 41 | notNum.i = F32_NEG_INF; | 41 | notNum.i = F32_NEG_INF; |
| 42 | 42 | ||
| 43 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 43 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 44 | - MicroAPI::MaskReg pregOnePt = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL1>(); | 44 | + Reg::MaskReg pregOnePt = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL1>(); |
| 45 | - MicroAPI::RegTensor<float> srcVreg; | 45 | + Reg::RegTensor<float> srcVreg; |
| 46 | - MicroAPI::RegTensor<float> maxVreg; | 46 | + Reg::RegTensor<float> maxVreg; |
| 47 | - MicroAPI::RegTensor<float> expMaxVreg; | 47 | + Reg::RegTensor<float> expMaxVreg; |
| 48 | - MicroAPI::RegTensor<float> sumVreg; | 48 | + Reg::RegTensor<float> sumVreg; |
| 49 | - MicroAPI::RegTensor<float> tmpVreg; | 49 | + Reg::RegTensor<float> tmpVreg; |
| 50 | - MicroAPI::RegTensor<float> dstVreg; | 50 | + Reg::RegTensor<float> dstVreg; |
| 51 | 51 | ||
| 52 | for (uint16_t i = 0; i < srcM; ++i) { | 52 | for (uint16_t i = 0; i < srcM; ++i) { |
| 53 | Duplicate(maxVreg, notNum.f); | 53 | Duplicate(maxVreg, notNum.f); |
| 54 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 54 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 55 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 55 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 56 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 56 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 57 | } | 57 | } |
| 58 | - MicroAPI::ReduceMax(maxVreg, maxVreg, pregFull); | 58 | + Reg::ReduceMax(maxVreg, maxVreg, pregFull); |
| 59 | if constexpr (isOutputReduceMax) { | 59 | if constexpr (isOutputReduceMax) { |
| 60 | StoreIfNeedCastM1<T2>(reduceMaxUb + i, maxVreg, pregOnePt); | 60 | StoreIfNeedCastM1<T2>(reduceMaxUb + i, maxVreg, pregOnePt); |
| 61 | } | 61 | } |
| 62 | LoadIfNeedCastM1<T2>(tmpVreg, inMaxUb + i, pregOnePt); | 62 | LoadIfNeedCastM1<T2>(tmpVreg, inMaxUb + i, pregOnePt); |
| 63 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, pregOnePt); | 63 | + Reg::Max(maxVreg, maxVreg, tmpVreg, pregOnePt); |
| 64 | StoreIfNeedCastM1<T2>(maxUb + i, maxVreg, pregOnePt); | 64 | StoreIfNeedCastM1<T2>(maxUb + i, maxVreg, pregOnePt); |
| 65 | 65 | ||
| 66 | - MicroAPI::FusedExpSub(expMaxVreg, tmpVreg, maxVreg, pregOnePt); | 66 | + Reg::FusedExpSub(expMaxVreg, tmpVreg, maxVreg, pregOnePt); |
| 67 | StoreIfNeedCastM1<T1>(expMaxUb + i, expMaxVreg, pregOnePt); | 67 | StoreIfNeedCastM1<T1>(expMaxUb + i, expMaxVreg, pregOnePt); |
| 68 | 68 | ||
| 69 | Duplicate(sumVreg, 0); | 69 | Duplicate(sumVreg, 0); |
| 70 | Duplicate(maxVreg, maxVreg, pregFull); | 70 | Duplicate(maxVreg, maxVreg, pregFull); |
| 71 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 71 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 72 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 72 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 73 | - MicroAPI::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregFull); | 73 | + Reg::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregFull); |
| 74 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregFull); | 74 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregFull); |
| 75 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 75 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 76 | } | 76 | } |
| 77 | - MicroAPI::ReduceSum(sumVreg, sumVreg, pregFull); | 77 | + Reg::ReduceSum(sumVreg, sumVreg, pregFull); |
| 78 | LoadIfNeedCastM1<T2>(tmpVreg, inExpSumUb + i, pregOnePt); | 78 | LoadIfNeedCastM1<T2>(tmpVreg, inExpSumUb + i, pregOnePt); |
| 79 | - MicroAPI::Mul(tmpVreg, expMaxVreg, tmpVreg, pregOnePt); | 79 | + Reg::Mul(tmpVreg, expMaxVreg, tmpVreg, pregOnePt); |
| 80 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregOnePt); | 80 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregOnePt); |
| 81 | StoreIfNeedCastM1<T2>(expSumUb + i, sumVreg, pregOnePt); | 81 | StoreIfNeedCastM1<T2>(expSumUb + i, sumVreg, pregOnePt); |
| 82 | } | 82 | } |
| 83 | } | 83 | } |
| @@ -113,56 +113,56 @@ __simd_vf__ inline void SoftmaxFlashV2M1NDWithTailUpdateVFImpl(__ubuf__ T1* dstU | |||
| 113 | NotNumUnion notNum; | 113 | NotNumUnion notNum; |
| 114 | notNum.i = F32_NEG_INF; | 114 | notNum.i = F32_NEG_INF; |
| 115 | 115 | ||
| 116 | - MicroAPI::MaskReg pregCnt; | 116 | + Reg::MaskReg pregCnt; |
| 117 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 117 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 118 | - MicroAPI::MaskReg pregOnePt = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL1>(); | 118 | + Reg::MaskReg pregOnePt = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL1>(); |
| 119 | - MicroAPI::RegTensor<float> srcVreg; | 119 | + Reg::RegTensor<float> srcVreg; |
| 120 | - MicroAPI::RegTensor<float> maxVreg; | 120 | + Reg::RegTensor<float> maxVreg; |
| 121 | - MicroAPI::RegTensor<float> expMaxVreg; | 121 | + Reg::RegTensor<float> expMaxVreg; |
| 122 | - MicroAPI::RegTensor<float> sumVreg; | 122 | + Reg::RegTensor<float> sumVreg; |
| 123 | - MicroAPI::RegTensor<float> tmpVreg; | 123 | + Reg::RegTensor<float> tmpVreg; |
| 124 | - MicroAPI::RegTensor<float> minVreg; | 124 | + Reg::RegTensor<float> minVreg; |
| 125 | - MicroAPI::RegTensor<float> dstVreg; | 125 | + Reg::RegTensor<float> dstVreg; |
| 126 | 126 | ||
| 127 | Duplicate(minVreg, notNum.f); | 127 | Duplicate(minVreg, notNum.f); |
| 128 | for (uint16_t i = 0; i < srcM; ++i) { | 128 | for (uint16_t i = 0; i < srcM; ++i) { |
| 129 | uint32_t sreg = originK; | 129 | uint32_t sreg = originK; |
| 130 | Duplicate(maxVreg, notNum.f); | 130 | Duplicate(maxVreg, notNum.f); |
| 131 | for (uint16_t j = 0; j < static_cast<uint16_t>(repeatTimes - 1); ++j) { | 131 | for (uint16_t j = 0; j < static_cast<uint16_t>(repeatTimes - 1); ++j) { |
| 132 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 132 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 133 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 133 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 134 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregCnt); | 134 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregCnt); |
| 135 | } | 135 | } |
| 136 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 136 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 137 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + (repeatTimes - 1) * FLOAT_REPEAT_SIZE, pregFull); | 137 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + (repeatTimes - 1) * FLOAT_REPEAT_SIZE, pregFull); |
| 138 | - MicroAPI::Select(srcVreg, srcVreg, minVreg, pregCnt); | 138 | + Reg::Select(srcVreg, srcVreg, minVreg, pregCnt); |
| 139 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 139 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 140 | 140 | ||
| 141 | - MicroAPI::ReduceMax(maxVreg, maxVreg, pregFull); | 141 | + Reg::ReduceMax(maxVreg, maxVreg, pregFull); |
| 142 | if constexpr (isOutputReduceMax) { | 142 | if constexpr (isOutputReduceMax) { |
| 143 | StoreIfNeedCastM1<T2>(reduceMaxUb + i, maxVreg, pregOnePt); | 143 | StoreIfNeedCastM1<T2>(reduceMaxUb + i, maxVreg, pregOnePt); |
| 144 | } | 144 | } |
| 145 | LoadIfNeedCastM1<T2>(tmpVreg, inMaxUb + i, pregOnePt); | 145 | LoadIfNeedCastM1<T2>(tmpVreg, inMaxUb + i, pregOnePt); |
| 146 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, pregOnePt); | 146 | + Reg::Max(maxVreg, maxVreg, tmpVreg, pregOnePt); |
| 147 | StoreIfNeedCastM1<T2>(maxUb + i, maxVreg, pregOnePt); | 147 | StoreIfNeedCastM1<T2>(maxUb + i, maxVreg, pregOnePt); |
| 148 | 148 | ||
| 149 | - MicroAPI::FusedExpSub(expMaxVreg, tmpVreg, maxVreg, pregOnePt); | 149 | + Reg::FusedExpSub(expMaxVreg, tmpVreg, maxVreg, pregOnePt); |
| 150 | StoreIfNeedCastM1<T1>(expMaxUb + i, expMaxVreg, pregOnePt); | 150 | StoreIfNeedCastM1<T1>(expMaxUb + i, expMaxVreg, pregOnePt); |
| 151 | 151 | ||
| 152 | Duplicate(sumVreg, 0); | 152 | Duplicate(sumVreg, 0); |
| 153 | Duplicate(maxVreg, maxVreg, pregFull); | 153 | Duplicate(maxVreg, maxVreg, pregFull); |
| 154 | sreg = originK; | 154 | sreg = originK; |
| 155 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 155 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 156 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 156 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 157 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 157 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 158 | - MicroAPI::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregCnt); | 158 | + Reg::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregCnt); |
| 159 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); | 159 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); |
| 160 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 160 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 161 | } | 161 | } |
| 162 | - MicroAPI::ReduceSum(sumVreg, sumVreg, pregFull); | 162 | + Reg::ReduceSum(sumVreg, sumVreg, pregFull); |
| 163 | LoadIfNeedCastM1<T2>(tmpVreg, inExpSumUb + i, pregOnePt); | 163 | LoadIfNeedCastM1<T2>(tmpVreg, inExpSumUb + i, pregOnePt); |
| 164 | - MicroAPI::Mul(tmpVreg, expMaxVreg, tmpVreg, pregOnePt); | 164 | + Reg::Mul(tmpVreg, expMaxVreg, tmpVreg, pregOnePt); |
| 165 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregOnePt); | 165 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregOnePt); |
| 166 | StoreIfNeedCastM1<T2>(expSumUb + i, sumVreg, pregOnePt); | 166 | StoreIfNeedCastM1<T2>(expSumUb + i, sumVreg, pregOnePt); |
| 167 | } | 167 | } |
| 168 | } | 168 | } |
| @@ -205,94 +205,94 @@ __simd_vf__ inline void SoftmaxFlashV2NZNoUpdateVFImpl(__ubuf__ T1* dstUb, __ubu | |||
| 205 | NotNumUnion notNum; | 205 | NotNumUnion notNum; |
| 206 | notNum.i = F32_NEG_INF; | 206 | notNum.i = F32_NEG_INF; |
| 207 | 207 | ||
| 208 | - MicroAPI::MaskReg pregCnt; | 208 | + Reg::MaskReg pregCnt; |
| 209 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 209 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 210 | - MicroAPI::MaskReg pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 210 | + Reg::MaskReg pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 211 | - MicroAPI::RegTensor<float> srcVreg; | 211 | + Reg::RegTensor<float> srcVreg; |
| 212 | - MicroAPI::RegTensor<float> maxVreg; | 212 | + Reg::RegTensor<float> maxVreg; |
| 213 | - MicroAPI::RegTensor<float> sumVreg; | 213 | + Reg::RegTensor<float> sumVreg; |
| 214 | - MicroAPI::RegTensor<float> tmpVreg; | 214 | + Reg::RegTensor<float> tmpVreg; |
| 215 | - MicroAPI::RegTensor<float> dstVreg; | 215 | + Reg::RegTensor<float> dstVreg; |
| 216 | - MicroAPI::RegTensor<T2> castReg; | 216 | + Reg::RegTensor<T2> castReg; |
| 217 | 217 | ||
| 218 | // reducemax | 218 | // reducemax |
| 219 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 219 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 220 | Duplicate(maxVreg, notNum.f); | 220 | Duplicate(maxVreg, notNum.f); |
| 221 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 221 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 222 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); | 222 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); |
| 223 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 223 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 224 | } | 224 | } |
| 225 | - MicroAPI::ReduceMaxWithDataBlock(maxVreg, maxVreg, pregFull); | 225 | + Reg::ReduceMaxWithDataBlock(maxVreg, maxVreg, pregFull); |
| 226 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, maxVreg, pregOneBlk); | 226 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, maxVreg, pregOneBlk); |
| 227 | } | 227 | } |
| 228 | 228 | ||
| 229 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 229 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 230 | 230 | ||
| 231 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { | 231 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { |
| 232 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 232 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 233 | maxVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 233 | maxVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 234 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, pregFull); | 234 | + Reg::Max(maxVreg, maxVreg, tmpVreg, pregFull); |
| 235 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); | 235 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); |
| 236 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 236 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 237 | workUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); | 237 | workUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); |
| 238 | } | 238 | } |
| 239 | 239 | ||
| 240 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 240 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 241 | 241 | ||
| 242 | uint32_t sreg = originM * dtypeBlkStride; | 242 | uint32_t sreg = originM * dtypeBlkStride; |
| 243 | for (uint16_t i = 0; i < e2bRep; ++i) { | 243 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 244 | - pregCnt = MicroAPI::UpdateMask<T2>(sreg); | 244 | + pregCnt = Reg::UpdateMask<T2>(sreg); |
| 245 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); | 245 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); |
| 246 | - MicroAPI::StoreAlign(maxUb + i * dtypeRepStride, castReg, pregCnt); | 246 | + Reg::StoreAlign(maxUb + i * dtypeRepStride, castReg, pregCnt); |
| 247 | } | 247 | } |
| 248 | 248 | ||
| 249 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 249 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 250 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; | 250 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; |
| 251 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 251 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 252 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 252 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 253 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); | 253 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); |
| 254 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); | 254 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); |
| 255 | - MicroAPI::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregFull); | 255 | + Reg::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregFull); |
| 256 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); | 256 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); |
| 257 | if constexpr (sizeof(T1) == 2) { | 257 | if constexpr (sizeof(T1) == 2) { |
| 258 | - MicroAPI::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); | 258 | + Reg::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); |
| 259 | } | 259 | } |
| 260 | } | 260 | } |
| 261 | } | 261 | } |
| 262 | 262 | ||
| 263 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 263 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 264 | 264 | ||
| 265 | // reducesum | 265 | // reducesum |
| 266 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 266 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 267 | Duplicate(sumVreg, 0); | 267 | Duplicate(sumVreg, 0); |
| 268 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 268 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 269 | if constexpr (sizeof(T1) == 2) { | 269 | if constexpr (sizeof(T1) == 2) { |
| 270 | - MicroAPI::LoadAlign(srcVreg, expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); | 270 | + Reg::LoadAlign(srcVreg, expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); |
| 271 | } else { | 271 | } else { |
| 272 | - MicroAPI::LoadAlign(srcVreg, dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); | 272 | + Reg::LoadAlign(srcVreg, dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); |
| 273 | } | 273 | } |
| 274 | - MicroAPI::Add(sumVreg, sumVreg, srcVreg, pregFull); | 274 | + Reg::Add(sumVreg, sumVreg, srcVreg, pregFull); |
| 275 | } | 275 | } |
| 276 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); | 276 | + Reg::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); |
| 277 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); | 277 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); |
| 278 | } | 278 | } |
| 279 | 279 | ||
| 280 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 280 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 281 | 281 | ||
| 282 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { | 282 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { |
| 283 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 283 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 284 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 284 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 285 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 285 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 286 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); | 286 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); |
| 287 | } | 287 | } |
| 288 | 288 | ||
| 289 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 289 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 290 | 290 | ||
| 291 | sreg = originM * dtypeBlkStride; | 291 | sreg = originM * dtypeBlkStride; |
| 292 | for (uint16_t i = 0; i < e2bRep; ++i) { | 292 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 293 | - pregCnt = MicroAPI::UpdateMask<T2>(sreg); | 293 | + pregCnt = Reg::UpdateMask<T2>(sreg); |
| 294 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); | 294 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); |
| 295 | - MicroAPI::StoreAlign(expSumUb + i * dtypeRepStride, castReg, pregCnt); | 295 | + Reg::StoreAlign(expSumUb + i * dtypeRepStride, castReg, pregCnt); |
| 296 | } | 296 | } |
| 297 | } | 297 | } |
| 298 | 298 | ||
| @@ -335,18 +335,18 @@ __simd_vf__ inline void SoftmaxFlashV2NZWithTailNoUpdateVFImpl(__ubuf__ T1* dstU | |||
| 335 | NotNumUnion notNum; | 335 | NotNumUnion notNum; |
| 336 | notNum.i = F32_NEG_INF; | 336 | notNum.i = F32_NEG_INF; |
| 337 | 337 | ||
| 338 | - MicroAPI::MaskReg pregDst; | 338 | + Reg::MaskReg pregDst; |
| 339 | - MicroAPI::MaskReg pregkTail = MicroAPI::MoveMask<uint32_t>(); | 339 | + Reg::MaskReg pregkTail = Reg::MoveMask<uint32_t>(); |
| 340 | - MicroAPI::MaskReg pregCnt; | 340 | + Reg::MaskReg pregCnt; |
| 341 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 341 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 342 | - MicroAPI::MaskReg pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 342 | + Reg::MaskReg pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 343 | - MicroAPI::RegTensor<float> srcVreg; | 343 | + Reg::RegTensor<float> srcVreg; |
| 344 | - MicroAPI::RegTensor<float> maxVreg; | 344 | + Reg::RegTensor<float> maxVreg; |
| 345 | - MicroAPI::RegTensor<float> sumVreg; | 345 | + Reg::RegTensor<float> sumVreg; |
| 346 | - MicroAPI::RegTensor<float> tmpVreg; | 346 | + Reg::RegTensor<float> tmpVreg; |
| 347 | - MicroAPI::RegTensor<float> minVreg; | 347 | + Reg::RegTensor<float> minVreg; |
| 348 | - MicroAPI::RegTensor<float> dstVreg; | 348 | + Reg::RegTensor<float> dstVreg; |
| 349 | - MicroAPI::RegTensor<T2> castReg; | 349 | + Reg::RegTensor<T2> castReg; |
| 350 | 350 | ||
| 351 | // reducemax | 351 | // reducemax |
| 352 | Duplicate(minVreg, notNum.f); | 352 | Duplicate(minVreg, notNum.f); |
| @@ -354,63 +354,63 @@ __simd_vf__ inline void SoftmaxFlashV2NZWithTailNoUpdateVFImpl(__ubuf__ T1* dstU | |||
| 354 | Duplicate(maxVreg, notNum.f); | 354 | Duplicate(maxVreg, notNum.f); |
| 355 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 355 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 356 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); | 356 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); |
| 357 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 357 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 358 | } | 358 | } |
| 359 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, pregFull); | 359 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, pregFull); |
| 360 | - MicroAPI::Select(srcVreg, srcVreg, minVreg, pregkTail); | 360 | + Reg::Select(srcVreg, srcVreg, minVreg, pregkTail); |
| 361 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 361 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 362 | 362 | ||
| 363 | - MicroAPI::ReduceMaxWithDataBlock(maxVreg, maxVreg, pregFull); | 363 | + Reg::ReduceMaxWithDataBlock(maxVreg, maxVreg, pregFull); |
| 364 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, maxVreg, pregOneBlk); | 364 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, maxVreg, pregOneBlk); |
| 365 | } | 365 | } |
| 366 | 366 | ||
| 367 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 367 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 368 | 368 | ||
| 369 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { | 369 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { |
| 370 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 370 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 371 | maxVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 371 | maxVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 372 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, pregFull); | 372 | + Reg::Max(maxVreg, maxVreg, tmpVreg, pregFull); |
| 373 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); | 373 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); |
| 374 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 374 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 375 | workUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); | 375 | workUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); |
| 376 | } | 376 | } |
| 377 | 377 | ||
| 378 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 378 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 379 | 379 | ||
| 380 | uint32_t sreg = originM * dtypeBlkStride; | 380 | uint32_t sreg = originM * dtypeBlkStride; |
| 381 | for (uint16_t i = 0; i < e2bRep; ++i) { | 381 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 382 | - pregCnt = MicroAPI::UpdateMask<T2>(sreg); | 382 | + pregCnt = Reg::UpdateMask<T2>(sreg); |
| 383 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); | 383 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); |
| 384 | - MicroAPI::StoreAlign(maxUb + i * dtypeRepStride, castReg, pregCnt); | 384 | + Reg::StoreAlign(maxUb + i * dtypeRepStride, castReg, pregCnt); |
| 385 | } | 385 | } |
| 386 | 386 | ||
| 387 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 387 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 388 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; | 388 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; |
| 389 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 389 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 390 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 390 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 391 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); | 391 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); |
| 392 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); | 392 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); |
| 393 | - MicroAPI::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregFull); | 393 | + Reg::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregFull); |
| 394 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); | 394 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); |
| 395 | if constexpr (sizeof(T1) == 2) { | 395 | if constexpr (sizeof(T1) == 2) { |
| 396 | - MicroAPI::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); | 396 | + Reg::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); |
| 397 | } | 397 | } |
| 398 | } | 398 | } |
| 399 | } | 399 | } |
| 400 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; | 400 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; |
| 401 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 401 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 402 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 402 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 403 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); | 403 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); |
| 404 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, pregFull); | 404 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, pregFull); |
| 405 | - MicroAPI::MaskAnd(pregDst, pregkTail, pregCnt, pregFull); | 405 | + Reg::MaskAnd(pregDst, pregkTail, pregCnt, pregFull); |
| 406 | - MicroAPI::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregDst); | 406 | + Reg::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregDst); |
| 407 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, tmpVreg, pregDst); | 407 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, tmpVreg, pregDst); |
| 408 | if constexpr (sizeof(T1) == 2) { | 408 | if constexpr (sizeof(T1) == 2) { |
| 409 | - MicroAPI::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, tmpVreg, pregDst); | 409 | + Reg::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, tmpVreg, pregDst); |
| 410 | } | 410 | } |
| 411 | } | 411 | } |
| 412 | 412 | ||
| 413 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 413 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 414 | 414 | ||
| 415 | // reducesum | 415 | // reducesum |
| 416 | Duplicate(minVreg, 0); | 416 | Duplicate(minVreg, 0); |
| @@ -418,40 +418,40 @@ __simd_vf__ inline void SoftmaxFlashV2NZWithTailNoUpdateVFImpl(__ubuf__ T1* dstU | |||
| 418 | Duplicate(sumVreg, 0); | 418 | Duplicate(sumVreg, 0); |
| 419 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 419 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 420 | if constexpr (sizeof(T1) == 2) { | 420 | if constexpr (sizeof(T1) == 2) { |
| 421 | - MicroAPI::LoadAlign(srcVreg, expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); | 421 | + Reg::LoadAlign(srcVreg, expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); |
| 422 | } else { | 422 | } else { |
| 423 | - MicroAPI::LoadAlign(srcVreg, dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); | 423 | + Reg::LoadAlign(srcVreg, dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); |
| 424 | } | 424 | } |
| 425 | - MicroAPI::Add(sumVreg, sumVreg, srcVreg, pregFull); | 425 | + Reg::Add(sumVreg, sumVreg, srcVreg, pregFull); |
| 426 | } | 426 | } |
| 427 | if constexpr (sizeof(T1) == 2) { | 427 | if constexpr (sizeof(T1) == 2) { |
| 428 | - MicroAPI::LoadAlign(srcVreg, expUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock); | 428 | + Reg::LoadAlign(srcVreg, expUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock); |
| 429 | } else { | 429 | } else { |
| 430 | - MicroAPI::LoadAlign(srcVreg, dstUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock); | 430 | + Reg::LoadAlign(srcVreg, dstUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock); |
| 431 | } | 431 | } |
| 432 | - MicroAPI::Select(srcVreg, srcVreg, minVreg, pregkTail); | 432 | + Reg::Select(srcVreg, srcVreg, minVreg, pregkTail); |
| 433 | - MicroAPI::Add(sumVreg, sumVreg, srcVreg, pregFull); | 433 | + Reg::Add(sumVreg, sumVreg, srcVreg, pregFull); |
| 434 | 434 | ||
| 435 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); | 435 | + Reg::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); |
| 436 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); | 436 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); |
| 437 | } | 437 | } |
| 438 | 438 | ||
| 439 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 439 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 440 | 440 | ||
| 441 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { | 441 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { |
| 442 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 442 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 443 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 443 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 444 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 444 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 445 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); | 445 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); |
| 446 | } | 446 | } |
| 447 | 447 | ||
| 448 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 448 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 449 | 449 | ||
| 450 | sreg = originM * dtypeBlkStride; | 450 | sreg = originM * dtypeBlkStride; |
| 451 | for (uint16_t i = 0; i < e2bRep; ++i) { | 451 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 452 | - pregCnt = MicroAPI::UpdateMask<T2>(sreg); | 452 | + pregCnt = Reg::UpdateMask<T2>(sreg); |
| 453 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); | 453 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); |
| 454 | - MicroAPI::StoreAlign(expSumUb + i * dtypeRepStride, castReg, pregCnt); | 454 | + Reg::StoreAlign(expSumUb + i * dtypeRepStride, castReg, pregCnt); |
| 455 | } | 455 | } |
| 456 | } | 456 | } |
| 457 | 457 | ||
| @@ -501,124 +501,124 @@ __simd_vf__ inline void SoftmaxFlashV2NZUpdateVFImpl(__ubuf__ T1* dstUb, __ubuf_ | |||
| 501 | NotNumUnion notNum; | 501 | NotNumUnion notNum; |
| 502 | notNum.i = F32_NEG_INF; | 502 | notNum.i = F32_NEG_INF; |
| 503 | 503 | ||
| 504 | - MicroAPI::MaskReg pregCnt; | 504 | + Reg::MaskReg pregCnt; |
| 505 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 505 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 506 | - MicroAPI::MaskReg pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 506 | + Reg::MaskReg pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 507 | - MicroAPI::RegTensor<float> srcVreg; | 507 | + Reg::RegTensor<float> srcVreg; |
| 508 | - MicroAPI::RegTensor<float> maxVreg; | 508 | + Reg::RegTensor<float> maxVreg; |
| 509 | - MicroAPI::RegTensor<float> inMaxVreg; | 509 | + Reg::RegTensor<float> inMaxVreg; |
| 510 | - MicroAPI::RegTensor<float> sumVreg; | 510 | + Reg::RegTensor<float> sumVreg; |
| 511 | - MicroAPI::RegTensor<float> tmpVreg; | 511 | + Reg::RegTensor<float> tmpVreg; |
| 512 | - MicroAPI::RegTensor<float> dstVreg; | 512 | + Reg::RegTensor<float> dstVreg; |
| 513 | - MicroAPI::RegTensor<T1> t1Reg; | 513 | + Reg::RegTensor<T1> t1Reg; |
| 514 | 514 | ||
| 515 | // reducemax | 515 | // reducemax |
| 516 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 516 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 517 | Duplicate(maxVreg, notNum.f); | 517 | Duplicate(maxVreg, notNum.f); |
| 518 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 518 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 519 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); | 519 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); |
| 520 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 520 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 521 | } | 521 | } |
| 522 | - MicroAPI::ReduceMaxWithDataBlock(maxVreg, maxVreg, pregFull); | 522 | + Reg::ReduceMaxWithDataBlock(maxVreg, maxVreg, pregFull); |
| 523 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, maxVreg, pregOneBlk); | 523 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, maxVreg, pregOneBlk); |
| 524 | } | 524 | } |
| 525 | 525 | ||
| 526 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 526 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 527 | 527 | ||
| 528 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { | 528 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { |
| 529 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 529 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 530 | maxVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 530 | maxVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 531 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, pregFull); | 531 | + Reg::Max(maxVreg, maxVreg, tmpVreg, pregFull); |
| 532 | if constexpr (sizeof(T2) == 4) { | 532 | if constexpr (sizeof(T2) == 4) { |
| 533 | - MicroAPI::StoreAlign(workUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); | 533 | + Reg::StoreAlign(workUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); |
| 534 | } else { | 534 | } else { |
| 535 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 535 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 536 | workUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); | 536 | workUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); |
| 537 | } | 537 | } |
| 538 | } | 538 | } |
| 539 | 539 | ||
| 540 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 540 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 541 | 541 | ||
| 542 | uint32_t sreg = originM * dtypeBlkStride; | 542 | uint32_t sreg = originM * dtypeBlkStride; |
| 543 | for (uint16_t i = 0; i < e2bRep; ++i) { | 543 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 544 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 544 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 545 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); | 545 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); |
| 546 | LoadIfNeedCast<T2>(inMaxVreg, inMaxUb + i * FLOAT_REPEAT_SIZE, pregCnt); | 546 | LoadIfNeedCast<T2>(inMaxVreg, inMaxUb + i * FLOAT_REPEAT_SIZE, pregCnt); |
| 547 | - MicroAPI::Max(maxVreg, inMaxVreg, maxVreg, pregCnt); | 547 | + Reg::Max(maxVreg, inMaxVreg, maxVreg, pregCnt); |
| 548 | StoreIfNeedCast<T2>(maxUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregCnt); | 548 | StoreIfNeedCast<T2>(maxUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregCnt); |
| 549 | if constexpr (sizeof(T2) == 4) { | 549 | if constexpr (sizeof(T2) == 4) { |
| 550 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 550 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 551 | tmpUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); | 551 | tmpUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); |
| 552 | } else { | 552 | } else { |
| 553 | - MicroAPI::StoreAlign(tmpUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); | 553 | + Reg::StoreAlign(tmpUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); |
| 554 | } | 554 | } |
| 555 | - MicroAPI::FusedExpSub(dstVreg, inMaxVreg, maxVreg, pregCnt); | 555 | + Reg::FusedExpSub(dstVreg, inMaxVreg, maxVreg, pregCnt); |
| 556 | - MicroAPI::StoreAlign(expMaxF32Ub + i * FLOAT_REPEAT_SIZE, dstVreg, pregCnt); | 556 | + Reg::StoreAlign(expMaxF32Ub + i * FLOAT_REPEAT_SIZE, dstVreg, pregCnt); |
| 557 | } | 557 | } |
| 558 | 558 | ||
| 559 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 559 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 560 | 560 | ||
| 561 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 561 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 562 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; | 562 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; |
| 563 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 563 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 564 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 564 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 565 | - MicroAPI::LoadAlign(maxVreg, tmpUb + i * FLOAT_REPEAT_SIZE); | 565 | + Reg::LoadAlign(maxVreg, tmpUb + i * FLOAT_REPEAT_SIZE); |
| 566 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); | 566 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); |
| 567 | - MicroAPI::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregCnt); | 567 | + Reg::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregCnt); |
| 568 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); | 568 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); |
| 569 | if constexpr (sizeof(T1) == 2) { | 569 | if constexpr (sizeof(T1) == 2) { |
| 570 | - MicroAPI::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); | 570 | + Reg::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); |
| 571 | } | 571 | } |
| 572 | } | 572 | } |
| 573 | } | 573 | } |
| 574 | 574 | ||
| 575 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 575 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 576 | 576 | ||
| 577 | // reducesum | 577 | // reducesum |
| 578 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 578 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 579 | Duplicate(sumVreg, 0); | 579 | Duplicate(sumVreg, 0); |
| 580 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 580 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 581 | if constexpr (sizeof(T1) == 2) { | 581 | if constexpr (sizeof(T1) == 2) { |
| 582 | - MicroAPI::LoadAlign(tmpVreg, expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); | 582 | + Reg::LoadAlign(tmpVreg, expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); |
| 583 | } else { | 583 | } else { |
| 584 | - MicroAPI::LoadAlign(tmpVreg, dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); | 584 | + Reg::LoadAlign(tmpVreg, dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); |
| 585 | } | 585 | } |
| 586 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 586 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 587 | } | 587 | } |
| 588 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); | 588 | + Reg::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); |
| 589 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); | 589 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); |
| 590 | } | 590 | } |
| 591 | 591 | ||
| 592 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 592 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 593 | 593 | ||
| 594 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { | 594 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { |
| 595 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 595 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 596 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 596 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 597 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 597 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 598 | if constexpr (sizeof(T2) == 4) { | 598 | if constexpr (sizeof(T2) == 4) { |
| 599 | - MicroAPI::StoreAlign(workUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); | 599 | + Reg::StoreAlign(workUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); |
| 600 | } else { | 600 | } else { |
| 601 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 601 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 602 | workUb + i * HALF_REPEAT_SIZE, sumVreg, sumVreg, pregFull); | 602 | workUb + i * HALF_REPEAT_SIZE, sumVreg, sumVreg, pregFull); |
| 603 | } | 603 | } |
| 604 | } | 604 | } |
| 605 | 605 | ||
| 606 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 606 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 607 | 607 | ||
| 608 | sreg = originM * dtypeBlkStride; | 608 | sreg = originM * dtypeBlkStride; |
| 609 | for (uint16_t i = 0; i < e2bRep; ++i) { | 609 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 610 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 610 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 611 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); | 611 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); |
| 612 | LoadIfNeedCast<T2>(tmpVreg, inExpSumUb + i * FLOAT_REPEAT_SIZE, pregCnt); | 612 | LoadIfNeedCast<T2>(tmpVreg, inExpSumUb + i * FLOAT_REPEAT_SIZE, pregCnt); |
| 613 | - MicroAPI::LoadAlign(maxVreg, expMaxF32Ub + i * FLOAT_REPEAT_SIZE); | 613 | + Reg::LoadAlign(maxVreg, expMaxF32Ub + i * FLOAT_REPEAT_SIZE); |
| 614 | - MicroAPI::Mul(dstVreg, maxVreg, tmpVreg, pregCnt); | 614 | + Reg::Mul(dstVreg, maxVreg, tmpVreg, pregCnt); |
| 615 | - MicroAPI::Add(sumVreg, sumVreg, dstVreg, pregCnt); | 615 | + Reg::Add(sumVreg, sumVreg, dstVreg, pregCnt); |
| 616 | StoreIfNeedCast<T2>(expSumUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregCnt); | 616 | StoreIfNeedCast<T2>(expSumUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregCnt); |
| 617 | if constexpr (sizeof(T1) == 2 && sizeof(T2) == 4) { | 617 | if constexpr (sizeof(T1) == 2 && sizeof(T2) == 4) { |
| 618 | - MicroAPI::Cast<T1, float, Internal::castTraitB32ToB16>(t1Reg, maxVreg, pregCnt); | 618 | + Reg::Cast<T1, float, Internal::castTraitB32ToB16>(t1Reg, maxVreg, pregCnt); |
| 619 | - MicroAPI::Pack<uint16_t, uint32_t>((MicroAPI::RegTensor<uint16_t> &)t1Reg, | 619 | + Reg::Pack<uint16_t, uint32_t>((Reg::RegTensor<uint16_t> &)t1Reg, |
| 620 | - (MicroAPI::RegTensor<uint32_t> &)t1Reg); | 620 | + (Reg::RegTensor<uint32_t> &)t1Reg); |
| 621 | - MicroAPI::StoreAlign<T1, MicroAPI::StoreDist::DIST_INTLV_B16>( | 621 | + Reg::StoreAlign<T1, Reg::StoreDist::DIST_INTLV_B16>( |
| 622 | expMaxUb + i * HALF_REPEAT_SIZE, t1Reg, t1Reg, pregCnt); | 622 | expMaxUb + i * HALF_REPEAT_SIZE, t1Reg, t1Reg, pregCnt); |
| 623 | } else if constexpr (sizeof(T1) == 2 && sizeof(T2) == 2) { | 623 | } else if constexpr (sizeof(T1) == 2 && sizeof(T2) == 2) { |
| 624 | StoreIfNeedCast<T1>(expMaxUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregCnt); | 624 | StoreIfNeedCast<T1>(expMaxUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregCnt); |
| @@ -680,19 +680,19 @@ __simd_vf__ inline void SoftmaxFlashV2NZWithTailUpdateVFImpl(__ubuf__ T1* dstUb, | |||
| 680 | NotNumUnion notNum; | 680 | NotNumUnion notNum; |
| 681 | notNum.i = F32_NEG_INF; | 681 | notNum.i = F32_NEG_INF; |
| 682 | 682 | ||
| 683 | - MicroAPI::MaskReg pregDst; | 683 | + Reg::MaskReg pregDst; |
| 684 | - MicroAPI::MaskReg pregkTail = MicroAPI::MoveMask<uint32_t>(); | 684 | + Reg::MaskReg pregkTail = Reg::MoveMask<uint32_t>(); |
| 685 | - MicroAPI::MaskReg pregCnt; | 685 | + Reg::MaskReg pregCnt; |
| 686 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 686 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 687 | - MicroAPI::MaskReg pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 687 | + Reg::MaskReg pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 688 | - MicroAPI::RegTensor<float> srcVreg; | 688 | + Reg::RegTensor<float> srcVreg; |
| 689 | - MicroAPI::RegTensor<float> maxVreg; | 689 | + Reg::RegTensor<float> maxVreg; |
| 690 | - MicroAPI::RegTensor<float> inMaxVreg; | 690 | + Reg::RegTensor<float> inMaxVreg; |
| 691 | - MicroAPI::RegTensor<float> sumVreg; | 691 | + Reg::RegTensor<float> sumVreg; |
| 692 | - MicroAPI::RegTensor<float> tmpVreg; | 692 | + Reg::RegTensor<float> tmpVreg; |
| 693 | - MicroAPI::RegTensor<float> minVreg; | 693 | + Reg::RegTensor<float> minVreg; |
| 694 | - MicroAPI::RegTensor<float> dstVreg; | 694 | + Reg::RegTensor<float> dstVreg; |
| 695 | - MicroAPI::RegTensor<T1> t1Reg; | 695 | + Reg::RegTensor<T1> t1Reg; |
| 696 | 696 | ||
| 697 | // reducemax | 697 | // reducemax |
| 698 | Duplicate(minVreg, notNum.f); | 698 | Duplicate(minVreg, notNum.f); |
| @@ -700,78 +700,78 @@ __simd_vf__ inline void SoftmaxFlashV2NZWithTailUpdateVFImpl(__ubuf__ T1* dstUb, | |||
| 700 | Duplicate(maxVreg, notNum.f); | 700 | Duplicate(maxVreg, notNum.f); |
| 701 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 701 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 702 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); | 702 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); |
| 703 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 703 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 704 | } | 704 | } |
| 705 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, pregFull); | 705 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, pregFull); |
| 706 | - MicroAPI::Select(srcVreg, srcVreg, minVreg, pregkTail); | 706 | + Reg::Select(srcVreg, srcVreg, minVreg, pregkTail); |
| 707 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 707 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 708 | 708 | ||
| 709 | - MicroAPI::ReduceMaxWithDataBlock(maxVreg, maxVreg, pregFull); | 709 | + Reg::ReduceMaxWithDataBlock(maxVreg, maxVreg, pregFull); |
| 710 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, maxVreg, pregOneBlk); | 710 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, maxVreg, pregOneBlk); |
| 711 | } | 711 | } |
| 712 | 712 | ||
| 713 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 713 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 714 | 714 | ||
| 715 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { | 715 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { |
| 716 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 716 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 717 | maxVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 717 | maxVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 718 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, pregFull); | 718 | + Reg::Max(maxVreg, maxVreg, tmpVreg, pregFull); |
| 719 | if constexpr (sizeof(T2) == 4) { | 719 | if constexpr (sizeof(T2) == 4) { |
| 720 | - MicroAPI::StoreAlign(workUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); | 720 | + Reg::StoreAlign(workUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); |
| 721 | } else { | 721 | } else { |
| 722 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 722 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 723 | workUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); | 723 | workUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); |
| 724 | } | 724 | } |
| 725 | } | 725 | } |
| 726 | 726 | ||
| 727 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 727 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 728 | 728 | ||
| 729 | uint32_t sreg = originM * dtypeBlkStride; | 729 | uint32_t sreg = originM * dtypeBlkStride; |
| 730 | for (uint16_t i = 0; i < e2bRep; ++i) { | 730 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 731 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 731 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 732 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); | 732 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); |
| 733 | LoadIfNeedCast<T2>(inMaxVreg, inMaxUb + i * FLOAT_REPEAT_SIZE, pregCnt); | 733 | LoadIfNeedCast<T2>(inMaxVreg, inMaxUb + i * FLOAT_REPEAT_SIZE, pregCnt); |
| 734 | - MicroAPI::Max(maxVreg, inMaxVreg, maxVreg, pregCnt); | 734 | + Reg::Max(maxVreg, inMaxVreg, maxVreg, pregCnt); |
| 735 | StoreIfNeedCast<T2>(maxUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregCnt); | 735 | StoreIfNeedCast<T2>(maxUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregCnt); |
| 736 | if constexpr (sizeof(T2) == 4) { | 736 | if constexpr (sizeof(T2) == 4) { |
| 737 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 737 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 738 | tmpUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); | 738 | tmpUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); |
| 739 | } else { | 739 | } else { |
| 740 | - MicroAPI::StoreAlign(tmpUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); | 740 | + Reg::StoreAlign(tmpUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); |
| 741 | } | 741 | } |
| 742 | - MicroAPI::FusedExpSub(dstVreg, inMaxVreg, maxVreg, pregCnt); | 742 | + Reg::FusedExpSub(dstVreg, inMaxVreg, maxVreg, pregCnt); |
| 743 | - MicroAPI::StoreAlign(expMaxF32Ub + i * FLOAT_REPEAT_SIZE, dstVreg, pregCnt); | 743 | + Reg::StoreAlign(expMaxF32Ub + i * FLOAT_REPEAT_SIZE, dstVreg, pregCnt); |
| 744 | } | 744 | } |
| 745 | 745 | ||
| 746 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 746 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 747 | 747 | ||
| 748 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 748 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 749 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; | 749 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; |
| 750 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 750 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 751 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 751 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 752 | - MicroAPI::LoadAlign(maxVreg, tmpUb + i * FLOAT_REPEAT_SIZE); | 752 | + Reg::LoadAlign(maxVreg, tmpUb + i * FLOAT_REPEAT_SIZE); |
| 753 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); | 753 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); |
| 754 | - MicroAPI::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregCnt); | 754 | + Reg::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregCnt); |
| 755 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); | 755 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); |
| 756 | if constexpr (sizeof(T1) == 2) { | 756 | if constexpr (sizeof(T1) == 2) { |
| 757 | - MicroAPI::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); | 757 | + Reg::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregCnt); |
| 758 | } | 758 | } |
| 759 | } | 759 | } |
| 760 | } | 760 | } |
| 761 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; | 761 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; |
| 762 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 762 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 763 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 763 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 764 | - MicroAPI::LoadAlign(maxVreg, tmpUb + i * FLOAT_REPEAT_SIZE); | 764 | + Reg::LoadAlign(maxVreg, tmpUb + i * FLOAT_REPEAT_SIZE); |
| 765 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, pregFull); | 765 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, pregFull); |
| 766 | - MicroAPI::MaskAnd(pregDst, pregkTail, pregCnt, pregFull); | 766 | + Reg::MaskAnd(pregDst, pregkTail, pregCnt, pregFull); |
| 767 | - MicroAPI::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregDst); | 767 | + Reg::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregDst); |
| 768 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, tmpVreg, pregDst); | 768 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, tmpVreg, pregDst); |
| 769 | if constexpr (sizeof(T1) == 2) { | 769 | if constexpr (sizeof(T1) == 2) { |
| 770 | - MicroAPI::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, tmpVreg, pregDst); | 770 | + Reg::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, tmpVreg, pregDst); |
| 771 | } | 771 | } |
| 772 | } | 772 | } |
| 773 | 773 | ||
| 774 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 774 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 775 | 775 | ||
| 776 | // reducesum | 776 | // reducesum |
| 777 | Duplicate(minVreg, 0); | 777 | Duplicate(minVreg, 0); |
| @@ -779,54 +779,54 @@ __simd_vf__ inline void SoftmaxFlashV2NZWithTailUpdateVFImpl(__ubuf__ T1* dstUb, | |||
| 779 | Duplicate(sumVreg, 0); | 779 | Duplicate(sumVreg, 0); |
| 780 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 780 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 781 | if constexpr (sizeof(T1) == 2) { | 781 | if constexpr (sizeof(T1) == 2) { |
| 782 | - MicroAPI::LoadAlign(tmpVreg, expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); | 782 | + Reg::LoadAlign(tmpVreg, expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); |
| 783 | } else { | 783 | } else { |
| 784 | - MicroAPI::LoadAlign(tmpVreg, dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); | 784 | + Reg::LoadAlign(tmpVreg, dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); |
| 785 | } | 785 | } |
| 786 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 786 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 787 | } | 787 | } |
| 788 | if constexpr (sizeof(T1) == 2) { | 788 | if constexpr (sizeof(T1) == 2) { |
| 789 | - MicroAPI::LoadAlign(tmpVreg, expUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock); | 789 | + Reg::LoadAlign(tmpVreg, expUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock); |
| 790 | } else { | 790 | } else { |
| 791 | - MicroAPI::LoadAlign(tmpVreg, dstUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock); | 791 | + Reg::LoadAlign(tmpVreg, dstUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock); |
| 792 | } | 792 | } |
| 793 | - MicroAPI::Select(tmpVreg, tmpVreg, minVreg, pregkTail); | 793 | + Reg::Select(tmpVreg, tmpVreg, minVreg, pregkTail); |
| 794 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 794 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 795 | 795 | ||
| 796 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); | 796 | + Reg::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); |
| 797 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); | 797 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); |
| 798 | } | 798 | } |
| 799 | 799 | ||
| 800 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 800 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 801 | 801 | ||
| 802 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { | 802 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { |
| 803 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 803 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 804 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 804 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 805 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 805 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 806 | if constexpr (sizeof(T2) == 4) { | 806 | if constexpr (sizeof(T2) == 4) { |
| 807 | - MicroAPI::StoreAlign(workUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); | 807 | + Reg::StoreAlign(workUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); |
| 808 | } else { | 808 | } else { |
| 809 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 809 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 810 | workUb + i * HALF_REPEAT_SIZE, sumVreg, sumVreg, pregFull); | 810 | workUb + i * HALF_REPEAT_SIZE, sumVreg, sumVreg, pregFull); |
| 811 | } | 811 | } |
| 812 | } | 812 | } |
| 813 | 813 | ||
| 814 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 814 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 815 | 815 | ||
| 816 | sreg = originM * dtypeBlkStride; | 816 | sreg = originM * dtypeBlkStride; |
| 817 | for (uint16_t i = 0; i < e2bRep; ++i) { | 817 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 818 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 818 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 819 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); | 819 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); |
| 820 | LoadIfNeedCast<T2>(tmpVreg, inExpSumUb + i * FLOAT_REPEAT_SIZE, pregCnt); | 820 | LoadIfNeedCast<T2>(tmpVreg, inExpSumUb + i * FLOAT_REPEAT_SIZE, pregCnt); |
| 821 | - MicroAPI::LoadAlign(maxVreg, expMaxF32Ub + i * FLOAT_REPEAT_SIZE); | 821 | + Reg::LoadAlign(maxVreg, expMaxF32Ub + i * FLOAT_REPEAT_SIZE); |
| 822 | - MicroAPI::Mul(dstVreg, maxVreg, tmpVreg, pregCnt); | 822 | + Reg::Mul(dstVreg, maxVreg, tmpVreg, pregCnt); |
| 823 | - MicroAPI::Add(sumVreg, sumVreg, dstVreg, pregCnt); | 823 | + Reg::Add(sumVreg, sumVreg, dstVreg, pregCnt); |
| 824 | StoreIfNeedCast<T2>(expSumUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregCnt); | 824 | StoreIfNeedCast<T2>(expSumUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregCnt); |
| 825 | if constexpr (sizeof(T1) == 2 && sizeof(T2) == 4) { | 825 | if constexpr (sizeof(T1) == 2 && sizeof(T2) == 4) { |
| 826 | - MicroAPI::Cast<T1, float, Internal::castTraitB32ToB16>(t1Reg, maxVreg, pregCnt); | 826 | + Reg::Cast<T1, float, Internal::castTraitB32ToB16>(t1Reg, maxVreg, pregCnt); |
| 827 | - MicroAPI::Pack<uint16_t, uint32_t>((MicroAPI::RegTensor<uint16_t> &)t1Reg, | 827 | + Reg::Pack<uint16_t, uint32_t>((Reg::RegTensor<uint16_t> &)t1Reg, |
| 828 | - (MicroAPI::RegTensor<uint32_t> &)t1Reg); | 828 | + (Reg::RegTensor<uint32_t> &)t1Reg); |
| 829 | - MicroAPI::StoreAlign<T1, MicroAPI::StoreDist::DIST_INTLV_B16>( | 829 | + Reg::StoreAlign<T1, Reg::StoreDist::DIST_INTLV_B16>( |
| 830 | expMaxUb + i * HALF_REPEAT_SIZE, t1Reg, t1Reg, pregCnt); | 830 | expMaxUb + i * HALF_REPEAT_SIZE, t1Reg, t1Reg, pregCnt); |
| 831 | } else if constexpr (sizeof(T1) == 2 && sizeof(T2) == 2) { | 831 | } else if constexpr (sizeof(T1) == 2 && sizeof(T2) == 2) { |
| 832 | StoreIfNeedCast<T1>(expMaxUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregCnt); | 832 | StoreIfNeedCast<T1>(expMaxUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregCnt); |
| @@ -916,39 +916,39 @@ __simd_vf__ inline void SoftmaxFlashV2NDUpdateVFImpl(__ubuf__ T1* dstUb, __ubuf_ | |||
| 916 | NotNumUnion notNum; | 916 | NotNumUnion notNum; |
| 917 | notNum.i = F32_NEG_INF; | 917 | notNum.i = F32_NEG_INF; |
| 918 | 918 | ||
| 919 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 919 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 920 | - MicroAPI::MaskReg pregOneBlk; | 920 | + Reg::MaskReg pregOneBlk; |
| 921 | if constexpr (IsSameType<T2, half>::value) { | 921 | if constexpr (IsSameType<T2, half>::value) { |
| 922 | - pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL16>(); | 922 | + pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL16>(); |
| 923 | } else { | 923 | } else { |
| 924 | - pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 924 | + pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 925 | } | 925 | } |
| 926 | - MicroAPI::RegTensor<float> srcVreg; | 926 | + Reg::RegTensor<float> srcVreg; |
| 927 | - MicroAPI::RegTensor<float> maxVreg; | 927 | + Reg::RegTensor<float> maxVreg; |
| 928 | - MicroAPI::RegTensor<float> expMaxVreg; | 928 | + Reg::RegTensor<float> expMaxVreg; |
| 929 | - MicroAPI::RegTensor<float> sumVreg; | 929 | + Reg::RegTensor<float> sumVreg; |
| 930 | - MicroAPI::RegTensor<float> tmpVreg; | 930 | + Reg::RegTensor<float> tmpVreg; |
| 931 | - MicroAPI::RegTensor<float> dstVreg; | 931 | + Reg::RegTensor<float> dstVreg; |
| 932 | - MicroAPI::RegTensor<T1> t1Reg; | 932 | + Reg::RegTensor<T1> t1Reg; |
| 933 | 933 | ||
| 934 | for (uint16_t i = 0; i < srcM; ++i) { | 934 | for (uint16_t i = 0; i < srcM; ++i) { |
| 935 | Duplicate(maxVreg, notNum.f); | 935 | Duplicate(maxVreg, notNum.f); |
| 936 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 936 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 937 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 937 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 938 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 938 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 939 | } | 939 | } |
| 940 | - MicroAPI::ReduceMax(maxVreg, maxVreg, pregFull); | 940 | + Reg::ReduceMax(maxVreg, maxVreg, pregFull); |
| 941 | Duplicate(maxVreg, maxVreg, pregOneBlk); | 941 | Duplicate(maxVreg, maxVreg, pregOneBlk); |
| 942 | LoadIfNeedCast<T2>(tmpVreg, inMaxUb + i * blockStride, pregOneBlk); | 942 | LoadIfNeedCast<T2>(tmpVreg, inMaxUb + i * blockStride, pregOneBlk); |
| 943 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, pregOneBlk); | 943 | + Reg::Max(maxVreg, maxVreg, tmpVreg, pregOneBlk); |
| 944 | StoreIfNeedCast<T2>(maxUb + i * blockStride, maxVreg, pregOneBlk); | 944 | StoreIfNeedCast<T2>(maxUb + i * blockStride, maxVreg, pregOneBlk); |
| 945 | 945 | ||
| 946 | - MicroAPI::FusedExpSub(expMaxVreg, tmpVreg, maxVreg, pregOneBlk); | 946 | + Reg::FusedExpSub(expMaxVreg, tmpVreg, maxVreg, pregOneBlk); |
| 947 | if constexpr (sizeof(T1) == 2 && sizeof (T2) == 4) { | 947 | if constexpr (sizeof(T1) == 2 && sizeof (T2) == 4) { |
| 948 | - MicroAPI::Cast<T1, float, Internal::castTraitB32ToB16>(t1Reg, expMaxVreg, pregOneBlk); | 948 | + Reg::Cast<T1, float, Internal::castTraitB32ToB16>(t1Reg, expMaxVreg, pregOneBlk); |
| 949 | - MicroAPI::Pack<uint16_t, uint32_t>((MicroAPI::RegTensor<uint16_t> &)t1Reg, | 949 | + Reg::Pack<uint16_t, uint32_t>((Reg::RegTensor<uint16_t> &)t1Reg, |
| 950 | - (MicroAPI::RegTensor<uint32_t> &)t1Reg); | 950 | + (Reg::RegTensor<uint32_t> &)t1Reg); |
| 951 | - MicroAPI::StoreAlign<T1, MicroAPI::StoreDist::DIST_INTLV_B16>( | 951 | + Reg::StoreAlign<T1, Reg::StoreDist::DIST_INTLV_B16>( |
| 952 | expMaxUb + i * blockStride * 2, t1Reg, t1Reg, pregOneBlk); | 952 | expMaxUb + i * blockStride * 2, t1Reg, t1Reg, pregOneBlk); |
| 953 | } else { | 953 | } else { |
| 954 | StoreIfNeedCast<T1>(expMaxUb + i * blockStride, expMaxVreg, pregOneBlk); | 954 | StoreIfNeedCast<T1>(expMaxUb + i * blockStride, expMaxVreg, pregOneBlk); |
| @@ -958,15 +958,15 @@ __simd_vf__ inline void SoftmaxFlashV2NDUpdateVFImpl(__ubuf__ T1* dstUb, __ubuf_ | |||
| 958 | Duplicate(maxVreg, maxVreg, pregFull); | 958 | Duplicate(maxVreg, maxVreg, pregFull); |
| 959 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 959 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 960 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 960 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 961 | - MicroAPI::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregFull); | 961 | + Reg::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregFull); |
| 962 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregFull); | 962 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregFull); |
| 963 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 963 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 964 | } | 964 | } |
| 965 | - MicroAPI::ReduceSum(sumVreg, sumVreg, pregFull); | 965 | + Reg::ReduceSum(sumVreg, sumVreg, pregFull); |
| 966 | Duplicate(sumVreg, sumVreg, pregOneBlk); | 966 | Duplicate(sumVreg, sumVreg, pregOneBlk); |
| 967 | LoadIfNeedCast<T2>(tmpVreg, inExpSumUb + i * blockStride, pregOneBlk); | 967 | LoadIfNeedCast<T2>(tmpVreg, inExpSumUb + i * blockStride, pregOneBlk); |
| 968 | - MicroAPI::Mul(tmpVreg, expMaxVreg, tmpVreg, pregOneBlk); | 968 | + Reg::Mul(tmpVreg, expMaxVreg, tmpVreg, pregOneBlk); |
| 969 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregOneBlk); | 969 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregOneBlk); |
| 970 | StoreIfNeedCast<T2>(expSumUb + i * blockStride, sumVreg, pregOneBlk); | 970 | StoreIfNeedCast<T2>(expSumUb + i * blockStride, sumVreg, pregOneBlk); |
| 971 | } | 971 | } |
| 972 | } | 972 | } |
| @@ -1003,49 +1003,49 @@ __simd_vf__ inline void SoftmaxFlashV2NDWithTailUpdateVFImpl(__ubuf__ T1* dstUb, | |||
| 1003 | NotNumUnion notNum; | 1003 | NotNumUnion notNum; |
| 1004 | notNum.i = F32_NEG_INF; | 1004 | notNum.i = F32_NEG_INF; |
| 1005 | 1005 | ||
| 1006 | - MicroAPI::MaskReg pregCnt; | 1006 | + Reg::MaskReg pregCnt; |
| 1007 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 1007 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 1008 | - MicroAPI::MaskReg pregOneBlk; | 1008 | + Reg::MaskReg pregOneBlk; |
| 1009 | if constexpr (IsSameType<T2, half>::value) { | 1009 | if constexpr (IsSameType<T2, half>::value) { |
| 1010 | - pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL16>(); | 1010 | + pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL16>(); |
| 1011 | } else { | 1011 | } else { |
| 1012 | - pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 1012 | + pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 1013 | } | 1013 | } |
| 1014 | - MicroAPI::RegTensor<float> srcVreg; | 1014 | + Reg::RegTensor<float> srcVreg; |
| 1015 | - MicroAPI::RegTensor<float> maxVreg; | 1015 | + Reg::RegTensor<float> maxVreg; |
| 1016 | - MicroAPI::RegTensor<float> expMaxVreg; | 1016 | + Reg::RegTensor<float> expMaxVreg; |
| 1017 | - MicroAPI::RegTensor<float> sumVreg; | 1017 | + Reg::RegTensor<float> sumVreg; |
| 1018 | - MicroAPI::RegTensor<float> tmpVreg; | 1018 | + Reg::RegTensor<float> tmpVreg; |
| 1019 | - MicroAPI::RegTensor<float> minVreg; | 1019 | + Reg::RegTensor<float> minVreg; |
| 1020 | - MicroAPI::RegTensor<float> dstVreg; | 1020 | + Reg::RegTensor<float> dstVreg; |
| 1021 | - MicroAPI::RegTensor<T1> t1Reg; | 1021 | + Reg::RegTensor<T1> t1Reg; |
| 1022 | 1022 | ||
| 1023 | Duplicate(minVreg, notNum.f); | 1023 | Duplicate(minVreg, notNum.f); |
| 1024 | for (uint16_t i = 0; i < srcM; ++i) { | 1024 | for (uint16_t i = 0; i < srcM; ++i) { |
| 1025 | uint32_t sreg = originK; | 1025 | uint32_t sreg = originK; |
| 1026 | Duplicate(maxVreg, notNum.f); | 1026 | Duplicate(maxVreg, notNum.f); |
| 1027 | for (uint16_t j = 0; j < static_cast<uint16_t>(repeatTimes - 1); ++j) { | 1027 | for (uint16_t j = 0; j < static_cast<uint16_t>(repeatTimes - 1); ++j) { |
| 1028 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 1028 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 1029 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 1029 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 1030 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregCnt); | 1030 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregCnt); |
| 1031 | } | 1031 | } |
| 1032 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 1032 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 1033 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + (repeatTimes - 1) * FLOAT_REPEAT_SIZE, pregFull); | 1033 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + (repeatTimes - 1) * FLOAT_REPEAT_SIZE, pregFull); |
| 1034 | - MicroAPI::Select(srcVreg, srcVreg, minVreg, pregCnt); | 1034 | + Reg::Select(srcVreg, srcVreg, minVreg, pregCnt); |
| 1035 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 1035 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 1036 | 1036 | ||
| 1037 | - MicroAPI::ReduceMax(maxVreg, maxVreg, pregFull); | 1037 | + Reg::ReduceMax(maxVreg, maxVreg, pregFull); |
| 1038 | Duplicate(maxVreg, maxVreg, pregOneBlk); | 1038 | Duplicate(maxVreg, maxVreg, pregOneBlk); |
| 1039 | LoadIfNeedCast<T2>(tmpVreg, inMaxUb + i *blockStride, pregOneBlk); | 1039 | LoadIfNeedCast<T2>(tmpVreg, inMaxUb + i *blockStride, pregOneBlk); |
| 1040 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, pregOneBlk); | 1040 | + Reg::Max(maxVreg, maxVreg, tmpVreg, pregOneBlk); |
| 1041 | StoreIfNeedCast<T2>(maxUb + i * blockStride, maxVreg, pregOneBlk); | 1041 | StoreIfNeedCast<T2>(maxUb + i * blockStride, maxVreg, pregOneBlk); |
| 1042 | 1042 | ||
| 1043 | - MicroAPI::FusedExpSub(expMaxVreg, tmpVreg, maxVreg, pregOneBlk); | 1043 | + Reg::FusedExpSub(expMaxVreg, tmpVreg, maxVreg, pregOneBlk); |
| 1044 | if constexpr (sizeof(T1) == 2 && sizeof (T2) == 4) { | 1044 | if constexpr (sizeof(T1) == 2 && sizeof (T2) == 4) { |
| 1045 | - MicroAPI::Cast<T1, float, Internal::castTraitB32ToB16>(t1Reg, expMaxVreg, pregOneBlk); | 1045 | + Reg::Cast<T1, float, Internal::castTraitB32ToB16>(t1Reg, expMaxVreg, pregOneBlk); |
| 1046 | - MicroAPI::Pack<uint16_t, uint32_t>((MicroAPI::RegTensor<uint16_t> &)t1Reg, | 1046 | + Reg::Pack<uint16_t, uint32_t>((Reg::RegTensor<uint16_t> &)t1Reg, |
| 1047 | - (MicroAPI::RegTensor<uint32_t> &)t1Reg); | 1047 | + (Reg::RegTensor<uint32_t> &)t1Reg); |
| 1048 | - MicroAPI::StoreAlign<T1, MicroAPI::StoreDist::DIST_INTLV_B16>( | 1048 | + Reg::StoreAlign<T1, Reg::StoreDist::DIST_INTLV_B16>( |
| 1049 | expMaxUb + i * blockStride * 2, t1Reg, t1Reg, pregOneBlk); | 1049 | expMaxUb + i * blockStride * 2, t1Reg, t1Reg, pregOneBlk); |
| 1050 | } else { | 1050 | } else { |
| 1051 | StoreIfNeedCast<T1>(expMaxUb + i * blockStride, expMaxVreg, pregOneBlk); | 1051 | StoreIfNeedCast<T1>(expMaxUb + i * blockStride, expMaxVreg, pregOneBlk); |
| @@ -1055,17 +1055,17 @@ __simd_vf__ inline void SoftmaxFlashV2NDWithTailUpdateVFImpl(__ubuf__ T1* dstUb, | |||
| 1055 | Duplicate(maxVreg, maxVreg, pregFull); | 1055 | Duplicate(maxVreg, maxVreg, pregFull); |
| 1056 | sreg = originK; | 1056 | sreg = originK; |
| 1057 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 1057 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 1058 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 1058 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 1059 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 1059 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 1060 | - MicroAPI::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregCnt); | 1060 | + Reg::FusedExpSub(tmpVreg, srcVreg, maxVreg, pregCnt); |
| 1061 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); | 1061 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); |
| 1062 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 1062 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 1063 | } | 1063 | } |
| 1064 | - MicroAPI::ReduceSum(sumVreg, sumVreg, pregFull); | 1064 | + Reg::ReduceSum(sumVreg, sumVreg, pregFull); |
| 1065 | Duplicate(sumVreg, sumVreg, pregOneBlk); | 1065 | Duplicate(sumVreg, sumVreg, pregOneBlk); |
| 1066 | LoadIfNeedCast<T2>(tmpVreg, inExpSumUb + i * blockStride, pregOneBlk); | 1066 | LoadIfNeedCast<T2>(tmpVreg, inExpSumUb + i * blockStride, pregOneBlk); |
| 1067 | - MicroAPI::Mul(tmpVreg, expMaxVreg, tmpVreg, pregOneBlk); | 1067 | + Reg::Mul(tmpVreg, expMaxVreg, tmpVreg, pregOneBlk); |
| 1068 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregOneBlk); | 1068 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregOneBlk); |
| 1069 | StoreIfNeedCast<T2>(expSumUb + i * blockStride, sumVreg, pregOneBlk); | 1069 | StoreIfNeedCast<T2>(expSumUb + i * blockStride, sumVreg, pregOneBlk); |
| 1070 | } | 1070 | } |
| 1071 | } | 1071 | } |
| @@ -37,111 +37,111 @@ __simd_vf__ __aicore__ inline void SoftmaxFlashV3NDNoUpdateImpl(__ubuf__ T* dstU | |||
| 37 | constexpr uint32_t blockStride = GetDataBlockSizeInBytes() / sizeof(U); | 37 | constexpr uint32_t blockStride = GetDataBlockSizeInBytes() / sizeof(U); |
| 38 | constexpr uint16_t repeatTime = static_cast<uint16_t>(repeatStride / blockStride); | 38 | constexpr uint16_t repeatTime = static_cast<uint16_t>(repeatStride / blockStride); |
| 39 | 39 | ||
| 40 | - MicroAPI::MaskReg maskCnt; | 40 | + Reg::MaskReg maskCnt; |
| 41 | - MicroAPI::MaskReg maskFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 41 | + Reg::MaskReg maskFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 42 | - MicroAPI::MaskReg maskOnePt = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL1>(); | 42 | + Reg::MaskReg maskOnePt = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL1>(); |
| 43 | - MicroAPI::MaskReg maskOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 43 | + Reg::MaskReg maskOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 44 | - MicroAPI::RegTensor<float> srcVreg; | 44 | + Reg::RegTensor<float> srcVreg; |
| 45 | - MicroAPI::RegTensor<float> maxVreg; | 45 | + Reg::RegTensor<float> maxVreg; |
| 46 | - MicroAPI::RegTensor<float> sumVreg; | 46 | + Reg::RegTensor<float> sumVreg; |
| 47 | - MicroAPI::RegTensor<float> meanVreg; | 47 | + Reg::RegTensor<float> meanVreg; |
| 48 | - MicroAPI::RegTensor<float> tmpVreg; | 48 | + Reg::RegTensor<float> tmpVreg; |
| 49 | - MicroAPI::RegTensor<float> dstVreg; | 49 | + Reg::RegTensor<float> dstVreg; |
| 50 | - MicroAPI::RegTensor<float> minVreg; | 50 | + Reg::RegTensor<float> minVreg; |
| 51 | - MicroAPI::RegTensor<T> castVreg; | 51 | + Reg::RegTensor<T> castVreg; |
| 52 | - MicroAPI::UnalignReg ureg0, ureg1; | 52 | + Reg::UnalignReg ureg0, ureg1; |
| 53 | NotNumUnion notNum; | 53 | NotNumUnion notNum; |
| 54 | notNum.i = F32_NEG_INF; | 54 | notNum.i = F32_NEG_INF; |
| 55 | 55 | ||
| 56 | - MicroAPI::Duplicate(minVreg, notNum.f); | 56 | + Reg::Duplicate(minVreg, notNum.f); |
| 57 | for (uint16_t i = 0; i < srcM; ++i) { | 57 | for (uint16_t i = 0; i < srcM; ++i) { |
| 58 | for (uint16_t j = 0; j < repeatTime; ++j) { | 58 | for (uint16_t j = 0; j < repeatTime; ++j) { |
| 59 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + j * repeatStride, maskFull); | 59 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + j * repeatStride, maskFull); |
| 60 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, srcVreg, maskFull); | 60 | + Reg::ReduceSumWithDataBlock(sumVreg, srcVreg, maskFull); |
| 61 | - MicroAPI::StoreAlign<float>(workUb + i * repeatStride + j * blockStride, sumVreg, maskOneBlk); | 61 | + Reg::StoreAlign<float>(workUb + i * repeatStride + j * blockStride, sumVreg, maskOneBlk); |
| 62 | } | 62 | } |
| 63 | } | 63 | } |
| 64 | 64 | ||
| 65 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 65 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 66 | 66 | ||
| 67 | for (uint16_t i = 0; i < srcM; ++i) { | 67 | for (uint16_t i = 0; i < srcM; ++i) { |
| 68 | - MicroAPI::LoadAlign<float>(sumVreg, workUb + i * repeatStride); | 68 | + Reg::LoadAlign<float>(sumVreg, workUb + i * repeatStride); |
| 69 | for (uint16_t j = 0; j < remainRepeatTime; ++j) { | 69 | for (uint16_t j = 0; j < remainRepeatTime; ++j) { |
| 70 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + repeatStride * splitMeanCnt + j * repeatStride, maskFull); | 70 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + repeatStride * splitMeanCnt + j * repeatStride, maskFull); |
| 71 | - MicroAPI::Add(sumVreg, srcVreg, sumVreg, maskFull); | 71 | + Reg::Add(sumVreg, srcVreg, sumVreg, maskFull); |
| 72 | } | 72 | } |
| 73 | - MicroAPI::StoreAlign<float>(workUb + i * repeatStride, sumVreg, maskFull); | 73 | + Reg::StoreAlign<float>(workUb + i * repeatStride, sumVreg, maskFull); |
| 74 | } | 74 | } |
| 75 | 75 | ||
| 76 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 76 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 77 | 77 | ||
| 78 | for (uint16_t i = 0; i < srcM; ++i) { | 78 | for (uint16_t i = 0; i < srcM; ++i) { |
| 79 | - MicroAPI::LoadAlign<float>(sumVreg, workUb + i * repeatStride); | 79 | + Reg::LoadAlign<float>(sumVreg, workUb + i * repeatStride); |
| 80 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, sumVreg, maskFull); | 80 | + Reg::ReduceSumWithDataBlock(sumVreg, sumVreg, maskFull); |
| 81 | - MicroAPI::Muls(meanVreg, sumVreg, r0, maskOneBlk); | 81 | + Reg::Muls(meanVreg, sumVreg, r0, maskOneBlk); |
| 82 | 82 | ||
| 83 | - MicroAPI::ReduceSum(tmpVreg, meanVreg, maskOneBlk); | 83 | + Reg::ReduceSum(tmpVreg, meanVreg, maskOneBlk); |
| 84 | - MicroAPI::Muls(tmpVreg, tmpVreg, r1, maskOnePt); | 84 | + Reg::Muls(tmpVreg, tmpVreg, r1, maskOnePt); |
| 85 | - MicroAPI::Duplicate(tmpVreg, tmpVreg, maskOneBlk); | 85 | + Reg::Duplicate(tmpVreg, tmpVreg, maskOneBlk); |
| 86 | StoreIfNeedCast<U>(meanUb + i * blockStride, tmpVreg, maskOneBlk); | 86 | StoreIfNeedCast<U>(meanUb + i * blockStride, tmpVreg, maskOneBlk); |
| 87 | - MicroAPI::Sub(tmpVreg, tmpVreg, meanVreg, maskOneBlk); | 87 | + Reg::Sub(tmpVreg, tmpVreg, meanVreg, maskOneBlk); |
| 88 | - MicroAPI::Muls(tmpVreg, tmpVreg, scalar, maskOneBlk); // scalar = alpha / (1 - alpha) | 88 | + Reg::Muls(tmpVreg, tmpVreg, scalar, maskOneBlk); // scalar = alpha / (1 - alpha) |
| 89 | - MicroAPI::StoreAlign<float>(workUb + i * blockStride, tmpVreg, maskOneBlk); | 89 | + Reg::StoreAlign<float>(workUb + i * blockStride, tmpVreg, maskOneBlk); |
| 90 | } | 90 | } |
| 91 | 91 | ||
| 92 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 92 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 93 | 93 | ||
| 94 | for (uint16_t i = 0; i < srcM; ++i) { | 94 | for (uint16_t i = 0; i < srcM; ++i) { |
| 95 | - MicroAPI::Duplicate(maxVreg, notNum.f); | 95 | + Reg::Duplicate(maxVreg, notNum.f); |
| 96 | for (uint16_t j = 0; j < splitMeanCnt; ++j) { // 8 | 96 | for (uint16_t j = 0; j < splitMeanCnt; ++j) { // 8 |
| 97 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(meanVreg, workUb + i * splitMeanCnt + j); | 97 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_BRC_B32>(meanVreg, workUb + i * splitMeanCnt + j); |
| 98 | uint32_t sreg = baseK; | 98 | uint32_t sreg = baseK; |
| 99 | for (uint16_t k = 0; k < baseKRepeatTime; ++k) { // baseK / 64 | 99 | for (uint16_t k = 0; k < baseKRepeatTime; ++k) { // baseK / 64 |
| 100 | - maskCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 100 | + maskCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 101 | __ubuf__ T *srcUbTmp = srcUb + i * srcK + j * baseK + k * repeatStride; | 101 | __ubuf__ T *srcUbTmp = srcUb + i * srcK + j * baseK + k * repeatStride; |
| 102 | - MicroAPI::LoadUnAlignPre(ureg0, srcUbTmp); | 102 | + Reg::LoadUnAlignPre(ureg0, srcUbTmp); |
| 103 | - MicroAPI::LoadUnAlign(castVreg, ureg0, srcUbTmp, repeatStride); | 103 | + Reg::LoadUnAlign(castVreg, ureg0, srcUbTmp, repeatStride); |
| 104 | - MicroAPI::UnPack<uint32_t, uint16_t>( | 104 | + Reg::UnPack<uint32_t, uint16_t>( |
| 105 | - (MicroAPI::RegTensor<uint32_t>&)castVreg, (MicroAPI::RegTensor<uint16_t>&)castVreg); | 105 | + (Reg::RegTensor<uint32_t>&)castVreg, (Reg::RegTensor<uint16_t>&)castVreg); |
| 106 | - MicroAPI::Cast<float, T, Internal::castTraitB16ToB32>(srcVreg, castVreg, maskCnt); | 106 | + Reg::Cast<float, T, Internal::castTraitB16ToB32>(srcVreg, castVreg, maskCnt); |
| 107 | - MicroAPI::Sub(srcVreg, srcVreg, meanVreg, maskCnt); | 107 | + Reg::Sub(srcVreg, srcVreg, meanVreg, maskCnt); |
| 108 | __ubuf__ float *newSrcUbTmp = newSrcUb + i * srcK + j * baseK + k * repeatStride; | 108 | __ubuf__ float *newSrcUbTmp = newSrcUb + i * srcK + j * baseK + k * repeatStride; |
| 109 | - MicroAPI::StoreUnAlign(newSrcUbTmp, srcVreg, ureg1, repeatStride); | 109 | + Reg::StoreUnAlign(newSrcUbTmp, srcVreg, ureg1, repeatStride); |
| 110 | - MicroAPI::StoreUnAlignPost(newSrcUbTmp, ureg1, 0); | 110 | + Reg::StoreUnAlignPost(newSrcUbTmp, ureg1, 0); |
| 111 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, maskFull); | 111 | + Reg::Max(maxVreg, maxVreg, srcVreg, maskFull); |
| 112 | } | 112 | } |
| 113 | - maskCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 113 | + maskCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 114 | __ubuf__ T *srcUbTmp = srcUb + i * srcK + j * baseK + baseKRepeatTime * repeatStride; | 114 | __ubuf__ T *srcUbTmp = srcUb + i * srcK + j * baseK + baseKRepeatTime * repeatStride; |
| 115 | - MicroAPI::LoadUnAlignPre(ureg0, srcUbTmp); | 115 | + Reg::LoadUnAlignPre(ureg0, srcUbTmp); |
| 116 | - MicroAPI::LoadUnAlign(castVreg, ureg0, srcUbTmp, tail); | 116 | + Reg::LoadUnAlign(castVreg, ureg0, srcUbTmp, tail); |
| 117 | - MicroAPI::UnPack<uint32_t, uint16_t>( | 117 | + Reg::UnPack<uint32_t, uint16_t>( |
| 118 | - (MicroAPI::RegTensor<uint32_t>&)castVreg, (MicroAPI::RegTensor<uint16_t>&)castVreg); | 118 | + (Reg::RegTensor<uint32_t>&)castVreg, (Reg::RegTensor<uint16_t>&)castVreg); |
| 119 | - MicroAPI::Cast<float, T, Internal::castTraitB16ToB32>(srcVreg, castVreg, maskCnt); | 119 | + Reg::Cast<float, T, Internal::castTraitB16ToB32>(srcVreg, castVreg, maskCnt); |
| 120 | - MicroAPI::Sub(srcVreg, srcVreg, meanVreg, maskCnt); | 120 | + Reg::Sub(srcVreg, srcVreg, meanVreg, maskCnt); |
| 121 | __ubuf__ float *newSrcUbTmp = newSrcUb + i * srcK + j * baseK + baseKRepeatTime * repeatStride; | 121 | __ubuf__ float *newSrcUbTmp = newSrcUb + i * srcK + j * baseK + baseKRepeatTime * repeatStride; |
| 122 | - MicroAPI::StoreUnAlign(newSrcUbTmp, srcVreg, ureg1, tail); | 122 | + Reg::StoreUnAlign(newSrcUbTmp, srcVreg, ureg1, tail); |
| 123 | - MicroAPI::StoreUnAlignPost(newSrcUbTmp, ureg1, 0); | 123 | + Reg::StoreUnAlignPost(newSrcUbTmp, ureg1, 0); |
| 124 | - MicroAPI::Select(srcVreg, srcVreg, minVreg, maskCnt); | 124 | + Reg::Select(srcVreg, srcVreg, minVreg, maskCnt); |
| 125 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, maskFull); | 125 | + Reg::Max(maxVreg, maxVreg, srcVreg, maskFull); |
| 126 | } | 126 | } |
| 127 | - MicroAPI::ReduceMax(maxVreg, maxVreg, maskFull); | 127 | + Reg::ReduceMax(maxVreg, maxVreg, maskFull); |
| 128 | - MicroAPI::Duplicate(maxVreg, maxVreg, maskOneBlk); | 128 | + Reg::Duplicate(maxVreg, maxVreg, maskOneBlk); |
| 129 | StoreIfNeedCast<U>(maxUb + i * blockStride, maxVreg, maskOneBlk); | 129 | StoreIfNeedCast<U>(maxUb + i * blockStride, maxVreg, maskOneBlk); |
| 130 | } | 130 | } |
| 131 | 131 | ||
| 132 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 132 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 133 | 133 | ||
| 134 | for (uint16_t i = 0; i < srcM; ++i) { | 134 | for (uint16_t i = 0; i < srcM; ++i) { |
| 135 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(maxVreg, maxUb + i * blockStride); | 135 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_BRC_B32>(maxVreg, maxUb + i * blockStride); |
| 136 | - MicroAPI::Duplicate(sumVreg, 0); | 136 | + Reg::Duplicate(sumVreg, 0); |
| 137 | for (uint16_t k = 0; k < kRepeatTime; ++k) { // k / 64 | 137 | for (uint16_t k = 0; k < kRepeatTime; ++k) { // k / 64 |
| 138 | - MicroAPI::LoadAlign<float>(srcVreg, newSrcUb + i * srcK + k * repeatStride); | 138 | + Reg::LoadAlign<float>(srcVreg, newSrcUb + i * srcK + k * repeatStride); |
| 139 | - MicroAPI::FusedExpSub(dstVreg, srcVreg, maxVreg, maskFull); | 139 | + Reg::FusedExpSub(dstVreg, srcVreg, maxVreg, maskFull); |
| 140 | StoreIfNeedCast<T>(dstUb + i * srcK + k * repeatStride, dstVreg, maskFull); | 140 | StoreIfNeedCast<T>(dstUb + i * srcK + k * repeatStride, dstVreg, maskFull); |
| 141 | - MicroAPI::Add(sumVreg, sumVreg, dstVreg, maskFull); | 141 | + Reg::Add(sumVreg, sumVreg, dstVreg, maskFull); |
| 142 | } | 142 | } |
| 143 | - MicroAPI::ReduceSum(sumVreg, sumVreg, maskFull); | 143 | + Reg::ReduceSum(sumVreg, sumVreg, maskFull); |
| 144 | - MicroAPI::Duplicate(sumVreg, sumVreg, maskOneBlk); | 144 | + Reg::Duplicate(sumVreg, sumVreg, maskOneBlk); |
| 145 | StoreIfNeedCast<U>(expSumUb + i * blockStride, sumVreg, maskOneBlk); | 145 | StoreIfNeedCast<U>(expSumUb + i * blockStride, sumVreg, maskOneBlk); |
| 146 | } | 146 | } |
| 147 | } | 147 | } |
| @@ -161,140 +161,140 @@ __simd_vf__ __aicore__ inline void SoftmaxFlashV3NDUpdateImpl(__ubuf__ T* dstUb, | |||
| 161 | constexpr uint32_t blockStride = GetDataBlockSizeInBytes() / sizeof(U); | 161 | constexpr uint32_t blockStride = GetDataBlockSizeInBytes() / sizeof(U); |
| 162 | constexpr uint16_t repeatTime = static_cast<uint16_t>(repeatStride / blockStride); | 162 | constexpr uint16_t repeatTime = static_cast<uint16_t>(repeatStride / blockStride); |
| 163 | 163 | ||
| 164 | - MicroAPI::MaskReg maskCnt; | 164 | + Reg::MaskReg maskCnt; |
| 165 | - MicroAPI::MaskReg maskFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 165 | + Reg::MaskReg maskFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 166 | - MicroAPI::MaskReg maskOnePt = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL1>(); | 166 | + Reg::MaskReg maskOnePt = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL1>(); |
| 167 | - MicroAPI::MaskReg maskOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 167 | + Reg::MaskReg maskOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 168 | - MicroAPI::MaskReg maskOut = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL16>(); | 168 | + Reg::MaskReg maskOut = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL16>(); |
| 169 | - MicroAPI::RegTensor<float> srcVreg; | 169 | + Reg::RegTensor<float> srcVreg; |
| 170 | - MicroAPI::RegTensor<float> maxVreg; | 170 | + Reg::RegTensor<float> maxVreg; |
| 171 | - MicroAPI::RegTensor<float> sumVreg; | 171 | + Reg::RegTensor<float> sumVreg; |
| 172 | - MicroAPI::RegTensor<float> meanVreg; | 172 | + Reg::RegTensor<float> meanVreg; |
| 173 | - MicroAPI::RegTensor<float> inputVreg; | 173 | + Reg::RegTensor<float> inputVreg; |
| 174 | - MicroAPI::RegTensor<float> shiftVreg; | 174 | + Reg::RegTensor<float> shiftVreg; |
| 175 | - MicroAPI::RegTensor<float> tmpVreg; | 175 | + Reg::RegTensor<float> tmpVreg; |
| 176 | - MicroAPI::RegTensor<float> dstVreg; | 176 | + Reg::RegTensor<float> dstVreg; |
| 177 | - MicroAPI::RegTensor<float> minVreg; | 177 | + Reg::RegTensor<float> minVreg; |
| 178 | - MicroAPI::RegTensor<T> castVreg; | 178 | + Reg::RegTensor<T> castVreg; |
| 179 | - MicroAPI::UnalignReg ureg0, ureg1; | 179 | + Reg::UnalignReg ureg0, ureg1; |
| 180 | NotNumUnion notNum; | 180 | NotNumUnion notNum; |
| 181 | notNum.i = F32_NEG_INF; | 181 | notNum.i = F32_NEG_INF; |
| 182 | 182 | ||
| 183 | - MicroAPI::Duplicate(minVreg, notNum.f); | 183 | + Reg::Duplicate(minVreg, notNum.f); |
| 184 | for (uint16_t i = 0; i < srcM; ++i) { | 184 | for (uint16_t i = 0; i < srcM; ++i) { |
| 185 | for (uint16_t j = 0; j < repeatTime; ++j) { | 185 | for (uint16_t j = 0; j < repeatTime; ++j) { |
| 186 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + j * repeatStride, maskFull); | 186 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + j * repeatStride, maskFull); |
| 187 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, srcVreg, maskFull); | 187 | + Reg::ReduceSumWithDataBlock(sumVreg, srcVreg, maskFull); |
| 188 | - MicroAPI::StoreAlign<float>(workUb + i * repeatStride + j * blockStride, sumVreg, maskOneBlk); | 188 | + Reg::StoreAlign<float>(workUb + i * repeatStride + j * blockStride, sumVreg, maskOneBlk); |
| 189 | } | 189 | } |
| 190 | } | 190 | } |
| 191 | 191 | ||
| 192 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 192 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 193 | 193 | ||
| 194 | for (uint16_t i = 0; i < srcM; ++i) { | 194 | for (uint16_t i = 0; i < srcM; ++i) { |
| 195 | - MicroAPI::LoadAlign<float>(sumVreg, workUb + i * repeatStride); | 195 | + Reg::LoadAlign<float>(sumVreg, workUb + i * repeatStride); |
| 196 | for (uint16_t j = 0; j < remainRepeatTime; ++j) { | 196 | for (uint16_t j = 0; j < remainRepeatTime; ++j) { |
| 197 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + repeatStride * splitMeanCnt + j * repeatStride, maskFull); | 197 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + repeatStride * splitMeanCnt + j * repeatStride, maskFull); |
| 198 | - MicroAPI::Add(sumVreg, srcVreg, sumVreg, maskFull); | 198 | + Reg::Add(sumVreg, srcVreg, sumVreg, maskFull); |
| 199 | } | 199 | } |
| 200 | - MicroAPI::StoreAlign<float>(workUb + i * repeatStride, sumVreg, maskFull); | 200 | + Reg::StoreAlign<float>(workUb + i * repeatStride, sumVreg, maskFull); |
| 201 | } | 201 | } |
| 202 | 202 | ||
| 203 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 203 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 204 | 204 | ||
| 205 | for (uint16_t i = 0; i < srcM; ++i) { | 205 | for (uint16_t i = 0; i < srcM; ++i) { |
| 206 | - MicroAPI::LoadAlign<float>(sumVreg, workUb + i * repeatStride); | 206 | + Reg::LoadAlign<float>(sumVreg, workUb + i * repeatStride); |
| 207 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, sumVreg, maskFull); | 207 | + Reg::ReduceSumWithDataBlock(sumVreg, sumVreg, maskFull); |
| 208 | - MicroAPI::Muls(meanVreg, sumVreg, r0, maskOneBlk); | 208 | + Reg::Muls(meanVreg, sumVreg, r0, maskOneBlk); |
| 209 | 209 | ||
| 210 | - MicroAPI::ReduceSum(tmpVreg, meanVreg, maskOneBlk); | 210 | + Reg::ReduceSum(tmpVreg, meanVreg, maskOneBlk); |
| 211 | - MicroAPI::Muls(tmpVreg, tmpVreg, r1, maskOnePt); | 211 | + Reg::Muls(tmpVreg, tmpVreg, r1, maskOnePt); |
| 212 | - MicroAPI::Duplicate(tmpVreg, tmpVreg, maskOneBlk); | 212 | + Reg::Duplicate(tmpVreg, tmpVreg, maskOneBlk); |
| 213 | - MicroAPI::StoreAlign<float>(tmpUb + i * blockStride, tmpVreg, maskOneBlk); | 213 | + Reg::StoreAlign<float>(tmpUb + i * blockStride, tmpVreg, maskOneBlk); |
| 214 | - MicroAPI::Sub(tmpVreg, tmpVreg, meanVreg, maskOneBlk); | 214 | + Reg::Sub(tmpVreg, tmpVreg, meanVreg, maskOneBlk); |
| 215 | - MicroAPI::Muls(tmpVreg, tmpVreg, scalar, maskOneBlk); // scalar = alpha / (1 - alpha) | 215 | + Reg::Muls(tmpVreg, tmpVreg, scalar, maskOneBlk); // scalar = alpha / (1 - alpha) |
| 216 | - MicroAPI::StoreAlign<float>(workUb + i * blockStride, tmpVreg, maskOneBlk); | 216 | + Reg::StoreAlign<float>(workUb + i * blockStride, tmpVreg, maskOneBlk); |
| 217 | } | 217 | } |
| 218 | 218 | ||
| 219 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 219 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 220 | 220 | ||
| 221 | for (uint16_t i = 0; i < srcM; ++i) { | 221 | for (uint16_t i = 0; i < srcM; ++i) { |
| 222 | LoadIfNeedCast<U>(inputVreg, inMeanUb + i * blockStride, maskOneBlk); | 222 | LoadIfNeedCast<U>(inputVreg, inMeanUb + i * blockStride, maskOneBlk); |
| 223 | - MicroAPI::LoadAlign<float>(tmpVreg, tmpUb + i * blockStride); | 223 | + Reg::LoadAlign<float>(tmpVreg, tmpUb + i * blockStride); |
| 224 | - MicroAPI::Muls(shiftVreg, inputVreg, r2, maskOneBlk); | 224 | + Reg::Muls(shiftVreg, inputVreg, r2, maskOneBlk); |
| 225 | - MicroAPI::Add(shiftVreg, shiftVreg, tmpVreg, maskOneBlk); | 225 | + Reg::Add(shiftVreg, shiftVreg, tmpVreg, maskOneBlk); |
| 226 | - MicroAPI::Muls(shiftVreg, shiftVreg, r3, maskOneBlk); | 226 | + Reg::Muls(shiftVreg, shiftVreg, r3, maskOneBlk); |
| 227 | StoreIfNeedCast<U>(meanUb + i * blockStride, shiftVreg, maskOneBlk); | 227 | StoreIfNeedCast<U>(meanUb + i * blockStride, shiftVreg, maskOneBlk); |
| 228 | 228 | ||
| 229 | - MicroAPI::Duplicate(maxVreg, notNum.f); | 229 | + Reg::Duplicate(maxVreg, notNum.f); |
| 230 | for (uint16_t j = 0; j < splitMeanCnt; ++j) { // 8 | 230 | for (uint16_t j = 0; j < splitMeanCnt; ++j) { // 8 |
| 231 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(meanVreg, workUb + i * splitMeanCnt + j); | 231 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_BRC_B32>(meanVreg, workUb + i * splitMeanCnt + j); |
| 232 | uint32_t sreg = baseK; | 232 | uint32_t sreg = baseK; |
| 233 | for (uint16_t k = 0; k < baseKRepeatTime; ++k) { // baseK / 64 | 233 | for (uint16_t k = 0; k < baseKRepeatTime; ++k) { // baseK / 64 |
| 234 | - maskCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 234 | + maskCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 235 | __ubuf__ T *srcUbTmp = srcUb + i * srcK + j * baseK + k * repeatStride; | 235 | __ubuf__ T *srcUbTmp = srcUb + i * srcK + j * baseK + k * repeatStride; |
| 236 | - MicroAPI::LoadUnAlignPre(ureg0, srcUbTmp); | 236 | + Reg::LoadUnAlignPre(ureg0, srcUbTmp); |
| 237 | - MicroAPI::LoadUnAlign(castVreg, ureg0, srcUbTmp, repeatStride); | 237 | + Reg::LoadUnAlign(castVreg, ureg0, srcUbTmp, repeatStride); |
| 238 | - MicroAPI::UnPack<uint32_t, uint16_t>( | 238 | + Reg::UnPack<uint32_t, uint16_t>( |
| 239 | - (MicroAPI::RegTensor<uint32_t>&)castVreg, (MicroAPI::RegTensor<uint16_t>&)castVreg); | 239 | + (Reg::RegTensor<uint32_t>&)castVreg, (Reg::RegTensor<uint16_t>&)castVreg); |
| 240 | - MicroAPI::Cast<float, T, Internal::castTraitB16ToB32>(srcVreg, castVreg, maskCnt); | 240 | + Reg::Cast<float, T, Internal::castTraitB16ToB32>(srcVreg, castVreg, maskCnt); |
| 241 | - MicroAPI::Sub(srcVreg, srcVreg, meanVreg, maskCnt); | 241 | + Reg::Sub(srcVreg, srcVreg, meanVreg, maskCnt); |
| 242 | __ubuf__ float *newSrcUbTmp = newSrcUb + i * srcK + j * baseK + k * repeatStride; | 242 | __ubuf__ float *newSrcUbTmp = newSrcUb + i * srcK + j * baseK + k * repeatStride; |
| 243 | - MicroAPI::StoreUnAlign(newSrcUbTmp, srcVreg, ureg1, repeatStride); | 243 | + Reg::StoreUnAlign(newSrcUbTmp, srcVreg, ureg1, repeatStride); |
| 244 | - MicroAPI::StoreUnAlignPost(newSrcUbTmp, ureg1, 0); | 244 | + Reg::StoreUnAlignPost(newSrcUbTmp, ureg1, 0); |
| 245 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, maskFull); | 245 | + Reg::Max(maxVreg, maxVreg, srcVreg, maskFull); |
| 246 | } | 246 | } |
| 247 | - maskCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 247 | + maskCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 248 | __ubuf__ T *srcUbTmp = srcUb + i * srcK + j * baseK + baseKRepeatTime * repeatStride; | 248 | __ubuf__ T *srcUbTmp = srcUb + i * srcK + j * baseK + baseKRepeatTime * repeatStride; |
| 249 | - MicroAPI::LoadUnAlignPre(ureg0, srcUbTmp); | 249 | + Reg::LoadUnAlignPre(ureg0, srcUbTmp); |
| 250 | - MicroAPI::LoadUnAlign(castVreg, ureg0, srcUbTmp, tail); | 250 | + Reg::LoadUnAlign(castVreg, ureg0, srcUbTmp, tail); |
| 251 | - MicroAPI::UnPack<uint32_t, uint16_t>( | 251 | + Reg::UnPack<uint32_t, uint16_t>( |
| 252 | - (MicroAPI::RegTensor<uint32_t>&)castVreg, (MicroAPI::RegTensor<uint16_t>&)castVreg); | 252 | + (Reg::RegTensor<uint32_t>&)castVreg, (Reg::RegTensor<uint16_t>&)castVreg); |
| 253 | - MicroAPI::Cast<float, T, Internal::castTraitB16ToB32>(srcVreg, castVreg, maskCnt); | 253 | + Reg::Cast<float, T, Internal::castTraitB16ToB32>(srcVreg, castVreg, maskCnt); |
| 254 | - MicroAPI::Sub(srcVreg, srcVreg, meanVreg, maskCnt); | 254 | + Reg::Sub(srcVreg, srcVreg, meanVreg, maskCnt); |
| 255 | __ubuf__ float *newSrcUbTmp = newSrcUb + i * srcK + j * baseK + baseKRepeatTime * repeatStride; | 255 | __ubuf__ float *newSrcUbTmp = newSrcUb + i * srcK + j * baseK + baseKRepeatTime * repeatStride; |
| 256 | - MicroAPI::StoreUnAlign(newSrcUbTmp, srcVreg, ureg1, tail); | 256 | + Reg::StoreUnAlign(newSrcUbTmp, srcVreg, ureg1, tail); |
| 257 | - MicroAPI::StoreUnAlignPost(newSrcUbTmp, ureg1, 0); | 257 | + Reg::StoreUnAlignPost(newSrcUbTmp, ureg1, 0); |
| 258 | - MicroAPI::Select(srcVreg, srcVreg, minVreg, maskCnt); | 258 | + Reg::Select(srcVreg, srcVreg, minVreg, maskCnt); |
| 259 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, maskFull); | 259 | + Reg::Max(maxVreg, maxVreg, srcVreg, maskFull); |
| 260 | } | 260 | } |
| 261 | - MicroAPI::ReduceMax(maxVreg, maxVreg, maskFull); | 261 | + Reg::ReduceMax(maxVreg, maxVreg, maskFull); |
| 262 | - MicroAPI::Duplicate(maxVreg, maxVreg, maskOneBlk); | 262 | + Reg::Duplicate(maxVreg, maxVreg, maskOneBlk); |
| 263 | 263 | ||
| 264 | - MicroAPI::Sub(dstVreg, tmpVreg, shiftVreg, maskOneBlk); | 264 | + Reg::Sub(dstVreg, tmpVreg, shiftVreg, maskOneBlk); |
| 265 | - MicroAPI::Muls(dstVreg, dstVreg, scalar, maskOneBlk); // scalar = alpha / (1 - alpha) | 265 | + Reg::Muls(dstVreg, dstVreg, scalar, maskOneBlk); // scalar = alpha / (1 - alpha) |
| 266 | - MicroAPI::Sub(tmpVreg, inputVreg, shiftVreg, maskOneBlk); | 266 | + Reg::Sub(tmpVreg, inputVreg, shiftVreg, maskOneBlk); |
| 267 | - MicroAPI::Muls(tmpVreg, tmpVreg, scalar, maskOneBlk); // scalar = alpha / (1 - alpha) | 267 | + Reg::Muls(tmpVreg, tmpVreg, scalar, maskOneBlk); // scalar = alpha / (1 - alpha) |
| 268 | - MicroAPI::Add(maxVreg, dstVreg, maxVreg, maskOneBlk); | 268 | + Reg::Add(maxVreg, dstVreg, maxVreg, maskOneBlk); |
| 269 | LoadIfNeedCast<U>(inputVreg, inMaxUb + i * blockStride, maskOneBlk); | 269 | LoadIfNeedCast<U>(inputVreg, inMaxUb + i * blockStride, maskOneBlk); |
| 270 | - MicroAPI::Add(tmpVreg, inputVreg, tmpVreg, maskOneBlk); | 270 | + Reg::Add(tmpVreg, inputVreg, tmpVreg, maskOneBlk); |
| 271 | - MicroAPI::Max(maxVreg, tmpVreg, maxVreg, maskOneBlk); | 271 | + Reg::Max(maxVreg, tmpVreg, maxVreg, maskOneBlk); |
| 272 | StoreIfNeedCast<U>(maxUb + i * blockStride, maxVreg, maskOneBlk); | 272 | StoreIfNeedCast<U>(maxUb + i * blockStride, maxVreg, maskOneBlk); |
| 273 | - MicroAPI::Sub(maxVreg, maxVreg, dstVreg, maskOneBlk); | 273 | + Reg::Sub(maxVreg, maxVreg, dstVreg, maskOneBlk); |
| 274 | - MicroAPI::StoreAlign<float>(tmpUb + i * blockStride, maxVreg, maskOneBlk); | 274 | + Reg::StoreAlign<float>(tmpUb + i * blockStride, maxVreg, maskOneBlk); |
| 275 | - MicroAPI::FusedExpSub(tmpVreg, tmpVreg, maxVreg, maskFull); | 275 | + Reg::FusedExpSub(tmpVreg, tmpVreg, maxVreg, maskFull); |
| 276 | LoadIfNeedCast<U>(inputVreg, inExpSumUb + i * blockStride, maskOneBlk); | 276 | LoadIfNeedCast<U>(inputVreg, inExpSumUb + i * blockStride, maskOneBlk); |
| 277 | - MicroAPI::Mul(sumVreg, tmpVreg, inputVreg, maskOneBlk); | 277 | + Reg::Mul(sumVreg, tmpVreg, inputVreg, maskOneBlk); |
| 278 | - MicroAPI::StoreAlign<float>(expSumUb + i * blockStride, sumVreg, maskOneBlk); | 278 | + Reg::StoreAlign<float>(expSumUb + i * blockStride, sumVreg, maskOneBlk); |
| 279 | - MicroAPI::Interleave(tmpVreg, dstVreg, tmpVreg, tmpVreg); | 279 | + Reg::Interleave(tmpVreg, dstVreg, tmpVreg, tmpVreg); |
| 280 | StoreIfNeedCast<T>(expMaxUb + i * blockStride * 2, tmpVreg, maskOut); | 280 | StoreIfNeedCast<T>(expMaxUb + i * blockStride * 2, tmpVreg, maskOut); |
| 281 | } | 281 | } |
| 282 | 282 | ||
| 283 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 283 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 284 | 284 | ||
| 285 | for (uint16_t i = 0; i < srcM; ++i) { | 285 | for (uint16_t i = 0; i < srcM; ++i) { |
| 286 | - MicroAPI::Duplicate(sumVreg, 0); | 286 | + Reg::Duplicate(sumVreg, 0); |
| 287 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(maxVreg, tmpUb + i * blockStride); | 287 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_BRC_B32>(maxVreg, tmpUb + i * blockStride); |
| 288 | for (uint16_t k = 0; k < kRepeatTime; ++k) { // k / 64 | 288 | for (uint16_t k = 0; k < kRepeatTime; ++k) { // k / 64 |
| 289 | - MicroAPI::LoadAlign<float>(srcVreg, newSrcUb + i * srcK + k * repeatStride); | 289 | + Reg::LoadAlign<float>(srcVreg, newSrcUb + i * srcK + k * repeatStride); |
| 290 | - MicroAPI::FusedExpSub(dstVreg, srcVreg, maxVreg, maskFull); | 290 | + Reg::FusedExpSub(dstVreg, srcVreg, maxVreg, maskFull); |
| 291 | StoreIfNeedCast<T>(dstUb + i * srcK + k * repeatStride, dstVreg, maskFull); | 291 | StoreIfNeedCast<T>(dstUb + i * srcK + k * repeatStride, dstVreg, maskFull); |
| 292 | - MicroAPI::Add(sumVreg, sumVreg, dstVreg, maskFull); | 292 | + Reg::Add(sumVreg, sumVreg, dstVreg, maskFull); |
| 293 | } | 293 | } |
| 294 | - MicroAPI::ReduceSum(sumVreg, sumVreg, maskFull); | 294 | + Reg::ReduceSum(sumVreg, sumVreg, maskFull); |
| 295 | - MicroAPI::Duplicate(sumVreg, sumVreg, maskOneBlk); | 295 | + Reg::Duplicate(sumVreg, sumVreg, maskOneBlk); |
| 296 | - MicroAPI::LoadAlign<float>(tmpVreg, expSumUb + i * blockStride); | 296 | + Reg::LoadAlign<float>(tmpVreg, expSumUb + i * blockStride); |
| 297 | - MicroAPI::Add(sumVreg, tmpVreg, sumVreg, maskOneBlk); | 297 | + Reg::Add(sumVreg, tmpVreg, sumVreg, maskOneBlk); |
| 298 | StoreIfNeedCast<U>(expSumUb + i * blockStride, sumVreg, maskOneBlk); | 298 | StoreIfNeedCast<U>(expSumUb + i * blockStride, sumVreg, maskOneBlk); |
| 299 | } | 299 | } |
| 300 | } | 300 | } |
| @@ -41,79 +41,79 @@ __simd_vf__ inline void SoftmaxGradGenericNZWithTailVFImpl(__ubuf__ T *dstUb, __ | |||
| 41 | uint16_t VcgFoldRepeat = (dataNumAfterVcg + HALF_REPEAT_SIZE - 1) / HALF_REPEAT_SIZE; | 41 | uint16_t VcgFoldRepeat = (dataNumAfterVcg + HALF_REPEAT_SIZE - 1) / HALF_REPEAT_SIZE; |
| 42 | uint16_t e2bRep = srcM / DEFAULT_BLK_NUM; | 42 | uint16_t e2bRep = srcM / DEFAULT_BLK_NUM; |
| 43 | 43 | ||
| 44 | - MicroAPI::MaskReg pregCnt; | 44 | + Reg::MaskReg pregCnt; |
| 45 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 45 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 46 | - MicroAPI::MaskReg pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 46 | + Reg::MaskReg pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 47 | - MicroAPI::MaskReg pregkTail = MicroAPI::MoveMask<uint32_t>(); | 47 | + Reg::MaskReg pregkTail = Reg::MoveMask<uint32_t>(); |
| 48 | - MicroAPI::RegTensor<float> srcVreg; | 48 | + Reg::RegTensor<float> srcVreg; |
| 49 | - MicroAPI::RegTensor<float> gradVreg; | 49 | + Reg::RegTensor<float> gradVreg; |
| 50 | - MicroAPI::RegTensor<float> sumVreg; | 50 | + Reg::RegTensor<float> sumVreg; |
| 51 | - MicroAPI::RegTensor<float> tmpVreg; | 51 | + Reg::RegTensor<float> tmpVreg; |
| 52 | - MicroAPI::RegTensor<T> castReg; | 52 | + Reg::RegTensor<T> castReg; |
| 53 | 53 | ||
| 54 | for (uint16_t i = 0; i < mRepeatInner; ++i) { | 54 | for (uint16_t i = 0; i < mRepeatInner; ++i) { |
| 55 | Duplicate(sumVreg, 0); | 55 | Duplicate(sumVreg, 0); |
| 56 | for (uint16_t j = 0; j < kOuter; ++j) { | 56 | for (uint16_t j = 0; j < kOuter; ++j) { |
| 57 | LoadIfNeedCast<T>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j *dataPadInner, pregFull); | 57 | LoadIfNeedCast<T>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j *dataPadInner, pregFull); |
| 58 | LoadIfNeedCast<T>(gradVreg, gradUb + i * FLOAT_REPEAT_SIZE + j *dataPadInner, pregFull); | 58 | LoadIfNeedCast<T>(gradVreg, gradUb + i * FLOAT_REPEAT_SIZE + j *dataPadInner, pregFull); |
| 59 | - MicroAPI::Mul(tmpVreg, gradVreg, srcVreg, pregFull); | 59 | + Reg::Mul(tmpVreg, gradVreg, srcVreg, pregFull); |
| 60 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 60 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 61 | } | 61 | } |
| 62 | uint32_t tailOffset = i * FLOAT_REPEAT_SIZE + kOuter * dataPadInner; | 62 | uint32_t tailOffset = i * FLOAT_REPEAT_SIZE + kOuter * dataPadInner; |
| 63 | LoadIfNeedCast<T>(srcVreg, srcUb + tailOffset, pregFull); | 63 | LoadIfNeedCast<T>(srcVreg, srcUb + tailOffset, pregFull); |
| 64 | LoadIfNeedCast<T>(gradVreg, gradUb + tailOffset, pregFull); | 64 | LoadIfNeedCast<T>(gradVreg, gradUb + tailOffset, pregFull); |
| 65 | - MicroAPI::Mul(tmpVreg, gradVreg, srcVreg, pregkTail); | 65 | + Reg::Mul(tmpVreg, gradVreg, srcVreg, pregkTail); |
| 66 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 66 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 67 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); | 67 | + Reg::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); |
| 68 | - MicroAPI::StoreAlign(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); | 68 | + Reg::StoreAlign(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); |
| 69 | } | 69 | } |
| 70 | 70 | ||
| 71 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 71 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 72 | 72 | ||
| 73 | for (uint16_t i = 0; i < VcgFoldRepeat; i++) { | 73 | for (uint16_t i = 0; i < VcgFoldRepeat; i++) { |
| 74 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 74 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 75 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 75 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 76 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 76 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 77 | if constexpr (isFront) { | 77 | if constexpr (isFront) { |
| 78 | StoreIfNeedCast<T>(workUbFront + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); | 78 | StoreIfNeedCast<T>(workUbFront + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); |
| 79 | } else { | 79 | } else { |
| 80 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 80 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 81 | workUb + i * HALF_REPEAT_SIZE, sumVreg, sumVreg, pregFull); | 81 | workUb + i * HALF_REPEAT_SIZE, sumVreg, sumVreg, pregFull); |
| 82 | } | 82 | } |
| 83 | } | 83 | } |
| 84 | 84 | ||
| 85 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 85 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 86 | 86 | ||
| 87 | if constexpr (isFront) { | 87 | if constexpr (isFront) { |
| 88 | uint32_t sreg = oriM * dtypeBlkStride; | 88 | uint32_t sreg = oriM * dtypeBlkStride; |
| 89 | for (uint16_t i = 0; i < e2bRep; ++i) { | 89 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 90 | - pregCnt = MicroAPI::UpdateMask<T>(sreg); | 90 | + pregCnt = Reg::UpdateMask<T>(sreg); |
| 91 | LoadE2B<T>(castReg, workUbFront + i * DEFAULT_BLK_NUM); | 91 | LoadE2B<T>(castReg, workUbFront + i * DEFAULT_BLK_NUM); |
| 92 | - MicroAPI::StoreAlign(dstUb + i * dtypeRepStride, castReg, pregCnt); | 92 | + Reg::StoreAlign(dstUb + i * dtypeRepStride, castReg, pregCnt); |
| 93 | } | 93 | } |
| 94 | } else { | 94 | } else { |
| 95 | for (uint16_t j = 0; j < kOuter; ++j) { | 95 | for (uint16_t j = 0; j < kOuter; ++j) { |
| 96 | uint32_t sreg = oriM * B16_DATA_NUM_PER_BLOCK; | 96 | uint32_t sreg = oriM * B16_DATA_NUM_PER_BLOCK; |
| 97 | for (uint16_t i = 0; i < mRepeatInner; ++i) { | 97 | for (uint16_t i = 0; i < mRepeatInner; ++i) { |
| 98 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 98 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 99 | LoadIfNeedCast<T>(srcVreg, srcUb + j * dataPadInner + i * FLOAT_REPEAT_SIZE, pregFull); | 99 | LoadIfNeedCast<T>(srcVreg, srcUb + j * dataPadInner + i * FLOAT_REPEAT_SIZE, pregFull); |
| 100 | LoadIfNeedCast<T>(gradVreg, gradUb + j * dataPadInner + i * FLOAT_REPEAT_SIZE, pregFull); | 100 | LoadIfNeedCast<T>(gradVreg, gradUb + j * dataPadInner + i * FLOAT_REPEAT_SIZE, pregFull); |
| 101 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); | 101 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); |
| 102 | - MicroAPI::Sub(tmpVreg, gradVreg, sumVreg, pregCnt); | 102 | + Reg::Sub(tmpVreg, gradVreg, sumVreg, pregCnt); |
| 103 | - MicroAPI::Mul(tmpVreg, srcVreg, tmpVreg, pregCnt); | 103 | + Reg::Mul(tmpVreg, srcVreg, tmpVreg, pregCnt); |
| 104 | StoreIfNeedCast<T>(dstUb + j * dataPadInner + i * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); | 104 | StoreIfNeedCast<T>(dstUb + j * dataPadInner + i * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); |
| 105 | } | 105 | } |
| 106 | } | 106 | } |
| 107 | uint32_t sreg = oriM * B16_DATA_NUM_PER_BLOCK; | 107 | uint32_t sreg = oriM * B16_DATA_NUM_PER_BLOCK; |
| 108 | for (uint16_t i = 0; i < mRepeatInner; ++i) { | 108 | for (uint16_t i = 0; i < mRepeatInner; ++i) { |
| 109 | uint32_t tailOffset = i * FLOAT_REPEAT_SIZE + kOuter * dataPadInner; | 109 | uint32_t tailOffset = i * FLOAT_REPEAT_SIZE + kOuter * dataPadInner; |
| 110 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 110 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 111 | - MicroAPI::MaskAnd(pregOneBlk, pregCnt, pregkTail, pregFull); | 111 | + Reg::MaskAnd(pregOneBlk, pregCnt, pregkTail, pregFull); |
| 112 | LoadIfNeedCast<T>(srcVreg, srcUb + tailOffset, pregFull); | 112 | LoadIfNeedCast<T>(srcVreg, srcUb + tailOffset, pregFull); |
| 113 | LoadIfNeedCast<T>(gradVreg, gradUb + tailOffset, pregFull); | 113 | LoadIfNeedCast<T>(gradVreg, gradUb + tailOffset, pregFull); |
| 114 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); | 114 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); |
| 115 | - MicroAPI::Sub(tmpVreg, gradVreg, sumVreg, pregOneBlk); | 115 | + Reg::Sub(tmpVreg, gradVreg, sumVreg, pregOneBlk); |
| 116 | - MicroAPI::Mul(tmpVreg, srcVreg, tmpVreg, pregOneBlk); | 116 | + Reg::Mul(tmpVreg, srcVreg, tmpVreg, pregOneBlk); |
| 117 | StoreIfNeedCast<T>(dstUb + tailOffset, tmpVreg, pregOneBlk); | 117 | StoreIfNeedCast<T>(dstUb + tailOffset, tmpVreg, pregOneBlk); |
| 118 | } | 118 | } |
| 119 | } | 119 | } |
| @@ -158,60 +158,60 @@ __simd_vf__ inline void SoftmaxGradGenericNZVFImpl(__ubuf__ T *dstUb, __ubuf__ T | |||
| 158 | uint16_t VcgFoldRepeat = (dataNumAfterVcg + HALF_REPEAT_SIZE - 1) / HALF_REPEAT_SIZE; | 158 | uint16_t VcgFoldRepeat = (dataNumAfterVcg + HALF_REPEAT_SIZE - 1) / HALF_REPEAT_SIZE; |
| 159 | uint16_t e2bRep = tiling.srcM / DEFAULT_BLK_NUM; | 159 | uint16_t e2bRep = tiling.srcM / DEFAULT_BLK_NUM; |
| 160 | 160 | ||
| 161 | - MicroAPI::MaskReg pregCnt; | 161 | + Reg::MaskReg pregCnt; |
| 162 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 162 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 163 | - MicroAPI::MaskReg pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 163 | + Reg::MaskReg pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 164 | - MicroAPI::RegTensor<float> srcVreg; | 164 | + Reg::RegTensor<float> srcVreg; |
| 165 | - MicroAPI::RegTensor<float> gradVreg; | 165 | + Reg::RegTensor<float> gradVreg; |
| 166 | - MicroAPI::RegTensor<float> sumVreg; | 166 | + Reg::RegTensor<float> sumVreg; |
| 167 | - MicroAPI::RegTensor<float> tmpVreg; | 167 | + Reg::RegTensor<float> tmpVreg; |
| 168 | - MicroAPI::RegTensor<T> castReg; | 168 | + Reg::RegTensor<T> castReg; |
| 169 | 169 | ||
| 170 | for (uint16_t i = 0; i < kRepeatInner; ++i) { | 170 | for (uint16_t i = 0; i < kRepeatInner; ++i) { |
| 171 | Duplicate(sumVreg, 0); | 171 | Duplicate(sumVreg, 0); |
| 172 | for (uint16_t j = 0; j < kOuter; ++j) { | 172 | for (uint16_t j = 0; j < kOuter; ++j) { |
| 173 | LoadIfNeedCast<T>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataPadInner, pregFull); | 173 | LoadIfNeedCast<T>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataPadInner, pregFull); |
| 174 | LoadIfNeedCast<T>(gradVreg, gradUb + i * FLOAT_REPEAT_SIZE + j * dataPadInner, pregFull); | 174 | LoadIfNeedCast<T>(gradVreg, gradUb + i * FLOAT_REPEAT_SIZE + j * dataPadInner, pregFull); |
| 175 | - MicroAPI::Mul(tmpVreg, gradVreg, srcVreg, pregFull); | 175 | + Reg::Mul(tmpVreg, gradVreg, srcVreg, pregFull); |
| 176 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 176 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 177 | } | 177 | } |
| 178 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); | 178 | + Reg::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); |
| 179 | - MicroAPI::StoreAlign(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); | 179 | + Reg::StoreAlign(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); |
| 180 | } | 180 | } |
| 181 | 181 | ||
| 182 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 182 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 183 | 183 | ||
| 184 | for (uint16_t i = 0; i < VcgFoldRepeat; i++) { | 184 | for (uint16_t i = 0; i < VcgFoldRepeat; i++) { |
| 185 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 185 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 186 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 186 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 187 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 187 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 188 | if constexpr (isFront) { | 188 | if constexpr (isFront) { |
| 189 | StoreIfNeedCast<T>(workUbFront + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); | 189 | StoreIfNeedCast<T>(workUbFront + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); |
| 190 | } else { | 190 | } else { |
| 191 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 191 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 192 | workUb + i * HALF_REPEAT_SIZE, sumVreg, sumVreg, pregFull); | 192 | workUb + i * HALF_REPEAT_SIZE, sumVreg, sumVreg, pregFull); |
| 193 | } | 193 | } |
| 194 | } | 194 | } |
| 195 | 195 | ||
| 196 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 196 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 197 | 197 | ||
| 198 | if constexpr (isFront) { | 198 | if constexpr (isFront) { |
| 199 | uint32_t sreg = oriM * dtypeBlkStride; | 199 | uint32_t sreg = oriM * dtypeBlkStride; |
| 200 | for (uint16_t i = 0; i < e2bRep; ++i) { | 200 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 201 | - pregCnt = MicroAPI::UpdateMask<T>(sreg); | 201 | + pregCnt = Reg::UpdateMask<T>(sreg); |
| 202 | LoadE2B<T>(castReg, workUbFront + i * DEFAULT_BLK_NUM); | 202 | LoadE2B<T>(castReg, workUbFront + i * DEFAULT_BLK_NUM); |
| 203 | - MicroAPI::StoreAlign(dstUb + i * dtypeRepStride, castReg, pregCnt); | 203 | + Reg::StoreAlign(dstUb + i * dtypeRepStride, castReg, pregCnt); |
| 204 | } | 204 | } |
| 205 | } else { | 205 | } else { |
| 206 | for (uint16_t j = 0; j < kOuter; ++j) { | 206 | for (uint16_t j = 0; j < kOuter; ++j) { |
| 207 | uint32_t sreg = dataLenInner; | 207 | uint32_t sreg = dataLenInner; |
| 208 | for (uint16_t i = 0; i < kRepeatInner; ++i) { | 208 | for (uint16_t i = 0; i < kRepeatInner; ++i) { |
| 209 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 209 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 210 | LoadIfNeedCast<T>(srcVreg, srcUb + j * dataPadInner + i * FLOAT_REPEAT_SIZE, pregFull); | 210 | LoadIfNeedCast<T>(srcVreg, srcUb + j * dataPadInner + i * FLOAT_REPEAT_SIZE, pregFull); |
| 211 | LoadIfNeedCast<T>(gradVreg, gradUb + j * dataPadInner + i * FLOAT_REPEAT_SIZE, pregFull); | 211 | LoadIfNeedCast<T>(gradVreg, gradUb + j * dataPadInner + i * FLOAT_REPEAT_SIZE, pregFull); |
| 212 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); | 212 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); |
| 213 | - MicroAPI::Sub(tmpVreg, gradVreg, sumVreg, pregCnt); | 213 | + Reg::Sub(tmpVreg, gradVreg, sumVreg, pregCnt); |
| 214 | - MicroAPI::Mul(tmpVreg, srcVreg, tmpVreg, pregCnt); | 214 | + Reg::Mul(tmpVreg, srcVreg, tmpVreg, pregCnt); |
| 215 | StoreIfNeedCast<T>(dstUb + j * dataPadInner + i * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); | 215 | StoreIfNeedCast<T>(dstUb + j * dataPadInner + i * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); |
| 216 | } | 216 | } |
| 217 | } | 217 | } |
| @@ -244,29 +244,29 @@ __simd_vf__ inline void SoftMaxGradGenericNDVFImpl(__ubuf__ T *dstUb, __ubuf__ T | |||
| 244 | uint16_t oriK = originalSrcShape.k; | 244 | uint16_t oriK = originalSrcShape.k; |
| 245 | uint16_t repeatTimes = (srcK + FLOAT_REPEAT_SIZE - 1) / FLOAT_REPEAT_SIZE; | 245 | uint16_t repeatTimes = (srcK + FLOAT_REPEAT_SIZE - 1) / FLOAT_REPEAT_SIZE; |
| 246 | 246 | ||
| 247 | - MicroAPI::MaskReg pregCnt; | 247 | + Reg::MaskReg pregCnt; |
| 248 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 248 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 249 | - MicroAPI::MaskReg pregOneBlk; | 249 | + Reg::MaskReg pregOneBlk; |
| 250 | if constexpr (IsSameType<T, half>::value) { | 250 | if constexpr (IsSameType<T, half>::value) { |
| 251 | - pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL16>(); | 251 | + pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL16>(); |
| 252 | } else { | 252 | } else { |
| 253 | - pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 253 | + pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 254 | } | 254 | } |
| 255 | - MicroAPI::RegTensor<float> srcVreg; | 255 | + Reg::RegTensor<float> srcVreg; |
| 256 | - MicroAPI::RegTensor<float> gradVreg; | 256 | + Reg::RegTensor<float> gradVreg; |
| 257 | - MicroAPI::RegTensor<float> sumVreg; | 257 | + Reg::RegTensor<float> sumVreg; |
| 258 | - MicroAPI::RegTensor<float> tmpVreg; | 258 | + Reg::RegTensor<float> tmpVreg; |
| 259 | for (uint16_t i = 0; i < srcM; ++i) { | 259 | for (uint16_t i = 0; i < srcM; ++i) { |
| 260 | uint32_t sreg = oriK; | 260 | uint32_t sreg = oriK; |
| 261 | Duplicate(sumVreg, 0); | 261 | Duplicate(sumVreg, 0); |
| 262 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 262 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 263 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 263 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 264 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 264 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 265 | LoadIfNeedCast<T>(gradVreg, gradUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 265 | LoadIfNeedCast<T>(gradVreg, gradUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 266 | - MicroAPI::Mul(tmpVreg, gradVreg, srcVreg, pregCnt); | 266 | + Reg::Mul(tmpVreg, gradVreg, srcVreg, pregCnt); |
| 267 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 267 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 268 | } | 268 | } |
| 269 | - MicroAPI::ReduceSum(tmpVreg, sumVreg, pregFull); | 269 | + Reg::ReduceSum(tmpVreg, sumVreg, pregFull); |
| 270 | if constexpr (isFront) { | 270 | if constexpr (isFront) { |
| 271 | Duplicate(tmpVreg, tmpVreg, pregOneBlk); | 271 | Duplicate(tmpVreg, tmpVreg, pregOneBlk); |
| 272 | StoreIfNeedCast<T>(dstUb + i * blockStride, tmpVreg, pregOneBlk); | 272 | StoreIfNeedCast<T>(dstUb + i * blockStride, tmpVreg, pregOneBlk); |
| @@ -274,11 +274,11 @@ __simd_vf__ inline void SoftMaxGradGenericNDVFImpl(__ubuf__ T *dstUb, __ubuf__ T | |||
| 274 | Duplicate(sumVreg, tmpVreg, pregFull); | 274 | Duplicate(sumVreg, tmpVreg, pregFull); |
| 275 | sreg = oriK; | 275 | sreg = oriK; |
| 276 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 276 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 277 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 277 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 278 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 278 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 279 | LoadIfNeedCast<T>(gradVreg, gradUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 279 | LoadIfNeedCast<T>(gradVreg, gradUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 280 | - MicroAPI::Sub(tmpVreg, gradVreg, sumVreg, pregCnt); | 280 | + Reg::Sub(tmpVreg, gradVreg, sumVreg, pregCnt); |
| 281 | - MicroAPI::Mul(tmpVreg, srcVreg, tmpVreg, pregCnt); | 281 | + Reg::Mul(tmpVreg, srcVreg, tmpVreg, pregCnt); |
| 282 | StoreIfNeedCast<T>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); | 282 | StoreIfNeedCast<T>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); |
| 283 | } | 283 | } |
| 284 | } | 284 | } |
| @@ -40,92 +40,92 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SoftMaxGenericNZVFImpl(__ubuf__ T1 | |||
| 40 | NotNumUnion notNum; | 40 | NotNumUnion notNum; |
| 41 | notNum.i = F32_NEG_INF; | 41 | notNum.i = F32_NEG_INF; |
| 42 | 42 | ||
| 43 | - MicroAPI::MaskReg pregCnt; | 43 | + Reg::MaskReg pregCnt; |
| 44 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 44 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 45 | - MicroAPI::MaskReg pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 45 | + Reg::MaskReg pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 46 | - MicroAPI::RegTensor<float> srcVreg; | 46 | + Reg::RegTensor<float> srcVreg; |
| 47 | - MicroAPI::RegTensor<float> maxVreg; | 47 | + Reg::RegTensor<float> maxVreg; |
| 48 | - MicroAPI::RegTensor<float> sumVreg; | 48 | + Reg::RegTensor<float> sumVreg; |
| 49 | - MicroAPI::RegTensor<float> tmpVreg; | 49 | + Reg::RegTensor<float> tmpVreg; |
| 50 | - MicroAPI::RegTensor<float> dstVreg; | 50 | + Reg::RegTensor<float> dstVreg; |
| 51 | - MicroAPI::RegTensor<T2> castReg; | 51 | + Reg::RegTensor<T2> castReg; |
| 52 | 52 | ||
| 53 | // reducemax | 53 | // reducemax |
| 54 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 54 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 55 | Duplicate(maxVreg, notNum.f); | 55 | Duplicate(maxVreg, notNum.f); |
| 56 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 56 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 57 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); | 57 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); |
| 58 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 58 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 59 | } | 59 | } |
| 60 | - MicroAPI::ReduceMaxWithDataBlock(maxVreg, maxVreg, pregFull); | 60 | + Reg::ReduceMaxWithDataBlock(maxVreg, maxVreg, pregFull); |
| 61 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, maxVreg, pregOneBlk); | 61 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, maxVreg, pregOneBlk); |
| 62 | } | 62 | } |
| 63 | 63 | ||
| 64 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 64 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 65 | 65 | ||
| 66 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { | 66 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { |
| 67 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 67 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 68 | maxVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 68 | maxVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 69 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, pregFull); | 69 | + Reg::Max(maxVreg, maxVreg, tmpVreg, pregFull); |
| 70 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); | 70 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); |
| 71 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 71 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 72 | workUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); | 72 | workUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); |
| 73 | } | 73 | } |
| 74 | 74 | ||
| 75 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 75 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 76 | 76 | ||
| 77 | uint32_t sreg = originM * dtypeBlkStride; | 77 | uint32_t sreg = originM * dtypeBlkStride; |
| 78 | for (uint16_t i = 0; i < e2bRep; ++i) { | 78 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 79 | - pregCnt = MicroAPI::UpdateMask<T2>(sreg); | 79 | + pregCnt = Reg::UpdateMask<T2>(sreg); |
| 80 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); | 80 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); |
| 81 | - MicroAPI::StoreAlign(maxUb + i * dtypeRepStride, castReg, pregCnt); | 81 | + Reg::StoreAlign(maxUb + i * dtypeRepStride, castReg, pregCnt); |
| 82 | } | 82 | } |
| 83 | 83 | ||
| 84 | // reducesum | 84 | // reducesum |
| 85 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 85 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 86 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 86 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 87 | Duplicate(sumVreg, 0); | 87 | Duplicate(sumVreg, 0); |
| 88 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); | 88 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); |
| 89 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 89 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 90 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); | 90 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); |
| 91 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, pregFull); | 91 | + Reg::Sub(dstVreg, srcVreg, maxVreg, pregFull); |
| 92 | - MicroAPI::Exp(tmpVreg, dstVreg, pregFull); | 92 | + Reg::Exp(tmpVreg, dstVreg, pregFull); |
| 93 | - MicroAPI::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregFull); | 93 | + Reg::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregFull); |
| 94 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 94 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 95 | } | 95 | } |
| 96 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); | 96 | + Reg::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); |
| 97 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); | 97 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); |
| 98 | } | 98 | } |
| 99 | 99 | ||
| 100 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 100 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 101 | 101 | ||
| 102 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { | 102 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { |
| 103 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 103 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 104 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 104 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 105 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 105 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 106 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); | 106 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); |
| 107 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 107 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 108 | workUb + i * HALF_REPEAT_SIZE, sumVreg, sumVreg, pregFull); | 108 | workUb + i * HALF_REPEAT_SIZE, sumVreg, sumVreg, pregFull); |
| 109 | } | 109 | } |
| 110 | 110 | ||
| 111 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 111 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 112 | 112 | ||
| 113 | sreg = originM * dtypeBlkStride; | 113 | sreg = originM * dtypeBlkStride; |
| 114 | for (uint16_t i = 0; i < e2bRep; ++i) { | 114 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 115 | - pregCnt = MicroAPI::UpdateMask<T2>(sreg); | 115 | + pregCnt = Reg::UpdateMask<T2>(sreg); |
| 116 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); | 116 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); |
| 117 | - MicroAPI::StoreAlign(sumUb + i * dtypeRepStride, castReg, pregCnt); | 117 | + Reg::StoreAlign(sumUb + i * dtypeRepStride, castReg, pregCnt); |
| 118 | } | 118 | } |
| 119 | 119 | ||
| 120 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 120 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 121 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; | 121 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; |
| 122 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 122 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 123 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 123 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 124 | - MicroAPI::LoadAlign(tmpVreg, expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); | 124 | + Reg::LoadAlign(tmpVreg, expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); |
| 125 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); | 125 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); |
| 126 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, pregCnt); | 126 | + Reg::Div(dstVreg, tmpVreg, sumVreg, pregCnt); |
| 127 | if constexpr (isLog) { | 127 | if constexpr (isLog) { |
| 128 | - MicroAPI::Log10(dstVreg, dstVreg, pregCnt); | 128 | + Reg::Log10(dstVreg, dstVreg, pregCnt); |
| 129 | } | 129 | } |
| 130 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, dstVreg, pregCnt); | 130 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, dstVreg, pregCnt); |
| 131 | } | 131 | } |
| @@ -168,17 +168,17 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SoftMaxGenericNZWithTailVFImpl(__u | |||
| 168 | NotNumUnion notNum; | 168 | NotNumUnion notNum; |
| 169 | notNum.i = F32_NEG_INF; | 169 | notNum.i = F32_NEG_INF; |
| 170 | 170 | ||
| 171 | - MicroAPI::MaskReg pregkTail = MicroAPI::MoveMask<uint32_t>(); | 171 | + Reg::MaskReg pregkTail = Reg::MoveMask<uint32_t>(); |
| 172 | - MicroAPI::MaskReg pregCnt; | 172 | + Reg::MaskReg pregCnt; |
| 173 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 173 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 174 | - MicroAPI::MaskReg pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 174 | + Reg::MaskReg pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 175 | - MicroAPI::RegTensor<float> srcVreg; | 175 | + Reg::RegTensor<float> srcVreg; |
| 176 | - MicroAPI::RegTensor<float> maxVreg; | 176 | + Reg::RegTensor<float> maxVreg; |
| 177 | - MicroAPI::RegTensor<float> sumVreg; | 177 | + Reg::RegTensor<float> sumVreg; |
| 178 | - MicroAPI::RegTensor<float> tmpVreg; | 178 | + Reg::RegTensor<float> tmpVreg; |
| 179 | - MicroAPI::RegTensor<float> minVreg; | 179 | + Reg::RegTensor<float> minVreg; |
| 180 | - MicroAPI::RegTensor<float> dstVreg; | 180 | + Reg::RegTensor<float> dstVreg; |
| 181 | - MicroAPI::RegTensor<T2> castReg; | 181 | + Reg::RegTensor<T2> castReg; |
| 182 | 182 | ||
| 183 | // reducemax | 183 | // reducemax |
| 184 | Duplicate(minVreg, notNum.f); | 184 | Duplicate(minVreg, notNum.f); |
| @@ -186,100 +186,100 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SoftMaxGenericNZWithTailVFImpl(__u | |||
| 186 | Duplicate(maxVreg, notNum.f); | 186 | Duplicate(maxVreg, notNum.f); |
| 187 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 187 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 188 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); | 188 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); |
| 189 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 189 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 190 | } | 190 | } |
| 191 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, pregFull); | 191 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, pregFull); |
| 192 | - MicroAPI::Select(srcVreg, srcVreg, minVreg, pregkTail); | 192 | + Reg::Select(srcVreg, srcVreg, minVreg, pregkTail); |
| 193 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 193 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 194 | 194 | ||
| 195 | - MicroAPI::ReduceMaxWithDataBlock(maxVreg, maxVreg, pregFull); | 195 | + Reg::ReduceMaxWithDataBlock(maxVreg, maxVreg, pregFull); |
| 196 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, maxVreg, pregOneBlk); | 196 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, maxVreg, pregOneBlk); |
| 197 | } | 197 | } |
| 198 | 198 | ||
| 199 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 199 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 200 | 200 | ||
| 201 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { | 201 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { |
| 202 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 202 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 203 | maxVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 203 | maxVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 204 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, pregFull); | 204 | + Reg::Max(maxVreg, maxVreg, tmpVreg, pregFull); |
| 205 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); | 205 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, maxVreg, pregFull); |
| 206 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 206 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 207 | workUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); | 207 | workUb + i * HALF_REPEAT_SIZE, maxVreg, maxVreg, pregFull); |
| 208 | } | 208 | } |
| 209 | 209 | ||
| 210 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 210 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 211 | 211 | ||
| 212 | uint32_t sreg = originM * dtypeBlkStride; | 212 | uint32_t sreg = originM * dtypeBlkStride; |
| 213 | for (uint16_t i = 0; i < e2bRep; ++i) { | 213 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 214 | - pregCnt = MicroAPI::UpdateMask<T2>(sreg); | 214 | + pregCnt = Reg::UpdateMask<T2>(sreg); |
| 215 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); | 215 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); |
| 216 | - MicroAPI::StoreAlign(maxUb + i * dtypeRepStride, castReg, pregCnt); | 216 | + Reg::StoreAlign(maxUb + i * dtypeRepStride, castReg, pregCnt); |
| 217 | } | 217 | } |
| 218 | 218 | ||
| 219 | // reducesum | 219 | // reducesum |
| 220 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 220 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 221 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 221 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 222 | Duplicate(sumVreg, 0); | 222 | Duplicate(sumVreg, 0); |
| 223 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); | 223 | LoadE2B<float>(maxVreg, workUb + i * DEFAULT_BLK_NUM); |
| 224 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 224 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 225 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); | 225 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, pregFull); |
| 226 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, pregFull); | 226 | + Reg::Sub(dstVreg, srcVreg, maxVreg, pregFull); |
| 227 | - MicroAPI::Exp(tmpVreg, dstVreg, pregFull); | 227 | + Reg::Exp(tmpVreg, dstVreg, pregFull); |
| 228 | - MicroAPI::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregFull); | 228 | + Reg::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, tmpVreg, pregFull); |
| 229 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 229 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 230 | } | 230 | } |
| 231 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, pregFull); | 231 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, pregFull); |
| 232 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, pregkTail); | 232 | + Reg::Sub(dstVreg, srcVreg, maxVreg, pregkTail); |
| 233 | - MicroAPI::Exp(tmpVreg, dstVreg, pregkTail); | 233 | + Reg::Exp(tmpVreg, dstVreg, pregkTail); |
| 234 | - MicroAPI::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, tmpVreg, pregkTail); | 234 | + Reg::StoreAlign(expUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, tmpVreg, pregkTail); |
| 235 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 235 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 236 | 236 | ||
| 237 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); | 237 | + Reg::ReduceSumWithDataBlock(sumVreg, sumVreg, pregFull); |
| 238 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); | 238 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_NORM>(workUb + i * DEFAULT_BLK_NUM, sumVreg, pregOneBlk); |
| 239 | } | 239 | } |
| 240 | 240 | ||
| 241 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 241 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 242 | 242 | ||
| 243 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { | 243 | for (uint16_t i = 0; i < VcgFoldRepeat; ++i) { |
| 244 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>( | 244 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_DINTLV_B32>( |
| 245 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); | 245 | sumVreg, tmpVreg, workUb + i * HALF_REPEAT_SIZE); |
| 246 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 246 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 247 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); | 247 | StoreIfNeedCast<T2>(tmpUb + i * FLOAT_REPEAT_SIZE, sumVreg, pregFull); |
| 248 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_INTLV_B32>( | 248 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_INTLV_B32>( |
| 249 | workUb + i * HALF_REPEAT_SIZE, sumVreg, sumVreg, pregFull); | 249 | workUb + i * HALF_REPEAT_SIZE, sumVreg, sumVreg, pregFull); |
| 250 | } | 250 | } |
| 251 | 251 | ||
| 252 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 252 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 253 | 253 | ||
| 254 | sreg = originM * dtypeBlkStride; | 254 | sreg = originM * dtypeBlkStride; |
| 255 | for (uint16_t i = 0; i < e2bRep; ++i) { | 255 | for (uint16_t i = 0; i < e2bRep; ++i) { |
| 256 | - pregCnt = MicroAPI::UpdateMask<T2>(sreg); | 256 | + pregCnt = Reg::UpdateMask<T2>(sreg); |
| 257 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); | 257 | LoadE2B<T2>(castReg, tmpUb + i * DEFAULT_BLK_NUM); |
| 258 | - MicroAPI::StoreAlign(sumUb + i * dtypeRepStride, castReg, pregCnt); | 258 | + Reg::StoreAlign(sumUb + i * dtypeRepStride, castReg, pregCnt); |
| 259 | } | 259 | } |
| 260 | 260 | ||
| 261 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 261 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 262 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; | 262 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; |
| 263 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 263 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 264 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 264 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 265 | - MicroAPI::LoadAlign(tmpVreg, expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); | 265 | + Reg::LoadAlign(tmpVreg, expUb + i * FLOAT_REPEAT_SIZE + j * dataBlock); |
| 266 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); | 266 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); |
| 267 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, pregCnt); | 267 | + Reg::Div(dstVreg, tmpVreg, sumVreg, pregCnt); |
| 268 | if constexpr (isLog) { | 268 | if constexpr (isLog) { |
| 269 | - MicroAPI::Log10(dstVreg, dstVreg, pregCnt); | 269 | + Reg::Log10(dstVreg, dstVreg, pregCnt); |
| 270 | } | 270 | } |
| 271 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, dstVreg, pregCnt); | 271 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + j * dataBlock, dstVreg, pregCnt); |
| 272 | } | 272 | } |
| 273 | } | 273 | } |
| 274 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; | 274 | sreg = originM * SOFTMAX_SHAPE_NZ_BASIC_COUNT; |
| 275 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 275 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 276 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 276 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 277 | - MicroAPI::LoadAlign(tmpVreg, expUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock); | 277 | + Reg::LoadAlign(tmpVreg, expUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock); |
| 278 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); | 278 | LoadE2B<float>(sumVreg, workUb + i * DEFAULT_BLK_NUM); |
| 279 | - MicroAPI::MaskAnd(pregOneBlk, pregkTail, pregCnt, pregFull); | 279 | + Reg::MaskAnd(pregOneBlk, pregkTail, pregCnt, pregFull); |
| 280 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, pregOneBlk); | 280 | + Reg::Div(dstVreg, tmpVreg, sumVreg, pregOneBlk); |
| 281 | if constexpr (isLog) { | 281 | if constexpr (isLog) { |
| 282 | - MicroAPI::Log10(dstVreg, dstVreg, pregOneBlk); | 282 | + Reg::Log10(dstVreg, dstVreg, pregOneBlk); |
| 283 | } | 283 | } |
| 284 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, dstVreg, pregOneBlk); | 284 | StoreIfNeedCast<T1>(dstUb + i * FLOAT_REPEAT_SIZE + kRepeatTimes * dataBlock, dstVreg, pregOneBlk); |
| 285 | } | 285 | } |
| @@ -339,27 +339,27 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SoftMaxGenericNDVFImpl(__ubuf__ T1 | |||
| 339 | NotNumUnion notNum; | 339 | NotNumUnion notNum; |
| 340 | notNum.i = F32_NEG_INF; | 340 | notNum.i = F32_NEG_INF; |
| 341 | 341 | ||
| 342 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 342 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 343 | - MicroAPI::MaskReg pregOneBlk; | 343 | + Reg::MaskReg pregOneBlk; |
| 344 | - MicroAPI::MaskReg pregOnePt = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL1>(); | 344 | + Reg::MaskReg pregOnePt = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL1>(); |
| 345 | if constexpr (IsSameType<T2, half>::value) { | 345 | if constexpr (IsSameType<T2, half>::value) { |
| 346 | - pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL16>(); | 346 | + pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL16>(); |
| 347 | } else { | 347 | } else { |
| 348 | - pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 348 | + pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 349 | } | 349 | } |
| 350 | - MicroAPI::RegTensor<float> srcVreg; | 350 | + Reg::RegTensor<float> srcVreg; |
| 351 | - MicroAPI::RegTensor<float> maxVreg; | 351 | + Reg::RegTensor<float> maxVreg; |
| 352 | - MicroAPI::RegTensor<float> sumVreg; | 352 | + Reg::RegTensor<float> sumVreg; |
| 353 | - MicroAPI::RegTensor<float> tmpVreg; | 353 | + Reg::RegTensor<float> tmpVreg; |
| 354 | - MicroAPI::RegTensor<float> dstVreg; | 354 | + Reg::RegTensor<float> dstVreg; |
| 355 | 355 | ||
| 356 | for (uint16_t i = 0; i < srcM; ++i) { | 356 | for (uint16_t i = 0; i < srcM; ++i) { |
| 357 | Duplicate(maxVreg, notNum.f); | 357 | Duplicate(maxVreg, notNum.f); |
| 358 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 358 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 359 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 359 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 360 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 360 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 361 | } | 361 | } |
| 362 | - MicroAPI::ReduceMax(maxVreg, maxVreg, pregFull); | 362 | + Reg::ReduceMax(maxVreg, maxVreg, pregFull); |
| 363 | if constexpr (outputBrc) { | 363 | if constexpr (outputBrc) { |
| 364 | Duplicate(maxVreg, maxVreg, pregOneBlk); | 364 | Duplicate(maxVreg, maxVreg, pregOneBlk); |
| 365 | StoreIfNeedCast<T2>(maxUb + i * blockStride, maxVreg, pregOneBlk); | 365 | StoreIfNeedCast<T2>(maxUb + i * blockStride, maxVreg, pregOneBlk); |
| @@ -371,16 +371,16 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SoftMaxGenericNDVFImpl(__ubuf__ T1 | |||
| 371 | Duplicate(maxVreg, maxVreg, pregFull); | 371 | Duplicate(maxVreg, maxVreg, pregFull); |
| 372 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 372 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 373 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 373 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 374 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, pregFull); | 374 | + Reg::Sub(dstVreg, srcVreg, maxVreg, pregFull); |
| 375 | - MicroAPI::Exp(tmpVreg, dstVreg, pregFull); | 375 | + Reg::Exp(tmpVreg, dstVreg, pregFull); |
| 376 | if constexpr (!isFlashV2) { | 376 | if constexpr (!isFlashV2) { |
| 377 | - MicroAPI::StoreAlign(workUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregFull); | 377 | + Reg::StoreAlign(workUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregFull); |
| 378 | } else { | 378 | } else { |
| 379 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregFull); | 379 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregFull); |
| 380 | } | 380 | } |
| 381 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 381 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 382 | } | 382 | } |
| 383 | - MicroAPI::ReduceSum(sumVreg, sumVreg, pregFull); | 383 | + Reg::ReduceSum(sumVreg, sumVreg, pregFull); |
| 384 | if constexpr (outputBrc) { | 384 | if constexpr (outputBrc) { |
| 385 | Duplicate(sumVreg, sumVreg, pregOneBlk); | 385 | Duplicate(sumVreg, sumVreg, pregOneBlk); |
| 386 | StoreIfNeedCast<T2>(sumUb + i * blockStride, sumVreg, pregOneBlk); | 386 | StoreIfNeedCast<T2>(sumUb + i * blockStride, sumVreg, pregOneBlk); |
| @@ -388,24 +388,24 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SoftMaxGenericNDVFImpl(__ubuf__ T1 | |||
| 388 | StoreIfNeedCastM1<T2>(sumUb + i, sumVreg, pregOnePt); | 388 | StoreIfNeedCastM1<T2>(sumUb + i, sumVreg, pregOnePt); |
| 389 | } | 389 | } |
| 390 | if constexpr (!isFlashV2 && sizeof(T2) == sizeof(half)) { | 390 | if constexpr (!isFlashV2 && sizeof(T2) == sizeof(half)) { |
| 391 | - MicroAPI::StoreAlign(tmpUb + i * blockStride, sumVreg, pregOneBlk); | 391 | + Reg::StoreAlign(tmpUb + i * blockStride, sumVreg, pregOneBlk); |
| 392 | } | 392 | } |
| 393 | } | 393 | } |
| 394 | 394 | ||
| 395 | if constexpr (!isFlashV2) { | 395 | if constexpr (!isFlashV2) { |
| 396 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 396 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 397 | for (uint16_t i = 0; i < srcM; ++i) { | 397 | for (uint16_t i = 0; i < srcM; ++i) { |
| 398 | if constexpr (sizeof(T2) == sizeof(half)) { | 398 | if constexpr (sizeof(T2) == sizeof(half)) { |
| 399 | - MicroAPI::LoadAlign(sumVreg, tmpUb + i * blockStride); | 399 | + Reg::LoadAlign(sumVreg, tmpUb + i * blockStride); |
| 400 | } else { | 400 | } else { |
| 401 | - MicroAPI::LoadAlign(sumVreg, sumUb + i * blockStride); | 401 | + Reg::LoadAlign(sumVreg, sumUb + i * blockStride); |
| 402 | } | 402 | } |
| 403 | Duplicate(sumVreg, sumVreg, pregFull); | 403 | Duplicate(sumVreg, sumVreg, pregFull); |
| 404 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 404 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 405 | - MicroAPI::LoadAlign(tmpVreg, workUb + i * srcK + j * FLOAT_REPEAT_SIZE); | 405 | + Reg::LoadAlign(tmpVreg, workUb + i * srcK + j * FLOAT_REPEAT_SIZE); |
| 406 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, pregFull); | 406 | + Reg::Div(dstVreg, tmpVreg, sumVreg, pregFull); |
| 407 | if constexpr (isLog) { | 407 | if constexpr (isLog) { |
| 408 | - MicroAPI::Log10(dstVreg, dstVreg, pregFull); | 408 | + Reg::Log10(dstVreg, dstVreg, pregFull); |
| 409 | } | 409 | } |
| 410 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, dstVreg, pregFull); | 410 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, dstVreg, pregFull); |
| 411 | } | 411 | } |
| @@ -444,37 +444,37 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SoftMaxGenericNDWithTailVFImpl(__u | |||
| 444 | NotNumUnion notNum; | 444 | NotNumUnion notNum; |
| 445 | notNum.i = F32_NEG_INF; | 445 | notNum.i = F32_NEG_INF; |
| 446 | 446 | ||
| 447 | - MicroAPI::MaskReg pregCnt; | 447 | + Reg::MaskReg pregCnt; |
| 448 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 448 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 449 | - MicroAPI::MaskReg pregOnePt = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL1>(); | 449 | + Reg::MaskReg pregOnePt = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL1>(); |
| 450 | - MicroAPI::MaskReg pregOneBlk; | 450 | + Reg::MaskReg pregOneBlk; |
| 451 | if constexpr (IsSameType<T2, half>::value) { | 451 | if constexpr (IsSameType<T2, half>::value) { |
| 452 | - pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL16>(); | 452 | + pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL16>(); |
| 453 | } else { | 453 | } else { |
| 454 | - pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 454 | + pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 455 | } | 455 | } |
| 456 | - MicroAPI::RegTensor<float> srcVreg; | 456 | + Reg::RegTensor<float> srcVreg; |
| 457 | - MicroAPI::RegTensor<float> maxVreg; | 457 | + Reg::RegTensor<float> maxVreg; |
| 458 | - MicroAPI::RegTensor<float> sumVreg; | 458 | + Reg::RegTensor<float> sumVreg; |
| 459 | - MicroAPI::RegTensor<float> tmpVreg; | 459 | + Reg::RegTensor<float> tmpVreg; |
| 460 | - MicroAPI::RegTensor<float> minVreg; | 460 | + Reg::RegTensor<float> minVreg; |
| 461 | - MicroAPI::RegTensor<float> dstVreg; | 461 | + Reg::RegTensor<float> dstVreg; |
| 462 | 462 | ||
| 463 | Duplicate(minVreg, notNum.f); | 463 | Duplicate(minVreg, notNum.f); |
| 464 | for (uint16_t i = 0; i < srcM; ++i) { | 464 | for (uint16_t i = 0; i < srcM; ++i) { |
| 465 | uint32_t sreg = originK; | 465 | uint32_t sreg = originK; |
| 466 | Duplicate(maxVreg, notNum.f); | 466 | Duplicate(maxVreg, notNum.f); |
| 467 | for (uint16_t j = 0; j < static_cast<uint16_t>(repeatTimes - 1); ++j) { | 467 | for (uint16_t j = 0; j < static_cast<uint16_t>(repeatTimes - 1); ++j) { |
| 468 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 468 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 469 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 469 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 470 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregCnt); | 470 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregCnt); |
| 471 | } | 471 | } |
| 472 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 472 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 473 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + (repeatTimes - 1) * FLOAT_REPEAT_SIZE, pregFull); | 473 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + (repeatTimes - 1) * FLOAT_REPEAT_SIZE, pregFull); |
| 474 | - MicroAPI::Select(srcVreg, srcVreg, minVreg, pregCnt); | 474 | + Reg::Select(srcVreg, srcVreg, minVreg, pregCnt); |
| 475 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 475 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 476 | 476 | ||
| 477 | - MicroAPI::ReduceMax(maxVreg, maxVreg, pregFull); | 477 | + Reg::ReduceMax(maxVreg, maxVreg, pregFull); |
| 478 | if constexpr (outputBrc) { | 478 | if constexpr (outputBrc) { |
| 479 | Duplicate(maxVreg, maxVreg, pregOneBlk); | 479 | Duplicate(maxVreg, maxVreg, pregOneBlk); |
| 480 | StoreIfNeedCast<T2>(maxUb + i * blockStride, maxVreg, pregOneBlk); | 480 | StoreIfNeedCast<T2>(maxUb + i * blockStride, maxVreg, pregOneBlk); |
| @@ -486,18 +486,18 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SoftMaxGenericNDWithTailVFImpl(__u | |||
| 486 | Duplicate(maxVreg, maxVreg, pregFull); | 486 | Duplicate(maxVreg, maxVreg, pregFull); |
| 487 | sreg = originK; | 487 | sreg = originK; |
| 488 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 488 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 489 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 489 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 490 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 490 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 491 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, pregCnt); | 491 | + Reg::Sub(dstVreg, srcVreg, maxVreg, pregCnt); |
| 492 | - MicroAPI::Exp(tmpVreg, dstVreg, pregCnt); | 492 | + Reg::Exp(tmpVreg, dstVreg, pregCnt); |
| 493 | if constexpr (!isFlashV2) { | 493 | if constexpr (!isFlashV2) { |
| 494 | - MicroAPI::StoreAlign(workUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); | 494 | + Reg::StoreAlign(workUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); |
| 495 | } else { | 495 | } else { |
| 496 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); | 496 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); |
| 497 | } | 497 | } |
| 498 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 498 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 499 | } | 499 | } |
| 500 | - MicroAPI::ReduceSum(sumVreg, sumVreg, pregFull); | 500 | + Reg::ReduceSum(sumVreg, sumVreg, pregFull); |
| 501 | if constexpr (outputBrc) { | 501 | if constexpr (outputBrc) { |
| 502 | Duplicate(sumVreg, sumVreg, pregOneBlk); | 502 | Duplicate(sumVreg, sumVreg, pregOneBlk); |
| 503 | StoreIfNeedCast<T2>(sumUb + i * blockStride, sumVreg, pregOneBlk); | 503 | StoreIfNeedCast<T2>(sumUb + i * blockStride, sumVreg, pregOneBlk); |
| @@ -505,26 +505,26 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SoftMaxGenericNDWithTailVFImpl(__u | |||
| 505 | StoreIfNeedCastM1<T2>(sumUb + i, sumVreg, pregOnePt); | 505 | StoreIfNeedCastM1<T2>(sumUb + i, sumVreg, pregOnePt); |
| 506 | } | 506 | } |
| 507 | if constexpr (!isFlashV2 && sizeof(T2) == sizeof(half)) { | 507 | if constexpr (!isFlashV2 && sizeof(T2) == sizeof(half)) { |
| 508 | - MicroAPI::StoreAlign(tmpUb + i * blockStride, sumVreg, pregOneBlk); | 508 | + Reg::StoreAlign(tmpUb + i * blockStride, sumVreg, pregOneBlk); |
| 509 | } | 509 | } |
| 510 | } | 510 | } |
| 511 | 511 | ||
| 512 | if constexpr (!isFlashV2) { | 512 | if constexpr (!isFlashV2) { |
| 513 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 513 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 514 | for (uint16_t i = 0; i < srcM; ++i) { | 514 | for (uint16_t i = 0; i < srcM; ++i) { |
| 515 | if constexpr (sizeof(T2) == sizeof(half)) { | 515 | if constexpr (sizeof(T2) == sizeof(half)) { |
| 516 | - MicroAPI::LoadAlign(sumVreg, tmpUb + i * blockStride); | 516 | + Reg::LoadAlign(sumVreg, tmpUb + i * blockStride); |
| 517 | } else { | 517 | } else { |
| 518 | - MicroAPI::LoadAlign(sumVreg, sumUb + i * blockStride); | 518 | + Reg::LoadAlign(sumVreg, sumUb + i * blockStride); |
| 519 | } | 519 | } |
| 520 | Duplicate(sumVreg, sumVreg, pregFull); | 520 | Duplicate(sumVreg, sumVreg, pregFull); |
| 521 | uint32_t sreg = originK; | 521 | uint32_t sreg = originK; |
| 522 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 522 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 523 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 523 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 524 | - MicroAPI::LoadAlign(tmpVreg, workUb + i * srcK + j * FLOAT_REPEAT_SIZE); | 524 | + Reg::LoadAlign(tmpVreg, workUb + i * srcK + j * FLOAT_REPEAT_SIZE); |
| 525 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, pregCnt); | 525 | + Reg::Div(dstVreg, tmpVreg, sumVreg, pregCnt); |
| 526 | if constexpr (isLog) { | 526 | if constexpr (isLog) { |
| 527 | - MicroAPI::Log10(dstVreg, dstVreg, pregCnt); | 527 | + Reg::Log10(dstVreg, dstVreg, pregCnt); |
| 528 | } | 528 | } |
| 529 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, dstVreg, pregCnt); | 529 | StoreIfNeedCast<T1>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, dstVreg, pregCnt); |
| 530 | } | 530 | } |
| @@ -561,78 +561,78 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SingleSoftMaxGenericNDForBlkVFImpl | |||
| 561 | uint16_t srcM, uint16_t srcK, uint16_t factorRow, uint16_t factor, uint16_t blockStride) | 561 | uint16_t srcM, uint16_t srcK, uint16_t factorRow, uint16_t factor, uint16_t blockStride) |
| 562 | { | 562 | { |
| 563 | uint32_t sreg = srcK * srcM; | 563 | uint32_t sreg = srcK * srcM; |
| 564 | - MicroAPI::MaskReg pregDst; | 564 | + Reg::MaskReg pregDst; |
| 565 | - MicroAPI::MaskReg pregOut; | 565 | + Reg::MaskReg pregOut; |
| 566 | - MicroAPI::MaskReg pregCnt = MicroAPI::MoveMask<uint32_t>(); | 566 | + Reg::MaskReg pregCnt = Reg::MoveMask<uint32_t>(); |
| 567 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 567 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 568 | - MicroAPI::MaskReg pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 568 | + Reg::MaskReg pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 569 | - MicroAPI::RegTensor<float> srcVreg; | 569 | + Reg::RegTensor<float> srcVreg; |
| 570 | - MicroAPI::RegTensor<float> maxVreg; | 570 | + Reg::RegTensor<float> maxVreg; |
| 571 | - MicroAPI::RegTensor<float> sumVreg; | 571 | + Reg::RegTensor<float> sumVreg; |
| 572 | - MicroAPI::RegTensor<float> tmpVreg; | 572 | + Reg::RegTensor<float> tmpVreg; |
| 573 | - MicroAPI::RegTensor<float> dstVreg; | 573 | + Reg::RegTensor<float> dstVreg; |
| 574 | - MicroAPI::UnalignReg ureg0; | 574 | + Reg::UnalignReg ureg0; |
| 575 | 575 | ||
| 576 | for (uint16_t i = 0; i < factor; ++i) { | 576 | for (uint16_t i = 0; i < factor; ++i) { |
| 577 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK * factorRow, pregFull); | 577 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK * factorRow, pregFull); |
| 578 | 578 | ||
| 579 | - MicroAPI::ReduceMaxWithDataBlock(maxVreg, srcVreg, pregCnt); | 579 | + Reg::ReduceMaxWithDataBlock(maxVreg, srcVreg, pregCnt); |
| 580 | - MicroAPI::StoreAlign(tmpUb0 + i * factorRow, maxVreg, pregOneBlk); | 580 | + Reg::StoreAlign(tmpUb0 + i * factorRow, maxVreg, pregOneBlk); |
| 581 | if constexpr (!outputBrc) { | 581 | if constexpr (!outputBrc) { |
| 582 | if constexpr (SupportType<T2, half>()) { | 582 | if constexpr (SupportType<T2, half>()) { |
| 583 | - MicroAPI::RegTensor<T2> castVreg; | 583 | + Reg::RegTensor<T2> castVreg; |
| 584 | - MicroAPI::Cast<T2, float, Internal::castTraitB32ToB16>(castVreg, maxVreg, pregOneBlk); | 584 | + Reg::Cast<T2, float, Internal::castTraitB32ToB16>(castVreg, maxVreg, pregOneBlk); |
| 585 | - MicroAPI::Pack<uint16_t, uint32_t>( | 585 | + Reg::Pack<uint16_t, uint32_t>( |
| 586 | - (MicroAPI::RegTensor<uint16_t>&)castVreg, (MicroAPI::RegTensor<uint32_t>&)castVreg); | 586 | + (Reg::RegTensor<uint16_t>&)castVreg, (Reg::RegTensor<uint32_t>&)castVreg); |
| 587 | - MicroAPI::StoreUnAlign(maxUb, castVreg, ureg0, factorRow); | 587 | + Reg::StoreUnAlign(maxUb, castVreg, ureg0, factorRow); |
| 588 | - MicroAPI::StoreUnAlignPost(maxUb, ureg0, 0); | 588 | + Reg::StoreUnAlignPost(maxUb, ureg0, 0); |
| 589 | } else { | 589 | } else { |
| 590 | - MicroAPI::StoreAlign<float>(maxUb + i * factorRow, maxVreg, pregOneBlk); | 590 | + Reg::StoreAlign<float>(maxUb + i * factorRow, maxVreg, pregOneBlk); |
| 591 | } | 591 | } |
| 592 | } | 592 | } |
| 593 | } | 593 | } |
| 594 | 594 | ||
| 595 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 595 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 596 | 596 | ||
| 597 | for (uint16_t i = 0; i < factor; ++i) { | 597 | for (uint16_t i = 0; i < factor; ++i) { |
| 598 | - pregOut = MicroAPI::UpdateMask<uint32_t>(sreg); | 598 | + pregOut = Reg::UpdateMask<uint32_t>(sreg); |
| 599 | LoadE2B<float>(maxVreg, tmpUb0 + i * factorRow); | 599 | LoadE2B<float>(maxVreg, tmpUb0 + i * factorRow); |
| 600 | if constexpr (outputBrc) { | 600 | if constexpr (outputBrc) { |
| 601 | StoreIfNeedCast<T2>(maxUb + i * blockStride * factorRow, maxVreg, pregOut); | 601 | StoreIfNeedCast<T2>(maxUb + i * blockStride * factorRow, maxVreg, pregOut); |
| 602 | } | 602 | } |
| 603 | 603 | ||
| 604 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK * factorRow, pregFull); | 604 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK * factorRow, pregFull); |
| 605 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, pregFull); | 605 | + Reg::Sub(dstVreg, srcVreg, maxVreg, pregFull); |
| 606 | - MicroAPI::Exp(tmpVreg, dstVreg, pregFull); | 606 | + Reg::Exp(tmpVreg, dstVreg, pregFull); |
| 607 | if constexpr (!isFlashV2) { | 607 | if constexpr (!isFlashV2) { |
| 608 | - MicroAPI::StoreAlign(workUb + i * srcK * factorRow, tmpVreg, pregOut); | 608 | + Reg::StoreAlign(workUb + i * srcK * factorRow, tmpVreg, pregOut); |
| 609 | } else { | 609 | } else { |
| 610 | - MicroAPI::MaskAnd(pregDst, pregCnt, pregOut, pregFull); | 610 | + Reg::MaskAnd(pregDst, pregCnt, pregOut, pregFull); |
| 611 | StoreIfNeedCast<T1>(dstUb + i * srcK * factorRow, tmpVreg, pregDst); | 611 | StoreIfNeedCast<T1>(dstUb + i * srcK * factorRow, tmpVreg, pregDst); |
| 612 | } | 612 | } |
| 613 | 613 | ||
| 614 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, tmpVreg, pregCnt); | 614 | + Reg::ReduceSumWithDataBlock(sumVreg, tmpVreg, pregCnt); |
| 615 | - MicroAPI::StoreAlign(tmpUb1 + i * factorRow, sumVreg, pregOneBlk); | 615 | + Reg::StoreAlign(tmpUb1 + i * factorRow, sumVreg, pregOneBlk); |
| 616 | if constexpr (!outputBrc) { | 616 | if constexpr (!outputBrc) { |
| 617 | if constexpr (SupportType<T2, half>()) { | 617 | if constexpr (SupportType<T2, half>()) { |
| 618 | - MicroAPI::RegTensor<T2> castVreg; | 618 | + Reg::RegTensor<T2> castVreg; |
| 619 | - MicroAPI::Cast<T2, float, Internal::castTraitB32ToB16>(castVreg, sumVreg, pregOneBlk); | 619 | + Reg::Cast<T2, float, Internal::castTraitB32ToB16>(castVreg, sumVreg, pregOneBlk); |
| 620 | - MicroAPI::Pack<uint16_t, uint32_t>( | 620 | + Reg::Pack<uint16_t, uint32_t>( |
| 621 | - (MicroAPI::RegTensor<uint16_t>&)castVreg, (MicroAPI::RegTensor<uint32_t>&)castVreg); | 621 | + (Reg::RegTensor<uint16_t>&)castVreg, (Reg::RegTensor<uint32_t>&)castVreg); |
| 622 | - MicroAPI::StoreUnAlign(sumUb, castVreg, ureg0, factorRow); | 622 | + Reg::StoreUnAlign(sumUb, castVreg, ureg0, factorRow); |
| 623 | - MicroAPI::StoreUnAlignPost(sumUb, ureg0, 0); | 623 | + Reg::StoreUnAlignPost(sumUb, ureg0, 0); |
| 624 | } else { | 624 | } else { |
| 625 | - MicroAPI::StoreAlign<float>(sumUb + i * factorRow, sumVreg, pregOneBlk); | 625 | + Reg::StoreAlign<float>(sumUb + i * factorRow, sumVreg, pregOneBlk); |
| 626 | } | 626 | } |
| 627 | } | 627 | } |
| 628 | } | 628 | } |
| 629 | 629 | ||
| 630 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 630 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 631 | 631 | ||
| 632 | if constexpr (isFlashV2 && outputBrc) { | 632 | if constexpr (isFlashV2 && outputBrc) { |
| 633 | sreg = srcK * srcM; | 633 | sreg = srcK * srcM; |
| 634 | for (uint16_t i = 0; i < factor; ++i) { | 634 | for (uint16_t i = 0; i < factor; ++i) { |
| 635 | - pregOut = MicroAPI::UpdateMask<uint32_t>(sreg); | 635 | + pregOut = Reg::UpdateMask<uint32_t>(sreg); |
| 636 | LoadE2B<float>(tmpVreg, tmpUb1 + i * factorRow); | 636 | LoadE2B<float>(tmpVreg, tmpUb1 + i * factorRow); |
| 637 | StoreIfNeedCast<T2>(sumUb + i * blockStride * factorRow, tmpVreg, pregOut); | 637 | StoreIfNeedCast<T2>(sumUb + i * blockStride * factorRow, tmpVreg, pregOut); |
| 638 | } | 638 | } |
| @@ -641,14 +641,14 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SingleSoftMaxGenericNDForBlkVFImpl | |||
| 641 | if constexpr (!isFlashV2) { | 641 | if constexpr (!isFlashV2) { |
| 642 | sreg = srcK * srcM; | 642 | sreg = srcK * srcM; |
| 643 | for (uint16_t i = 0; i < factor; ++i) { | 643 | for (uint16_t i = 0; i < factor; ++i) { |
| 644 | - pregOut = MicroAPI::UpdateMask<uint32_t>(sreg); | 644 | + pregOut = Reg::UpdateMask<uint32_t>(sreg); |
| 645 | LoadE2B<float>(sumVreg, tmpUb1 + i * factorRow); | 645 | LoadE2B<float>(sumVreg, tmpUb1 + i * factorRow); |
| 646 | StoreIfNeedCast<T2>(sumUb + i * blockStride * factorRow, sumVreg, pregOut); | 646 | StoreIfNeedCast<T2>(sumUb + i * blockStride * factorRow, sumVreg, pregOut); |
| 647 | - MicroAPI::LoadAlign(tmpVreg, workUb + i * srcK * factorRow); | 647 | + Reg::LoadAlign(tmpVreg, workUb + i * srcK * factorRow); |
| 648 | - MicroAPI::MaskAnd(pregDst, pregCnt, pregOut, pregFull); | 648 | + Reg::MaskAnd(pregDst, pregCnt, pregOut, pregFull); |
| 649 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, pregDst); | 649 | + Reg::Div(dstVreg, tmpVreg, sumVreg, pregDst); |
| 650 | if constexpr (isLog) { | 650 | if constexpr (isLog) { |
| 651 | - MicroAPI::Log10(dstVreg, dstVreg, pregDst); | 651 | + Reg::Log10(dstVreg, dstVreg, pregDst); |
| 652 | } | 652 | } |
| 653 | StoreIfNeedCast<T1>(dstUb + i * srcK * factorRow, dstVreg, pregDst); | 653 | StoreIfNeedCast<T1>(dstUb + i * srcK * factorRow, dstVreg, pregDst); |
| 654 | } | 654 | } |
| @@ -694,57 +694,57 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SingleSoftMaxGenericNDAlignedWithB | |||
| 694 | __ubuf__ float* tmpUb0Tmp0 = tmpUb0; | 694 | __ubuf__ float* tmpUb0Tmp0 = tmpUb0; |
| 695 | __ubuf__ float* tmpUb0Tmp1 = tmpUb0; | 695 | __ubuf__ float* tmpUb0Tmp1 = tmpUb0; |
| 696 | 696 | ||
| 697 | - MicroAPI::MaskReg pregDst; | 697 | + Reg::MaskReg pregDst; |
| 698 | - MicroAPI::MaskReg pregTmp; | 698 | + Reg::MaskReg pregTmp; |
| 699 | - MicroAPI::MaskReg pregOut; | 699 | + Reg::MaskReg pregOut; |
| 700 | - MicroAPI::MaskReg pregCnt = MicroAPI::MoveMask<uint32_t>(); | 700 | + Reg::MaskReg pregCnt = Reg::MoveMask<uint32_t>(); |
| 701 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 701 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 702 | - MicroAPI::MaskReg pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 702 | + Reg::MaskReg pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 703 | - MicroAPI::RegTensor<float> srcVreg; | 703 | + Reg::RegTensor<float> srcVreg; |
| 704 | - MicroAPI::RegTensor<float> maxVreg; | 704 | + Reg::RegTensor<float> maxVreg; |
| 705 | - MicroAPI::RegTensor<float> sumVreg; | 705 | + Reg::RegTensor<float> sumVreg; |
| 706 | - MicroAPI::RegTensor<float> tmpVreg; | 706 | + Reg::RegTensor<float> tmpVreg; |
| 707 | - MicroAPI::RegTensor<float> dstVreg; | 707 | + Reg::RegTensor<float> dstVreg; |
| 708 | - MicroAPI::UnalignReg ureg0; | 708 | + Reg::UnalignReg ureg0; |
| 709 | - MicroAPI::UnalignReg ureg1; | 709 | + Reg::UnalignReg ureg1; |
| 710 | - MicroAPI::UnalignReg ureg2; | 710 | + Reg::UnalignReg ureg2; |
| 711 | 711 | ||
| 712 | for (uint16_t i = 0; i < factor; ++i) { | 712 | for (uint16_t i = 0; i < factor; ++i) { |
| 713 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK * factorRow, pregFull); | 713 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK * factorRow, pregFull); |
| 714 | 714 | ||
| 715 | - MicroAPI::ReduceMaxWithDataBlock(maxVreg, srcVreg, pregCnt); | 715 | + Reg::ReduceMaxWithDataBlock(maxVreg, srcVreg, pregCnt); |
| 716 | 716 | ||
| 717 | Duplicate(tmpVreg, 0); | 717 | Duplicate(tmpVreg, 0); |
| 718 | - MicroAPI::DeInterleave(maxVreg, tmpVreg, maxVreg, tmpVreg); | 718 | + Reg::DeInterleave(maxVreg, tmpVreg, maxVreg, tmpVreg); |
| 719 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, pregFull); | 719 | + Reg::Max(maxVreg, maxVreg, tmpVreg, pregFull); |
| 720 | if constexpr (!outputBrc) { | 720 | if constexpr (!outputBrc) { |
| 721 | if constexpr (SupportType<T2, half>()) { | 721 | if constexpr (SupportType<T2, half>()) { |
| 722 | - MicroAPI::RegTensor<T2> castVreg; | 722 | + Reg::RegTensor<T2> castVreg; |
| 723 | - MicroAPI::Cast<T2, float, Internal::castTraitB32ToB16>(castVreg, maxVreg, pregOneBlk); | 723 | + Reg::Cast<T2, float, Internal::castTraitB32ToB16>(castVreg, maxVreg, pregOneBlk); |
| 724 | - MicroAPI::Pack<uint16_t, uint32_t>( | 724 | + Reg::Pack<uint16_t, uint32_t>( |
| 725 | - (MicroAPI::RegTensor<uint16_t>&)castVreg, (MicroAPI::RegTensor<uint32_t>&)castVreg); | 725 | + (Reg::RegTensor<uint16_t>&)castVreg, (Reg::RegTensor<uint32_t>&)castVreg); |
| 726 | - MicroAPI::StoreUnAlign(maxUb, castVreg, ureg2, factorRow); | 726 | + Reg::StoreUnAlign(maxUb, castVreg, ureg2, factorRow); |
| 727 | } else { | 727 | } else { |
| 728 | - MicroAPI::StoreUnAlign(maxUb, maxVreg, ureg2, factorRow); | 728 | + Reg::StoreUnAlign(maxUb, maxVreg, ureg2, factorRow); |
| 729 | } | 729 | } |
| 730 | } | 730 | } |
| 731 | if constexpr (sizeof(T2) == sizeof(float) && outputBrc) { | 731 | if constexpr (sizeof(T2) == sizeof(float) && outputBrc) { |
| 732 | - MicroAPI::StoreUnAlign(tmpUb0Tmp0, maxVreg, ureg0, factorRow); | 732 | + Reg::StoreUnAlign(tmpUb0Tmp0, maxVreg, ureg0, factorRow); |
| 733 | } | 733 | } |
| 734 | - MicroAPI::Interleave(maxVreg, tmpVreg, maxVreg, maxVreg); | 734 | + Reg::Interleave(maxVreg, tmpVreg, maxVreg, maxVreg); |
| 735 | - MicroAPI::StoreAlign(tmpUb1 + i * 2 * factorRow, maxVreg, pregOneBlk); | 735 | + Reg::StoreAlign(tmpUb1 + i * 2 * factorRow, maxVreg, pregOneBlk); |
| 736 | } | 736 | } |
| 737 | if constexpr (sizeof(T2) == sizeof(float) && outputBrc) { | 737 | if constexpr (sizeof(T2) == sizeof(float) && outputBrc) { |
| 738 | - MicroAPI::StoreUnAlignPost(tmpUb0Tmp0, ureg0, 0); | 738 | + Reg::StoreUnAlignPost(tmpUb0Tmp0, ureg0, 0); |
| 739 | } else if constexpr (!outputBrc) { | 739 | } else if constexpr (!outputBrc) { |
| 740 | - MicroAPI::StoreUnAlignPost(maxUb, ureg2, 0); | 740 | + Reg::StoreUnAlignPost(maxUb, ureg2, 0); |
| 741 | } | 741 | } |
| 742 | 742 | ||
| 743 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 743 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 744 | 744 | ||
| 745 | if constexpr (sizeof(T2) == sizeof(float) && outputBrc) { | 745 | if constexpr (sizeof(T2) == sizeof(float) && outputBrc) { |
| 746 | for (uint16_t i = 0; i < halfFactor; ++i) { | 746 | for (uint16_t i = 0; i < halfFactor; ++i) { |
| 747 | - pregTmp = MicroAPI::UpdateMask<uint32_t>(sreg1); | 747 | + pregTmp = Reg::UpdateMask<uint32_t>(sreg1); |
| 748 | LoadE2B<float>(tmpVreg, tmpUb0 + i * DEFAULT_BLK_NUM); | 748 | LoadE2B<float>(tmpVreg, tmpUb0 + i * DEFAULT_BLK_NUM); |
| 749 | StoreIfNeedCast<T2>(maxUb + i * blockStride * factorRow * 2, tmpVreg, pregTmp); | 749 | StoreIfNeedCast<T2>(maxUb + i * blockStride * factorRow * 2, tmpVreg, pregTmp); |
| 750 | } | 750 | } |
| @@ -752,64 +752,64 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SingleSoftMaxGenericNDAlignedWithB | |||
| 752 | 752 | ||
| 753 | sreg = srcK * srcM; | 753 | sreg = srcK * srcM; |
| 754 | for (uint16_t i = 0; i < factor; ++i) { | 754 | for (uint16_t i = 0; i < factor; ++i) { |
| 755 | - pregOut = MicroAPI::UpdateMask<uint32_t>(sreg); | 755 | + pregOut = Reg::UpdateMask<uint32_t>(sreg); |
| 756 | LoadE2B<float>(maxVreg, tmpUb1 + i * DEFAULT_BLK_NUM); | 756 | LoadE2B<float>(maxVreg, tmpUb1 + i * DEFAULT_BLK_NUM); |
| 757 | if constexpr (sizeof(T2) == sizeof(half) && outputBrc) { | 757 | if constexpr (sizeof(T2) == sizeof(half) && outputBrc) { |
| 758 | StoreIfNeedCast<T2>(maxUb + i * blockStride * factorRow, maxVreg, pregOut); | 758 | StoreIfNeedCast<T2>(maxUb + i * blockStride * factorRow, maxVreg, pregOut); |
| 759 | } | 759 | } |
| 760 | 760 | ||
| 761 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK * factorRow, pregFull); | 761 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK * factorRow, pregFull); |
| 762 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, pregFull); | 762 | + Reg::Sub(dstVreg, srcVreg, maxVreg, pregFull); |
| 763 | - MicroAPI::Exp(tmpVreg, dstVreg, pregFull); | 763 | + Reg::Exp(tmpVreg, dstVreg, pregFull); |
| 764 | if constexpr (!isFlashV2) { | 764 | if constexpr (!isFlashV2) { |
| 765 | - MicroAPI::StoreAlign(workUb + i * srcK * factorRow, tmpVreg, pregOut); | 765 | + Reg::StoreAlign(workUb + i * srcK * factorRow, tmpVreg, pregOut); |
| 766 | } else { | 766 | } else { |
| 767 | - MicroAPI::MaskAnd(pregDst, pregCnt, pregOut, pregFull); | 767 | + Reg::MaskAnd(pregDst, pregCnt, pregOut, pregFull); |
| 768 | StoreIfNeedCast<T1>(dstUb + i * srcK * factorRow, tmpVreg, pregDst); | 768 | StoreIfNeedCast<T1>(dstUb + i * srcK * factorRow, tmpVreg, pregDst); |
| 769 | } | 769 | } |
| 770 | 770 | ||
| 771 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, tmpVreg, pregCnt); | 771 | + Reg::ReduceSumWithDataBlock(sumVreg, tmpVreg, pregCnt); |
| 772 | 772 | ||
| 773 | Duplicate(tmpVreg, 0); | 773 | Duplicate(tmpVreg, 0); |
| 774 | - MicroAPI::DeInterleave(sumVreg, tmpVreg, sumVreg, tmpVreg); | 774 | + Reg::DeInterleave(sumVreg, tmpVreg, sumVreg, tmpVreg); |
| 775 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 775 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 776 | if constexpr (!outputBrc) { | 776 | if constexpr (!outputBrc) { |
| 777 | if constexpr (SupportType<T2, half>()) { | 777 | if constexpr (SupportType<T2, half>()) { |
| 778 | - MicroAPI::RegTensor<T2> castVreg; | 778 | + Reg::RegTensor<T2> castVreg; |
| 779 | - MicroAPI::Cast<T2, float, Internal::castTraitB32ToB16>(castVreg, sumVreg, pregOneBlk); | 779 | + Reg::Cast<T2, float, Internal::castTraitB32ToB16>(castVreg, sumVreg, pregOneBlk); |
| 780 | - MicroAPI::Pack<uint16_t, uint32_t>( | 780 | + Reg::Pack<uint16_t, uint32_t>( |
| 781 | - (MicroAPI::RegTensor<uint16_t>&)castVreg, (MicroAPI::RegTensor<uint32_t>&)castVreg); | 781 | + (Reg::RegTensor<uint16_t>&)castVreg, (Reg::RegTensor<uint32_t>&)castVreg); |
| 782 | - MicroAPI::StoreUnAlign(sumUb, castVreg, ureg2, factorRow); | 782 | + Reg::StoreUnAlign(sumUb, castVreg, ureg2, factorRow); |
| 783 | } else { | 783 | } else { |
| 784 | - MicroAPI::StoreUnAlign(sumUb, sumVreg, ureg2, factorRow); | 784 | + Reg::StoreUnAlign(sumUb, sumVreg, ureg2, factorRow); |
| 785 | } | 785 | } |
| 786 | } | 786 | } |
| 787 | if constexpr (sizeof(T2) == sizeof(float) && outputBrc) { | 787 | if constexpr (sizeof(T2) == sizeof(float) && outputBrc) { |
| 788 | - MicroAPI::StoreUnAlign(tmpUb0Tmp1, sumVreg, ureg1, factorRow); | 788 | + Reg::StoreUnAlign(tmpUb0Tmp1, sumVreg, ureg1, factorRow); |
| 789 | } | 789 | } |
| 790 | - MicroAPI::Interleave(sumVreg, tmpVreg, sumVreg, sumVreg); | 790 | + Reg::Interleave(sumVreg, tmpVreg, sumVreg, sumVreg); |
| 791 | - MicroAPI::StoreAlign(tmpUb + i * 2 * factorRow, sumVreg, pregOneBlk); | 791 | + Reg::StoreAlign(tmpUb + i * 2 * factorRow, sumVreg, pregOneBlk); |
| 792 | } | 792 | } |
| 793 | if constexpr (sizeof(T2) == sizeof(float) && outputBrc) { | 793 | if constexpr (sizeof(T2) == sizeof(float) && outputBrc) { |
| 794 | - MicroAPI::StoreUnAlignPost(tmpUb0Tmp1, ureg1, 0); | 794 | + Reg::StoreUnAlignPost(tmpUb0Tmp1, ureg1, 0); |
| 795 | } else if (!outputBrc) { | 795 | } else if (!outputBrc) { |
| 796 | - MicroAPI::StoreUnAlignPost(sumUb, ureg2, 0); | 796 | + Reg::StoreUnAlignPost(sumUb, ureg2, 0); |
| 797 | } | 797 | } |
| 798 | 798 | ||
| 799 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 799 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 800 | 800 | ||
| 801 | if constexpr (outputBrc) { | 801 | if constexpr (outputBrc) { |
| 802 | if constexpr (sizeof(T2) == sizeof(float)) { | 802 | if constexpr (sizeof(T2) == sizeof(float)) { |
| 803 | sreg1 = srcM * blockStride; | 803 | sreg1 = srcM * blockStride; |
| 804 | for (uint16_t i = 0; i < halfFactor; ++i) { | 804 | for (uint16_t i = 0; i < halfFactor; ++i) { |
| 805 | - pregTmp = MicroAPI::UpdateMask<uint32_t>(sreg1); | 805 | + pregTmp = Reg::UpdateMask<uint32_t>(sreg1); |
| 806 | LoadE2B<float>(tmpVreg, tmpUb0 + i * DEFAULT_BLK_NUM); | 806 | LoadE2B<float>(tmpVreg, tmpUb0 + i * DEFAULT_BLK_NUM); |
| 807 | StoreIfNeedCast<T2>(sumUb + i * blockStride * factorRow * 2, tmpVreg, pregTmp); | 807 | StoreIfNeedCast<T2>(sumUb + i * blockStride * factorRow * 2, tmpVreg, pregTmp); |
| 808 | } | 808 | } |
| 809 | } else if constexpr (sizeof(T2) == sizeof(half)) { | 809 | } else if constexpr (sizeof(T2) == sizeof(half)) { |
| 810 | sreg = srcM * blockStride; | 810 | sreg = srcM * blockStride; |
| 811 | for (uint16_t i = 0; i < factor; ++i) { | 811 | for (uint16_t i = 0; i < factor; ++i) { |
| 812 | - pregOut = MicroAPI::UpdateMask<uint32_t>(sreg); | 812 | + pregOut = Reg::UpdateMask<uint32_t>(sreg); |
| 813 | LoadE2B<float>(tmpVreg, tmpUb + i * DEFAULT_BLK_NUM); | 813 | LoadE2B<float>(tmpVreg, tmpUb + i * DEFAULT_BLK_NUM); |
| 814 | StoreIfNeedCast<T2>(sumUb + i * blockStride * factorRow, tmpVreg, pregOut); | 814 | StoreIfNeedCast<T2>(sumUb + i * blockStride * factorRow, tmpVreg, pregOut); |
| 815 | } | 815 | } |
| @@ -819,13 +819,13 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SingleSoftMaxGenericNDAlignedWithB | |||
| 819 | if constexpr (!isFlashV2) { | 819 | if constexpr (!isFlashV2) { |
| 820 | sreg = srcK * srcM; | 820 | sreg = srcK * srcM; |
| 821 | for (uint16_t i = 0; i < factor; ++i) { | 821 | for (uint16_t i = 0; i < factor; ++i) { |
| 822 | - pregOut = MicroAPI::UpdateMask<uint32_t>(sreg); | 822 | + pregOut = Reg::UpdateMask<uint32_t>(sreg); |
| 823 | LoadE2B<float>(sumVreg, tmpUb + i * DEFAULT_BLK_NUM); | 823 | LoadE2B<float>(sumVreg, tmpUb + i * DEFAULT_BLK_NUM); |
| 824 | - MicroAPI::LoadAlign(tmpVreg, workUb + i * srcK * factorRow); | 824 | + Reg::LoadAlign(tmpVreg, workUb + i * srcK * factorRow); |
| 825 | - MicroAPI::MaskAnd(pregDst, pregCnt, pregOut, pregFull); | 825 | + Reg::MaskAnd(pregDst, pregCnt, pregOut, pregFull); |
| 826 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, pregDst); | 826 | + Reg::Div(dstVreg, tmpVreg, sumVreg, pregDst); |
| 827 | if constexpr (isLog) { | 827 | if constexpr (isLog) { |
| 828 | - MicroAPI::Log10(dstVreg, dstVreg, pregDst); | 828 | + Reg::Log10(dstVreg, dstVreg, pregDst); |
| 829 | } | 829 | } |
| 830 | StoreIfNeedCast<T1>(dstUb + i * srcK * factorRow, dstVreg, pregDst); | 830 | StoreIfNeedCast<T1>(dstUb + i * srcK * factorRow, dstVreg, pregDst); |
| 831 | } | 831 | } |
| @@ -872,25 +872,25 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SingleSoftMaxGenericNDVFImpl(__ubu | |||
| 872 | NotNumUnion notNum; | 872 | NotNumUnion notNum; |
| 873 | notNum.i = F32_NEG_INF; | 873 | notNum.i = F32_NEG_INF; |
| 874 | 874 | ||
| 875 | - MicroAPI::MaskReg pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 875 | + Reg::MaskReg pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 876 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 876 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 877 | - MicroAPI::MaskReg pregOnePt = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL1>(); | 877 | + Reg::MaskReg pregOnePt = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL1>(); |
| 878 | - MicroAPI::MaskReg pregOneBlk; | 878 | + Reg::MaskReg pregOneBlk; |
| 879 | if constexpr (IsSameType<T2, half>::value) { | 879 | if constexpr (IsSameType<T2, half>::value) { |
| 880 | - pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL16>(); | 880 | + pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL16>(); |
| 881 | } else { | 881 | } else { |
| 882 | - pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 882 | + pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 883 | } | 883 | } |
| 884 | - MicroAPI::RegTensor<float> srcVreg; | 884 | + Reg::RegTensor<float> srcVreg; |
| 885 | - MicroAPI::RegTensor<float> maxVreg; | 885 | + Reg::RegTensor<float> maxVreg; |
| 886 | - MicroAPI::RegTensor<float> sumVreg; | 886 | + Reg::RegTensor<float> sumVreg; |
| 887 | - MicroAPI::RegTensor<float> tmpVreg; | 887 | + Reg::RegTensor<float> tmpVreg; |
| 888 | - MicroAPI::RegTensor<float> dstVreg; | 888 | + Reg::RegTensor<float> dstVreg; |
| 889 | 889 | ||
| 890 | for (uint16_t i = 0; i < srcM; ++i) { | 890 | for (uint16_t i = 0; i < srcM; ++i) { |
| 891 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK, pregFull); | 891 | LoadIfNeedCast<T1>(srcVreg, srcUb + i * srcK, pregFull); |
| 892 | 892 | ||
| 893 | - MicroAPI::ReduceMax(maxVreg, srcVreg, pregCnt); | 893 | + Reg::ReduceMax(maxVreg, srcVreg, pregCnt); |
| 894 | if constexpr (!outputBrc) { | 894 | if constexpr (!outputBrc) { |
| 895 | StoreIfNeedCastM1<T2>(maxUb + i, maxVreg, pregOnePt); | 895 | StoreIfNeedCastM1<T2>(maxUb + i, maxVreg, pregOnePt); |
| 896 | } else { | 896 | } else { |
| @@ -899,14 +899,14 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SingleSoftMaxGenericNDVFImpl(__ubu | |||
| 899 | } | 899 | } |
| 900 | 900 | ||
| 901 | Duplicate(maxVreg, maxVreg, pregFull); | 901 | Duplicate(maxVreg, maxVreg, pregFull); |
| 902 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, pregCnt); | 902 | + Reg::Sub(dstVreg, srcVreg, maxVreg, pregCnt); |
| 903 | - MicroAPI::Exp(tmpVreg, dstVreg, pregCnt); | 903 | + Reg::Exp(tmpVreg, dstVreg, pregCnt); |
| 904 | if constexpr (!isFlashV2) { | 904 | if constexpr (!isFlashV2) { |
| 905 | - MicroAPI::StoreAlign(workUb + i * srcK, tmpVreg, pregCnt); | 905 | + Reg::StoreAlign(workUb + i * srcK, tmpVreg, pregCnt); |
| 906 | } else { | 906 | } else { |
| 907 | StoreIfNeedCast<T1>(dstUb + i * srcK, tmpVreg, pregCnt); | 907 | StoreIfNeedCast<T1>(dstUb + i * srcK, tmpVreg, pregCnt); |
| 908 | } | 908 | } |
| 909 | - MicroAPI::ReduceSum(sumVreg, tmpVreg, pregCnt); | 909 | + Reg::ReduceSum(sumVreg, tmpVreg, pregCnt); |
| 910 | if constexpr (!outputBrc) { | 910 | if constexpr (!outputBrc) { |
| 911 | StoreIfNeedCastM1<T2>(sumUb + i, sumVreg, pregOnePt); | 911 | StoreIfNeedCastM1<T2>(sumUb + i, sumVreg, pregOnePt); |
| 912 | } else { | 912 | } else { |
| @@ -914,23 +914,23 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SingleSoftMaxGenericNDVFImpl(__ubu | |||
| 914 | StoreIfNeedCast<T2>(sumUb + i * blockStride, sumVreg, pregOneBlk); | 914 | StoreIfNeedCast<T2>(sumUb + i * blockStride, sumVreg, pregOneBlk); |
| 915 | } | 915 | } |
| 916 | if constexpr (!isFlashV2 && sizeof(T2) == sizeof(half)) { | 916 | if constexpr (!isFlashV2 && sizeof(T2) == sizeof(half)) { |
| 917 | - MicroAPI::StoreAlign(tmpUb + i * blockStride, sumVreg, pregOneBlk); | 917 | + Reg::StoreAlign(tmpUb + i * blockStride, sumVreg, pregOneBlk); |
| 918 | } | 918 | } |
| 919 | } | 919 | } |
| 920 | 920 | ||
| 921 | if constexpr (!isFlashV2) { | 921 | if constexpr (!isFlashV2) { |
| 922 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 922 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 923 | for (uint16_t i = 0; i < srcM; ++i) { | 923 | for (uint16_t i = 0; i < srcM; ++i) { |
| 924 | if constexpr (sizeof(T2) == sizeof(half)) { | 924 | if constexpr (sizeof(T2) == sizeof(half)) { |
| 925 | - MicroAPI::LoadAlign(sumVreg, tmpUb + i * blockStride); | 925 | + Reg::LoadAlign(sumVreg, tmpUb + i * blockStride); |
| 926 | } else { | 926 | } else { |
| 927 | - MicroAPI::LoadAlign(sumVreg, sumUb + i * blockStride); | 927 | + Reg::LoadAlign(sumVreg, sumUb + i * blockStride); |
| 928 | } | 928 | } |
| 929 | Duplicate(sumVreg, sumVreg, pregFull); | 929 | Duplicate(sumVreg, sumVreg, pregFull); |
| 930 | - MicroAPI::LoadAlign(tmpVreg, workUb + i * srcK); | 930 | + Reg::LoadAlign(tmpVreg, workUb + i * srcK); |
| 931 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, pregCnt); | 931 | + Reg::Div(dstVreg, tmpVreg, sumVreg, pregCnt); |
| 932 | if constexpr (isLog) { | 932 | if constexpr (isLog) { |
| 933 | - MicroAPI::Log10(dstVreg, dstVreg, pregCnt); | 933 | + Reg::Log10(dstVreg, dstVreg, pregCnt); |
| 934 | } | 934 | } |
| 935 | StoreIfNeedCast<T1>(dstUb + i * srcK, dstVreg, pregCnt); | 935 | StoreIfNeedCast<T1>(dstUb + i * srcK, dstVreg, pregCnt); |
| 936 | } | 936 | } |
| @@ -994,57 +994,57 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SoftMaxGenericNDWithTailVFImpl(__u | |||
| 994 | notNum.i = F32_NEG_INF; | 994 | notNum.i = F32_NEG_INF; |
| 995 | __ubuf__ float* tmpUb = sumUb; | 995 | __ubuf__ float* tmpUb = sumUb; |
| 996 | 996 | ||
| 997 | - MicroAPI::MaskReg pregCnt; | 997 | + Reg::MaskReg pregCnt; |
| 998 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 998 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 999 | - MicroAPI::RegTensor<float> srcVreg; | 999 | + Reg::RegTensor<float> srcVreg; |
| 1000 | - MicroAPI::RegTensor<float> maxVreg; | 1000 | + Reg::RegTensor<float> maxVreg; |
| 1001 | - MicroAPI::RegTensor<float> sumVreg; | 1001 | + Reg::RegTensor<float> sumVreg; |
| 1002 | - MicroAPI::RegTensor<float> minVreg; | 1002 | + Reg::RegTensor<float> minVreg; |
| 1003 | - MicroAPI::RegTensor<float> tmpVreg; | 1003 | + Reg::RegTensor<float> tmpVreg; |
| 1004 | - MicroAPI::RegTensor<float> dstVreg; | 1004 | + Reg::RegTensor<float> dstVreg; |
| 1005 | - MicroAPI::UnalignReg ureg0; | 1005 | + Reg::UnalignReg ureg0; |
| 1006 | 1006 | ||
| 1007 | Duplicate(minVreg, notNum.f); | 1007 | Duplicate(minVreg, notNum.f); |
| 1008 | for (uint16_t i = 0; i < srcM; ++i) { | 1008 | for (uint16_t i = 0; i < srcM; ++i) { |
| 1009 | uint32_t sreg = originK; | 1009 | uint32_t sreg = originK; |
| 1010 | Duplicate(maxVreg, notNum.f); | 1010 | Duplicate(maxVreg, notNum.f); |
| 1011 | for (uint16_t j = 0; j < static_cast<uint16_t>(repeatTimes - 1); ++j) { | 1011 | for (uint16_t j = 0; j < static_cast<uint16_t>(repeatTimes - 1); ++j) { |
| 1012 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 1012 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 1013 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 1013 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 1014 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregCnt); | 1014 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregCnt); |
| 1015 | } | 1015 | } |
| 1016 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 1016 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 1017 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + (repeatTimes - 1) * FLOAT_REPEAT_SIZE, pregFull); | 1017 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + (repeatTimes - 1) * FLOAT_REPEAT_SIZE, pregFull); |
| 1018 | - MicroAPI::Select(srcVreg, srcVreg, minVreg, pregCnt); | 1018 | + Reg::Select(srcVreg, srcVreg, minVreg, pregCnt); |
| 1019 | - MicroAPI::Max(maxVreg, maxVreg, srcVreg, pregFull); | 1019 | + Reg::Max(maxVreg, maxVreg, srcVreg, pregFull); |
| 1020 | 1020 | ||
| 1021 | - MicroAPI::ReduceMax(maxVreg, maxVreg, pregFull); | 1021 | + Reg::ReduceMax(maxVreg, maxVreg, pregFull); |
| 1022 | 1022 | ||
| 1023 | Duplicate(sumVreg, 0); | 1023 | Duplicate(sumVreg, 0); |
| 1024 | Duplicate(maxVreg, maxVreg, pregFull); | 1024 | Duplicate(maxVreg, maxVreg, pregFull); |
| 1025 | sreg = originK; | 1025 | sreg = originK; |
| 1026 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 1026 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 1027 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 1027 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 1028 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 1028 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 1029 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, pregCnt); | 1029 | + Reg::Sub(dstVreg, srcVreg, maxVreg, pregCnt); |
| 1030 | - MicroAPI::Exp(tmpVreg, dstVreg, pregCnt); | 1030 | + Reg::Exp(tmpVreg, dstVreg, pregCnt); |
| 1031 | - MicroAPI::StoreAlign(workUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); | 1031 | + Reg::StoreAlign(workUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg, pregCnt); |
| 1032 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 1032 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 1033 | } | 1033 | } |
| 1034 | - MicroAPI::ReduceSum(sumVreg, sumVreg, pregFull); | 1034 | + Reg::ReduceSum(sumVreg, sumVreg, pregFull); |
| 1035 | - MicroAPI::StoreUnAlign(sumUb, sumVreg, ureg0, 1); | 1035 | + Reg::StoreUnAlign(sumUb, sumVreg, ureg0, 1); |
| 1036 | } | 1036 | } |
| 1037 | - MicroAPI::StoreUnAlignPost(sumUb, ureg0, 0); | 1037 | + Reg::StoreUnAlignPost(sumUb, ureg0, 0); |
| 1038 | 1038 | ||
| 1039 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 1039 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 1040 | 1040 | ||
| 1041 | for (uint16_t i = 0; i < srcM; ++i) { | 1041 | for (uint16_t i = 0; i < srcM; ++i) { |
| 1042 | - MicroAPI::LoadAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE, MicroAPI::LoadDist::DIST_BRC_B32>(sumVreg, tmpUb, 1); | 1042 | + Reg::LoadAlign<float, Reg::PostLiteral::POST_MODE_UPDATE, Reg::LoadDist::DIST_BRC_B32>(sumVreg, tmpUb, 1); |
| 1043 | uint32_t sreg = originK; | 1043 | uint32_t sreg = originK; |
| 1044 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 1044 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 1045 | - pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 1045 | + pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 1046 | - MicroAPI::LoadAlign(tmpVreg, workUb + i * srcK + j * FLOAT_REPEAT_SIZE); | 1046 | + Reg::LoadAlign(tmpVreg, workUb + i * srcK + j * FLOAT_REPEAT_SIZE); |
| 1047 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, pregCnt); | 1047 | + Reg::Div(dstVreg, tmpVreg, sumVreg, pregCnt); |
| 1048 | StoreIfNeedCast<T>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, dstVreg, pregCnt); | 1048 | StoreIfNeedCast<T>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, dstVreg, pregCnt); |
| 1049 | } | 1049 | } |
| 1050 | } | 1050 | } |
| @@ -1080,19 +1080,19 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SoftMaxGenericNDVFImpl(__ubuf__ T* | |||
| 1080 | NotNumUnion notNum; | 1080 | NotNumUnion notNum; |
| 1081 | notNum.i = F32_NEG_INF; | 1081 | notNum.i = F32_NEG_INF; |
| 1082 | 1082 | ||
| 1083 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 1083 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 1084 | - MicroAPI::MaskReg pregOne = MicroAPI::CreateMask<float, MicroAPI::MaskPattern::VL1>(); | 1084 | + Reg::MaskReg pregOne = Reg::CreateMask<float, Reg::MaskPattern::VL1>(); |
| 1085 | - MicroAPI::RegTensor<float> srcVreg0; | 1085 | + Reg::RegTensor<float> srcVreg0; |
| 1086 | - MicroAPI::RegTensor<float> maxVreg0; | 1086 | + Reg::RegTensor<float> maxVreg0; |
| 1087 | - MicroAPI::RegTensor<float> sumVreg0; | 1087 | + Reg::RegTensor<float> sumVreg0; |
| 1088 | - MicroAPI::RegTensor<float> tmpVreg0; | 1088 | + Reg::RegTensor<float> tmpVreg0; |
| 1089 | - MicroAPI::RegTensor<float> dstVreg0; | 1089 | + Reg::RegTensor<float> dstVreg0; |
| 1090 | 1090 | ||
| 1091 | - MicroAPI::RegTensor<float> srcVreg1; | 1091 | + Reg::RegTensor<float> srcVreg1; |
| 1092 | - MicroAPI::RegTensor<float> maxVreg1; | 1092 | + Reg::RegTensor<float> maxVreg1; |
| 1093 | - MicroAPI::RegTensor<float> sumVreg1; | 1093 | + Reg::RegTensor<float> sumVreg1; |
| 1094 | - MicroAPI::RegTensor<float> tmpVreg1; | 1094 | + Reg::RegTensor<float> tmpVreg1; |
| 1095 | - MicroAPI::RegTensor<float> dstVreg1; | 1095 | + Reg::RegTensor<float> dstVreg1; |
| 1096 | 1096 | ||
| 1097 | for (uint16_t i = 0; i < halfM; ++i) { | 1097 | for (uint16_t i = 0; i < halfM; ++i) { |
| 1098 | Duplicate(maxVreg0, notNum.f); | 1098 | Duplicate(maxVreg0, notNum.f); |
| @@ -1100,77 +1100,77 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SoftMaxGenericNDVFImpl(__ubuf__ T* | |||
| 1100 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 1100 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 1101 | LoadIfNeedCast<T>(srcVreg0, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 1101 | LoadIfNeedCast<T>(srcVreg0, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 1102 | LoadIfNeedCast<T>(srcVreg1, srcUb + (i + halfM) * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 1102 | LoadIfNeedCast<T>(srcVreg1, srcUb + (i + halfM) * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 1103 | - MicroAPI::Max(maxVreg0, maxVreg0, srcVreg0, pregFull); | 1103 | + Reg::Max(maxVreg0, maxVreg0, srcVreg0, pregFull); |
| 1104 | - MicroAPI::Max(maxVreg1, maxVreg1, srcVreg1, pregFull); | 1104 | + Reg::Max(maxVreg1, maxVreg1, srcVreg1, pregFull); |
| 1105 | } | 1105 | } |
| 1106 | - MicroAPI::ReduceMax(maxVreg0, maxVreg0, pregFull); | 1106 | + Reg::ReduceMax(maxVreg0, maxVreg0, pregFull); |
| 1107 | - MicroAPI::ReduceMax(maxVreg1, maxVreg1, pregFull); | 1107 | + Reg::ReduceMax(maxVreg1, maxVreg1, pregFull); |
| 1108 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>((maxUb + i), maxVreg0, pregOne); | 1108 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_FIRST_ELEMENT_B32>((maxUb + i), maxVreg0, pregOne); |
| 1109 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>((maxUb + i + halfM), maxVreg1, pregOne); | 1109 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_FIRST_ELEMENT_B32>((maxUb + i + halfM), maxVreg1, pregOne); |
| 1110 | } | 1110 | } |
| 1111 | for (uint16_t i = 0; i < tailM; ++i) { | 1111 | for (uint16_t i = 0; i < tailM; ++i) { |
| 1112 | Duplicate(maxVreg0, notNum.f); | 1112 | Duplicate(maxVreg0, notNum.f); |
| 1113 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 1113 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 1114 | LoadIfNeedCast<T>(srcVreg0, srcUb + mainM * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 1114 | LoadIfNeedCast<T>(srcVreg0, srcUb + mainM * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 1115 | - MicroAPI::Max(maxVreg0, maxVreg0, srcVreg0, pregFull); | 1115 | + Reg::Max(maxVreg0, maxVreg0, srcVreg0, pregFull); |
| 1116 | } | 1116 | } |
| 1117 | - MicroAPI::ReduceMax(maxVreg0, maxVreg0, pregFull); | 1117 | + Reg::ReduceMax(maxVreg0, maxVreg0, pregFull); |
| 1118 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>((maxUb + mainM), maxVreg0, pregOne); | 1118 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_FIRST_ELEMENT_B32>((maxUb + mainM), maxVreg0, pregOne); |
| 1119 | } | 1119 | } |
| 1120 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 1120 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 1121 | for (uint16_t i = 0; i < halfM; i++) { | 1121 | for (uint16_t i = 0; i < halfM; i++) { |
| 1122 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(maxVreg0, maxUb + i); | 1122 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_BRC_B32>(maxVreg0, maxUb + i); |
| 1123 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(maxVreg1, maxUb + (i + halfM)); | 1123 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_BRC_B32>(maxVreg1, maxUb + (i + halfM)); |
| 1124 | Duplicate(sumVreg0, 0); | 1124 | Duplicate(sumVreg0, 0); |
| 1125 | Duplicate(sumVreg1, 0); | 1125 | Duplicate(sumVreg1, 0); |
| 1126 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 1126 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 1127 | LoadIfNeedCast<T>(srcVreg0, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 1127 | LoadIfNeedCast<T>(srcVreg0, srcUb + i * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 1128 | LoadIfNeedCast<T>(srcVreg1, srcUb + (i + halfM) * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 1128 | LoadIfNeedCast<T>(srcVreg1, srcUb + (i + halfM) * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 1129 | - MicroAPI::FusedExpSub(tmpVreg0, srcVreg0, maxVreg0, pregFull); | 1129 | + Reg::FusedExpSub(tmpVreg0, srcVreg0, maxVreg0, pregFull); |
| 1130 | - MicroAPI::FusedExpSub(tmpVreg1, srcVreg1, maxVreg1, pregFull); | 1130 | + Reg::FusedExpSub(tmpVreg1, srcVreg1, maxVreg1, pregFull); |
| 1131 | - MicroAPI::StoreAlign(workUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg0, pregFull); | 1131 | + Reg::StoreAlign(workUb + i * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg0, pregFull); |
| 1132 | - MicroAPI::StoreAlign(workUb + (i + halfM) * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg1, pregFull); | 1132 | + Reg::StoreAlign(workUb + (i + halfM) * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg1, pregFull); |
| 1133 | - MicroAPI::Add(sumVreg0, sumVreg0, tmpVreg0, pregFull); | 1133 | + Reg::Add(sumVreg0, sumVreg0, tmpVreg0, pregFull); |
| 1134 | - MicroAPI::Add(sumVreg1, sumVreg1, tmpVreg1, pregFull); | 1134 | + Reg::Add(sumVreg1, sumVreg1, tmpVreg1, pregFull); |
| 1135 | } | 1135 | } |
| 1136 | - MicroAPI::ReduceSum(sumVreg0, sumVreg0, pregFull); | 1136 | + Reg::ReduceSum(sumVreg0, sumVreg0, pregFull); |
| 1137 | - MicroAPI::ReduceSum(sumVreg1, sumVreg1, pregFull); | 1137 | + Reg::ReduceSum(sumVreg1, sumVreg1, pregFull); |
| 1138 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>((sumUb0 + i), sumVreg0, pregOne); | 1138 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_FIRST_ELEMENT_B32>((sumUb0 + i), sumVreg0, pregOne); |
| 1139 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>((sumUb1 + i), sumVreg1, pregOne); | 1139 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_FIRST_ELEMENT_B32>((sumUb1 + i), sumVreg1, pregOne); |
| 1140 | } | 1140 | } |
| 1141 | 1141 | ||
| 1142 | for (uint16_t i = 0; i < tailM; ++i) { | 1142 | for (uint16_t i = 0; i < tailM; ++i) { |
| 1143 | Duplicate(sumVreg0, 0); | 1143 | Duplicate(sumVreg0, 0); |
| 1144 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(maxVreg0, maxUb + mainM); | 1144 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_BRC_B32>(maxVreg0, maxUb + mainM); |
| 1145 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 1145 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 1146 | LoadIfNeedCast<T>(srcVreg0, srcUb + mainM * srcK + j * FLOAT_REPEAT_SIZE, pregFull); | 1146 | LoadIfNeedCast<T>(srcVreg0, srcUb + mainM * srcK + j * FLOAT_REPEAT_SIZE, pregFull); |
| 1147 | - MicroAPI::FusedExpSub(tmpVreg0, srcVreg0, maxVreg0, pregFull); | 1147 | + Reg::FusedExpSub(tmpVreg0, srcVreg0, maxVreg0, pregFull); |
| 1148 | - MicroAPI::StoreAlign(workUb + mainM * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg0, pregFull); | 1148 | + Reg::StoreAlign(workUb + mainM * srcK + j * FLOAT_REPEAT_SIZE, tmpVreg0, pregFull); |
| 1149 | - MicroAPI::Add(sumVreg0, sumVreg0, tmpVreg0, pregFull); | 1149 | + Reg::Add(sumVreg0, sumVreg0, tmpVreg0, pregFull); |
| 1150 | } | 1150 | } |
| 1151 | - MicroAPI::ReduceSum(sumVreg0, sumVreg0, pregFull); | 1151 | + Reg::ReduceSum(sumVreg0, sumVreg0, pregFull); |
| 1152 | - MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>((sumUb + mainM), sumVreg0, pregOne); | 1152 | + Reg::StoreAlign<float, Reg::StoreDist::DIST_FIRST_ELEMENT_B32>((sumUb + mainM), sumVreg0, pregOne); |
| 1153 | } | 1153 | } |
| 1154 | 1154 | ||
| 1155 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 1155 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 1156 | 1156 | ||
| 1157 | for (uint16_t i = 0; i < halfM; ++i) { | 1157 | for (uint16_t i = 0; i < halfM; ++i) { |
| 1158 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(sumVreg0, sumUb + i); | 1158 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_BRC_B32>(sumVreg0, sumUb + i); |
| 1159 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(sumVreg1, sumUb + (i + halfM)); | 1159 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_BRC_B32>(sumVreg1, sumUb + (i + halfM)); |
| 1160 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 1160 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 1161 | - MicroAPI::LoadAlign(tmpVreg0, workUb + i * srcK + j * FLOAT_REPEAT_SIZE); | 1161 | + Reg::LoadAlign(tmpVreg0, workUb + i * srcK + j * FLOAT_REPEAT_SIZE); |
| 1162 | - MicroAPI::LoadAlign(tmpVreg1, workUb + (i + halfM) * srcK + j * FLOAT_REPEAT_SIZE); | 1162 | + Reg::LoadAlign(tmpVreg1, workUb + (i + halfM) * srcK + j * FLOAT_REPEAT_SIZE); |
| 1163 | - MicroAPI::Div(dstVreg0, tmpVreg0, sumVreg0, pregFull); | 1163 | + Reg::Div(dstVreg0, tmpVreg0, sumVreg0, pregFull); |
| 1164 | - MicroAPI::Div(dstVreg1, tmpVreg1, sumVreg1, pregFull); | 1164 | + Reg::Div(dstVreg1, tmpVreg1, sumVreg1, pregFull); |
| 1165 | StoreIfNeedCast<T>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, dstVreg0, pregFull); | 1165 | StoreIfNeedCast<T>(dstUb + i * srcK + j * FLOAT_REPEAT_SIZE, dstVreg0, pregFull); |
| 1166 | StoreIfNeedCast<T>(dstUb + (i + halfM) * srcK + j * FLOAT_REPEAT_SIZE, dstVreg1, pregFull); | 1166 | StoreIfNeedCast<T>(dstUb + (i + halfM) * srcK + j * FLOAT_REPEAT_SIZE, dstVreg1, pregFull); |
| 1167 | } | 1167 | } |
| 1168 | } | 1168 | } |
| 1169 | for (uint16_t i = 0; i < tailM; ++i) { | 1169 | for (uint16_t i = 0; i < tailM; ++i) { |
| 1170 | - MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(sumVreg0, sumUb + mainM); | 1170 | + Reg::LoadAlign<float, Reg::LoadDist::DIST_BRC_B32>(sumVreg0, sumUb + mainM); |
| 1171 | for (uint16_t j = 0; j < repeatTimes; ++j) { | 1171 | for (uint16_t j = 0; j < repeatTimes; ++j) { |
| 1172 | - MicroAPI::LoadAlign(tmpVreg0, workUb + mainM * srcK + j * FLOAT_REPEAT_SIZE); | 1172 | + Reg::LoadAlign(tmpVreg0, workUb + mainM * srcK + j * FLOAT_REPEAT_SIZE); |
| 1173 | - MicroAPI::Div(dstVreg0, tmpVreg0, sumVreg0, pregFull); | 1173 | + Reg::Div(dstVreg0, tmpVreg0, sumVreg0, pregFull); |
| 1174 | StoreIfNeedCast<T>(dstUb + mainM * srcK + j * FLOAT_REPEAT_SIZE, dstVreg0, pregFull); | 1174 | StoreIfNeedCast<T>(dstUb + mainM * srcK + j * FLOAT_REPEAT_SIZE, dstVreg0, pregFull); |
| 1175 | } | 1175 | } |
| 1176 | } | 1176 | } |
| @@ -1212,45 +1212,45 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SingleSoftMaxGenericNDForBlkVFImpl | |||
| 1212 | uint16_t factorRow, uint16_t factor) | 1212 | uint16_t factorRow, uint16_t factor) |
| 1213 | { | 1213 | { |
| 1214 | uint32_t sreg = srcK * srcM; | 1214 | uint32_t sreg = srcK * srcM; |
| 1215 | - MicroAPI::MaskReg pregOut; | 1215 | + Reg::MaskReg pregOut; |
| 1216 | - MicroAPI::MaskReg pregCnt = MicroAPI::MoveMask<uint32_t>(); | 1216 | + Reg::MaskReg pregCnt = Reg::MoveMask<uint32_t>(); |
| 1217 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 1217 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 1218 | - MicroAPI::MaskReg pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 1218 | + Reg::MaskReg pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 1219 | - MicroAPI::RegTensor<float> srcVreg; | 1219 | + Reg::RegTensor<float> srcVreg; |
| 1220 | - MicroAPI::RegTensor<float> maxVreg; | 1220 | + Reg::RegTensor<float> maxVreg; |
| 1221 | - MicroAPI::RegTensor<float> sumVreg; | 1221 | + Reg::RegTensor<float> sumVreg; |
| 1222 | - MicroAPI::RegTensor<float> tmpVreg; | 1222 | + Reg::RegTensor<float> tmpVreg; |
| 1223 | - MicroAPI::RegTensor<float> dstVreg; | 1223 | + Reg::RegTensor<float> dstVreg; |
| 1224 | 1224 | ||
| 1225 | for (uint16_t i = 0; i < factor; ++i) { | 1225 | for (uint16_t i = 0; i < factor; ++i) { |
| 1226 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK * factorRow, pregFull); | 1226 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK * factorRow, pregFull); |
| 1227 | 1227 | ||
| 1228 | - MicroAPI::ReduceMaxWithDataBlock(maxVreg, srcVreg, pregCnt); | 1228 | + Reg::ReduceMaxWithDataBlock(maxVreg, srcVreg, pregCnt); |
| 1229 | - MicroAPI::StoreAlign(tmpUb + i * factorRow, maxVreg, pregOneBlk); | 1229 | + Reg::StoreAlign(tmpUb + i * factorRow, maxVreg, pregOneBlk); |
| 1230 | } | 1230 | } |
| 1231 | 1231 | ||
| 1232 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 1232 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 1233 | 1233 | ||
| 1234 | for (uint16_t i = 0; i < factor; ++i) { | 1234 | for (uint16_t i = 0; i < factor; ++i) { |
| 1235 | LoadE2B<float>(maxVreg, tmpUb + i * factorRow); | 1235 | LoadE2B<float>(maxVreg, tmpUb + i * factorRow); |
| 1236 | 1236 | ||
| 1237 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK * factorRow, pregFull); | 1237 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK * factorRow, pregFull); |
| 1238 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, pregFull); | 1238 | + Reg::Sub(dstVreg, srcVreg, maxVreg, pregFull); |
| 1239 | - MicroAPI::Exp(tmpVreg, dstVreg, pregFull); | 1239 | + Reg::Exp(tmpVreg, dstVreg, pregFull); |
| 1240 | - MicroAPI::StoreAlign(workUb + i * srcK * factorRow, tmpVreg, pregFull); | 1240 | + Reg::StoreAlign(workUb + i * srcK * factorRow, tmpVreg, pregFull); |
| 1241 | 1241 | ||
| 1242 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, tmpVreg, pregCnt); | 1242 | + Reg::ReduceSumWithDataBlock(sumVreg, tmpVreg, pregCnt); |
| 1243 | - MicroAPI::StoreAlign(sumUb + i * factorRow, sumVreg, pregOneBlk); | 1243 | + Reg::StoreAlign(sumUb + i * factorRow, sumVreg, pregOneBlk); |
| 1244 | } | 1244 | } |
| 1245 | 1245 | ||
| 1246 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 1246 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 1247 | 1247 | ||
| 1248 | for (uint16_t i = 0; i < factor; ++i) { | 1248 | for (uint16_t i = 0; i < factor; ++i) { |
| 1249 | - pregOut = MicroAPI::UpdateMask<uint32_t>(sreg); | 1249 | + pregOut = Reg::UpdateMask<uint32_t>(sreg); |
| 1250 | LoadE2B<float>(sumVreg, sumUb + i * factorRow); | 1250 | LoadE2B<float>(sumVreg, sumUb + i * factorRow); |
| 1251 | - MicroAPI::LoadAlign(tmpVreg, workUb + i * srcK * factorRow); | 1251 | + Reg::LoadAlign(tmpVreg, workUb + i * srcK * factorRow); |
| 1252 | - MicroAPI::MaskAnd(pregOneBlk, pregCnt, pregOut, pregFull); | 1252 | + Reg::MaskAnd(pregOneBlk, pregCnt, pregOut, pregFull); |
| 1253 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, pregOneBlk); | 1253 | + Reg::Div(dstVreg, tmpVreg, sumVreg, pregOneBlk); |
| 1254 | StoreIfNeedCast<T>(dstUb + i * srcK * factorRow, dstVreg, pregOneBlk); | 1254 | StoreIfNeedCast<T>(dstUb + i * srcK * factorRow, dstVreg, pregOneBlk); |
| 1255 | } | 1255 | } |
| 1256 | } | 1256 | } |
| @@ -1287,55 +1287,55 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SingleSoftMaxGenericNDAlignedWithB | |||
| 1287 | uint16_t factorRow, uint16_t factor) | 1287 | uint16_t factorRow, uint16_t factor) |
| 1288 | { | 1288 | { |
| 1289 | uint32_t sreg = srcK * srcM; | 1289 | uint32_t sreg = srcK * srcM; |
| 1290 | - MicroAPI::MaskReg pregOut; | 1290 | + Reg::MaskReg pregOut; |
| 1291 | - MicroAPI::MaskReg pregCnt = MicroAPI::MoveMask<uint32_t>(); | 1291 | + Reg::MaskReg pregCnt = Reg::MoveMask<uint32_t>(); |
| 1292 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 1292 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 1293 | - MicroAPI::MaskReg pregOneBlk = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::VL8>(); | 1293 | + Reg::MaskReg pregOneBlk = Reg::CreateMask<uint32_t, Reg::MaskPattern::VL8>(); |
| 1294 | - MicroAPI::RegTensor<float> srcVreg; | 1294 | + Reg::RegTensor<float> srcVreg; |
| 1295 | - MicroAPI::RegTensor<float> maxVreg; | 1295 | + Reg::RegTensor<float> maxVreg; |
| 1296 | - MicroAPI::RegTensor<float> sumVreg; | 1296 | + Reg::RegTensor<float> sumVreg; |
| 1297 | - MicroAPI::RegTensor<float> tmpVreg; | 1297 | + Reg::RegTensor<float> tmpVreg; |
| 1298 | - MicroAPI::RegTensor<float> dstVreg; | 1298 | + Reg::RegTensor<float> dstVreg; |
| 1299 | 1299 | ||
| 1300 | for (uint16_t i = 0; i < factor; ++i) { | 1300 | for (uint16_t i = 0; i < factor; ++i) { |
| 1301 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK * factorRow, pregFull); | 1301 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK * factorRow, pregFull); |
| 1302 | 1302 | ||
| 1303 | - MicroAPI::ReduceMaxWithDataBlock(maxVreg, srcVreg, pregCnt); | 1303 | + Reg::ReduceMaxWithDataBlock(maxVreg, srcVreg, pregCnt); |
| 1304 | 1304 | ||
| 1305 | Duplicate(tmpVreg, 0); | 1305 | Duplicate(tmpVreg, 0); |
| 1306 | - MicroAPI::DeInterleave(maxVreg, tmpVreg, maxVreg, tmpVreg); | 1306 | + Reg::DeInterleave(maxVreg, tmpVreg, maxVreg, tmpVreg); |
| 1307 | - MicroAPI::Max(maxVreg, maxVreg, tmpVreg, pregFull); | 1307 | + Reg::Max(maxVreg, maxVreg, tmpVreg, pregFull); |
| 1308 | - MicroAPI::Interleave(maxVreg, tmpVreg, maxVreg, maxVreg); | 1308 | + Reg::Interleave(maxVreg, tmpVreg, maxVreg, maxVreg); |
| 1309 | - MicroAPI::StoreAlign(tmpUb + i * 2 * factorRow, maxVreg, pregOneBlk); | 1309 | + Reg::StoreAlign(tmpUb + i * 2 * factorRow, maxVreg, pregOneBlk); |
| 1310 | } | 1310 | } |
| 1311 | 1311 | ||
| 1312 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 1312 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 1313 | 1313 | ||
| 1314 | for (uint16_t i = 0; i < factor; ++i) { | 1314 | for (uint16_t i = 0; i < factor; ++i) { |
| 1315 | LoadE2B<float>(maxVreg, tmpUb + i * DEFAULT_BLK_NUM); | 1315 | LoadE2B<float>(maxVreg, tmpUb + i * DEFAULT_BLK_NUM); |
| 1316 | 1316 | ||
| 1317 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK * factorRow, pregFull); | 1317 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK * factorRow, pregFull); |
| 1318 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, pregFull); | 1318 | + Reg::Sub(dstVreg, srcVreg, maxVreg, pregFull); |
| 1319 | - MicroAPI::Exp(tmpVreg, dstVreg, pregFull); | 1319 | + Reg::Exp(tmpVreg, dstVreg, pregFull); |
| 1320 | - MicroAPI::StoreAlign(workUb + i * srcK * factorRow, tmpVreg, pregFull); | 1320 | + Reg::StoreAlign(workUb + i * srcK * factorRow, tmpVreg, pregFull); |
| 1321 | 1321 | ||
| 1322 | - MicroAPI::ReduceSumWithDataBlock(sumVreg, tmpVreg, pregCnt); | 1322 | + Reg::ReduceSumWithDataBlock(sumVreg, tmpVreg, pregCnt); |
| 1323 | 1323 | ||
| 1324 | Duplicate(tmpVreg, 0); | 1324 | Duplicate(tmpVreg, 0); |
| 1325 | - MicroAPI::DeInterleave(sumVreg, tmpVreg, sumVreg, tmpVreg); | 1325 | + Reg::DeInterleave(sumVreg, tmpVreg, sumVreg, tmpVreg); |
| 1326 | - MicroAPI::Add(sumVreg, sumVreg, tmpVreg, pregFull); | 1326 | + Reg::Add(sumVreg, sumVreg, tmpVreg, pregFull); |
| 1327 | - MicroAPI::Interleave(sumVreg, tmpVreg, sumVreg, sumVreg); | 1327 | + Reg::Interleave(sumVreg, tmpVreg, sumVreg, sumVreg); |
| 1328 | - MicroAPI::StoreAlign(sumUb + i * 2 * factorRow, sumVreg, pregOneBlk); | 1328 | + Reg::StoreAlign(sumUb + i * 2 * factorRow, sumVreg, pregOneBlk); |
| 1329 | } | 1329 | } |
| 1330 | 1330 | ||
| 1331 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 1331 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 1332 | 1332 | ||
| 1333 | for (uint16_t i = 0; i < factor; ++i) { | 1333 | for (uint16_t i = 0; i < factor; ++i) { |
| 1334 | - pregOut = MicroAPI::UpdateMask<uint32_t>(sreg); | 1334 | + pregOut = Reg::UpdateMask<uint32_t>(sreg); |
| 1335 | LoadE2B<float>(sumVreg, sumUb + i * DEFAULT_BLK_NUM); | 1335 | LoadE2B<float>(sumVreg, sumUb + i * DEFAULT_BLK_NUM); |
| 1336 | - MicroAPI::LoadAlign(tmpVreg, workUb + i * srcK * factorRow); | 1336 | + Reg::LoadAlign(tmpVreg, workUb + i * srcK * factorRow); |
| 1337 | - MicroAPI::MaskAnd(pregOneBlk, pregCnt, pregOut, pregFull); | 1337 | + Reg::MaskAnd(pregOneBlk, pregCnt, pregOut, pregFull); |
| 1338 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, pregOneBlk); | 1338 | + Reg::Div(dstVreg, tmpVreg, sumVreg, pregOneBlk); |
| 1339 | StoreIfNeedCast<T>(dstUb + i * srcK * factorRow, dstVreg, pregOneBlk); | 1339 | StoreIfNeedCast<T>(dstUb + i * srcK * factorRow, dstVreg, pregOneBlk); |
| 1340 | } | 1340 | } |
| 1341 | } | 1341 | } |
| @@ -1374,37 +1374,37 @@ __no_simd_vf_fusion__ __simd_vf__ inline void SingleSoftMaxGenericNDVFImpl(__ubu | |||
| 1374 | uint16_t originK) | 1374 | uint16_t originK) |
| 1375 | { | 1375 | { |
| 1376 | uint32_t sreg = originK; | 1376 | uint32_t sreg = originK; |
| 1377 | - MicroAPI::MaskReg pregCnt = MicroAPI::UpdateMask<uint32_t>(sreg); | 1377 | + Reg::MaskReg pregCnt = Reg::UpdateMask<uint32_t>(sreg); |
| 1378 | - MicroAPI::MaskReg pregFull = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>(); | 1378 | + Reg::MaskReg pregFull = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>(); |
| 1379 | - MicroAPI::RegTensor<float> srcVreg; | 1379 | + Reg::RegTensor<float> srcVreg; |
| 1380 | - MicroAPI::RegTensor<float> maxVreg; | 1380 | + Reg::RegTensor<float> maxVreg; |
| 1381 | - MicroAPI::RegTensor<float> sumVreg; | 1381 | + Reg::RegTensor<float> sumVreg; |
| 1382 | - MicroAPI::RegTensor<float> tmpVreg; | 1382 | + Reg::RegTensor<float> tmpVreg; |
| 1383 | - MicroAPI::RegTensor<float> dstVreg; | 1383 | + Reg::RegTensor<float> dstVreg; |
| 1384 | - MicroAPI::UnalignReg ureg0; | 1384 | + Reg::UnalignReg ureg0; |
| 1385 | 1385 | ||
| 1386 | for (uint16_t i = 0; i < srcM; ++i) { | 1386 | for (uint16_t i = 0; i < srcM; ++i) { |
| 1387 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK, pregFull); | 1387 | LoadIfNeedCast<T>(srcVreg, srcUb + i * srcK, pregFull); |
| 1388 | 1388 | ||
| 1389 | - MicroAPI::ReduceMax(maxVreg, srcVreg, pregCnt); | 1389 | + Reg::ReduceMax(maxVreg, srcVreg, pregCnt); |
| 1390 | 1390 | ||
| 1391 | Duplicate(maxVreg, maxVreg, pregFull); | 1391 | Duplicate(maxVreg, maxVreg, pregFull); |
| 1392 | 1392 | ||
| 1393 | - MicroAPI::Sub(dstVreg, srcVreg, maxVreg, pregCnt); | 1393 | + Reg::Sub(dstVreg, srcVreg, maxVreg, pregCnt); |
| 1394 | - MicroAPI::Exp(tmpVreg, dstVreg, pregCnt); | 1394 | + Reg::Exp(tmpVreg, dstVreg, pregCnt); |
| 1395 | - MicroAPI::StoreAlign(workUb + i * srcK, tmpVreg, pregCnt); | 1395 | + Reg::StoreAlign(workUb + i * srcK, tmpVreg, pregCnt); |
| 1396 | 1396 | ||
| 1397 | - MicroAPI::ReduceSum(sumVreg, tmpVreg, pregCnt); | 1397 | + Reg::ReduceSum(sumVreg, tmpVreg, pregCnt); |
| 1398 | - MicroAPI::StoreUnAlign(sumUb, sumVreg, ureg0, 1); | 1398 | + Reg::StoreUnAlign(sumUb, sumVreg, ureg0, 1); |
| 1399 | } | 1399 | } |
| 1400 | - MicroAPI::StoreUnAlignPost(sumUb, ureg0, 0); | 1400 | + Reg::StoreUnAlignPost(sumUb, ureg0, 0); |
| 1401 | 1401 | ||
| 1402 | - MicroAPI::LocalMemBar<MicroAPI::MemType::VEC_STORE, MicroAPI::MemType::VEC_LOAD>(); | 1402 | + Reg::LocalMemBar<Reg::MemType::VEC_STORE, Reg::MemType::VEC_LOAD>(); |
| 1403 | 1403 | ||
| 1404 | for (uint16_t i = 0; i < srcM; ++i) { | 1404 | for (uint16_t i = 0; i < srcM; ++i) { |
| 1405 | - MicroAPI::LoadAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE, MicroAPI::LoadDist::DIST_BRC_B32>(sumVreg, tmpUb, 1); | 1405 | + Reg::LoadAlign<float, Reg::PostLiteral::POST_MODE_UPDATE, Reg::LoadDist::DIST_BRC_B32>(sumVreg, tmpUb, 1); |
| 1406 | - MicroAPI::LoadAlign(tmpVreg, workUb + i * srcK); | 1406 | + Reg::LoadAlign(tmpVreg, workUb + i * srcK); |
| 1407 | - MicroAPI::Div(dstVreg, tmpVreg, sumVreg, pregCnt); | 1407 | + Reg::Div(dstVreg, tmpVreg, sumVreg, pregCnt); |
| 1408 | StoreIfNeedCast<T>(dstUb + i * srcK, dstVreg, pregCnt); | 1408 | StoreIfNeedCast<T>(dstUb + i * srcK, dstVreg, pregCnt); |
| 1409 | } | 1409 | } |
| 1410 | } | 1410 | } |
| @@ -1459,41 +1459,41 @@ __simd_vf__ inline void AdjustSoftMaxResNZImpl(__ubuf__ T1* resUb, __ubuf__ T2* | |||
| 1459 | __ubuf__ uint64_t* maskUb, const uint32_t from, const T1 to, const uint32_t dataBlock, | 1459 | __ubuf__ uint64_t* maskUb, const uint32_t from, const T1 to, const uint32_t dataBlock, |
| 1460 | const uint16_t mRepeatTimes, const uint16_t kRepeatTimes) | 1460 | const uint16_t mRepeatTimes, const uint16_t kRepeatTimes) |
| 1461 | { | 1461 | { |
| 1462 | - MicroAPI::RegTensor<T1> srcVreg; | 1462 | + Reg::RegTensor<T1> srcVreg; |
| 1463 | - MicroAPI::RegTensor<T1> tmpVreg; | 1463 | + Reg::RegTensor<T1> tmpVreg; |
| 1464 | - MicroAPI::RegTensor<T1> dstVreg; | 1464 | + Reg::RegTensor<T1> dstVreg; |
| 1465 | - MicroAPI::RegTensor<T2> maxVreg; | 1465 | + Reg::RegTensor<T2> maxVreg; |
| 1466 | - MicroAPI::MaskReg cmpMaskReg; | 1466 | + Reg::MaskReg cmpMaskReg; |
| 1467 | - MicroAPI::MaskReg cmpMaskReg0; | 1467 | + Reg::MaskReg cmpMaskReg0; |
| 1468 | - MicroAPI::MaskReg cmpMaskReg1; | 1468 | + Reg::MaskReg cmpMaskReg1; |
| 1469 | - MicroAPI::MaskReg maskFull = MicroAPI::CreateMask<T1, MicroAPI::MaskPattern::ALL>(); | 1469 | + Reg::MaskReg maskFull = Reg::CreateMask<T1, Reg::MaskPattern::ALL>(); |
| 1470 | - MicroAPI::MaskReg dstMask = MicroAPI::CreateMask<T1, MicroAPI::MaskPattern::ALLF>(); | 1470 | + Reg::MaskReg dstMask = Reg::CreateMask<T1, Reg::MaskPattern::ALLF>(); |
| 1471 | 1471 | ||
| 1472 | bool isUpdateNeedCheck = false; | 1472 | bool isUpdateNeedCheck = false; |
| 1473 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { | 1473 | for (uint16_t i = 0; i < mRepeatTimes; ++i) { |
| 1474 | if constexpr (sizeof(T2) == sizeof(float)) { | 1474 | if constexpr (sizeof(T2) == sizeof(float)) { |
| 1475 | - MicroAPI::LoadAlign<T2, MicroAPI::LoadDist::DIST_BRC_B32>(maxVreg, maxUb + i * stepSize); | 1475 | + Reg::LoadAlign<T2, Reg::LoadDist::DIST_BRC_B32>(maxVreg, maxUb + i * stepSize); |
| 1476 | // either full mask or zero mask | 1476 | // either full mask or zero mask |
| 1477 | - MicroAPI::CompareScalar(cmpMaskReg, (MicroAPI::RegTensor<uint32_t>&)maxVreg, from, maskFull); | 1477 | + Reg::CompareScalar(cmpMaskReg, (Reg::RegTensor<uint32_t>&)maxVreg, from, maskFull); |
| 1478 | } else if constexpr (sizeof(T2) == sizeof(half)) { | 1478 | } else if constexpr (sizeof(T2) == sizeof(half)) { |
| 1479 | - MicroAPI::LoadAlign<T2, MicroAPI::LoadDist::DIST_BRC_B16>(maxVreg, maxUb + i * stepSize); | 1479 | + Reg::LoadAlign<T2, Reg::LoadDist::DIST_BRC_B16>(maxVreg, maxUb + i * stepSize); |
| 1480 | // either full mask or zero mask | 1480 | // either full mask or zero mask |
| 1481 | - MicroAPI::CompareScalar(cmpMaskReg, (MicroAPI::RegTensor<uint16_t>&)maxVreg, (uint16_t)from, maskFull); | 1481 | + Reg::CompareScalar(cmpMaskReg, (Reg::RegTensor<uint16_t>&)maxVreg, (uint16_t)from, maskFull); |
| 1482 | } | 1482 | } |
| 1483 | if constexpr (sizeof(T1) != sizeof(T2)) { | 1483 | if constexpr (sizeof(T1) != sizeof(T2)) { |
| 1484 | - MicroAPI::MaskPack(cmpMaskReg0, cmpMaskReg); | 1484 | + Reg::MaskPack(cmpMaskReg0, cmpMaskReg); |
| 1485 | - MicroAPI::MaskPack<MicroAPI::HighLowPart::HIGHEST>(cmpMaskReg1, cmpMaskReg); | 1485 | + Reg::MaskPack<Reg::HighLowPart::HIGHEST>(cmpMaskReg1, cmpMaskReg); |
| 1486 | - MicroAPI::MaskOr(cmpMaskReg, cmpMaskReg0, cmpMaskReg1, maskFull); | 1486 | + Reg::MaskOr(cmpMaskReg, cmpMaskReg0, cmpMaskReg1, maskFull); |
| 1487 | } | 1487 | } |
| 1488 | - MicroAPI::MaskOr(dstMask, dstMask, cmpMaskReg, maskFull); | 1488 | + Reg::MaskOr(dstMask, dstMask, cmpMaskReg, maskFull); |
| 1489 | - MicroAPI::Duplicate(tmpVreg, to, maskFull); | 1489 | + Reg::Duplicate(tmpVreg, to, maskFull); |
| 1490 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { | 1490 | for (uint16_t j = 0; j < kRepeatTimes; ++j) { |
| 1491 | - MicroAPI::LoadAlign(srcVreg, resUb + i * stride + j * dataBlock); | 1491 | + Reg::LoadAlign(srcVreg, resUb + i * stride + j * dataBlock); |
| 1492 | - MicroAPI::Select(dstVreg, tmpVreg, srcVreg, cmpMaskReg); | 1492 | + Reg::Select(dstVreg, tmpVreg, srcVreg, cmpMaskReg); |
| 1493 | - MicroAPI::StoreAlign(resUb + i * stride + j * dataBlock, dstVreg, maskFull); | 1493 | + Reg::StoreAlign(resUb + i * stride + j * dataBlock, dstVreg, maskFull); |
| 1494 | } | 1494 | } |
| 1495 | } | 1495 | } |
| 1496 | - MicroAPI::StoreAlign((__ubuf__ uint8_t*)maskUb, dstMask); | 1496 | + Reg::StoreAlign((__ubuf__ uint8_t*)maskUb, dstMask); |
| 1497 | } | 1497 | } |
| 1498 | 1498 | ||
| 1499 | template <typename T1, typename T2, uint32_t stepSize, uint32_t stride> | 1499 | template <typename T1, typename T2, uint32_t stepSize, uint32_t stride> |
| @@ -1501,43 +1501,43 @@ __simd_vf__ inline void AdjustSoftMaxResNDImpl(__ubuf__ T1* resUb, __ubuf__ T2* | |||
| 1501 | __ubuf__ uint64_t* maskUb, const uint32_t from, const T1 to, const uint32_t srcK, const uint16_t srcM, | 1501 | __ubuf__ uint64_t* maskUb, const uint32_t from, const T1 to, const uint32_t srcK, const uint16_t srcM, |
| 1502 | const uint16_t repeatTimes) | 1502 | const uint16_t repeatTimes) |
| 1503 | { | 1503 | { |
| 1504 | - MicroAPI::RegTensor<T1> srcVreg; | 1504 | + Reg::RegTensor<T1> srcVreg; |
| 1505 | - MicroAPI::RegTensor<T1> tmpVreg; | 1505 | + Reg::RegTensor<T1> tmpVreg; |
| 1506 | - MicroAPI::RegTensor<T1> dstVreg; | 1506 | + Reg::RegTensor<T1> dstVreg; |
| 1507 | - MicroAPI::RegTensor<T2> maxVreg; | 1507 | + Reg::RegTensor<T2> maxVreg; |
| 1508 | - MicroAPI::MaskReg maskReg; | 1508 | + Reg::MaskReg maskReg; |
| 1509 | - MicroAPI::MaskReg cmpMaskReg; | 1509 | + Reg::MaskReg cmpMaskReg; |
| 1510 | - MicroAPI::MaskReg cmpMaskReg0; | 1510 | + Reg::MaskReg cmpMaskReg0; |
| 1511 | - MicroAPI::MaskReg cmpMaskReg1; | 1511 | + Reg::MaskReg cmpMaskReg1; |
| 1512 | - MicroAPI::MaskReg maskFull = MicroAPI::CreateMask<T1, MicroAPI::MaskPattern::ALL>(); | 1512 | + Reg::MaskReg maskFull = Reg::CreateMask<T1, Reg::MaskPattern::ALL>(); |
| 1513 | - MicroAPI::MaskReg dstMask = MicroAPI::CreateMask<T1, MicroAPI::MaskPattern::ALLF>(); | 1513 | + Reg::MaskReg dstMask = Reg::CreateMask<T1, Reg::MaskPattern::ALLF>(); |
| 1514 | 1514 | ||
| 1515 | for (uint16_t i = 0; i < srcM; i++) { | 1515 | for (uint16_t i = 0; i < srcM; i++) { |
| 1516 | if constexpr (sizeof(T2) == sizeof(float)) { | 1516 | if constexpr (sizeof(T2) == sizeof(float)) { |
| 1517 | - MicroAPI::LoadAlign<T2, MicroAPI::LoadDist::DIST_BRC_B32>(maxVreg, maxUb + i * stepSize); | 1517 | + Reg::LoadAlign<T2, Reg::LoadDist::DIST_BRC_B32>(maxVreg, maxUb + i * stepSize); |
| 1518 | // either full mask or zero mask | 1518 | // either full mask or zero mask |
| 1519 | - MicroAPI::CompareScalar(cmpMaskReg, (MicroAPI::RegTensor<uint32_t>&)maxVreg, from, maskFull); | 1519 | + Reg::CompareScalar(cmpMaskReg, (Reg::RegTensor<uint32_t>&)maxVreg, from, maskFull); |
| 1520 | } else if constexpr (sizeof(T2) == sizeof(half)) { | 1520 | } else if constexpr (sizeof(T2) == sizeof(half)) { |
| 1521 | - MicroAPI::LoadAlign<T2, MicroAPI::LoadDist::DIST_BRC_B16>(maxVreg, maxUb + i * stepSize); | 1521 | + Reg::LoadAlign<T2, Reg::LoadDist::DIST_BRC_B16>(maxVreg, maxUb + i * stepSize); |
| 1522 | // either full mask or zero mask | 1522 | // either full mask or zero mask |
| 1523 | - MicroAPI::CompareScalar(cmpMaskReg, (MicroAPI::RegTensor<uint16_t>&)maxVreg, (uint16_t)from, maskFull); | 1523 | + Reg::CompareScalar(cmpMaskReg, (Reg::RegTensor<uint16_t>&)maxVreg, (uint16_t)from, maskFull); |
| 1524 | } | 1524 | } |
| 1525 | if constexpr (sizeof(T1) == sizeof(half) && sizeof(T2) == sizeof(float)) { | 1525 | if constexpr (sizeof(T1) == sizeof(half) && sizeof(T2) == sizeof(float)) { |
| 1526 | - MicroAPI::MaskPack(cmpMaskReg0, cmpMaskReg); | 1526 | + Reg::MaskPack(cmpMaskReg0, cmpMaskReg); |
| 1527 | - MicroAPI::MaskPack<MicroAPI::HighLowPart::HIGHEST>(cmpMaskReg1, cmpMaskReg); | 1527 | + Reg::MaskPack<Reg::HighLowPart::HIGHEST>(cmpMaskReg1, cmpMaskReg); |
| 1528 | - MicroAPI::MaskOr(cmpMaskReg, cmpMaskReg0, cmpMaskReg1, maskFull); | 1528 | + Reg::MaskOr(cmpMaskReg, cmpMaskReg0, cmpMaskReg1, maskFull); |
| 1529 | } | 1529 | } |
| 1530 | - MicroAPI::MaskOr(dstMask, dstMask, cmpMaskReg, maskFull); | 1530 | + Reg::MaskOr(dstMask, dstMask, cmpMaskReg, maskFull); |
| 1531 | - MicroAPI::Duplicate(tmpVreg, to, maskFull); | 1531 | + Reg::Duplicate(tmpVreg, to, maskFull); |
| 1532 | uint32_t sreg = srcK; | 1532 | uint32_t sreg = srcK; |
| 1533 | for (uint16_t j = 0; j < repeatTimes; j++) { | 1533 | for (uint16_t j = 0; j < repeatTimes; j++) { |
| 1534 | - maskReg = MicroAPI::UpdateMask<T1>(sreg); | 1534 | + maskReg = Reg::UpdateMask<T1>(sreg); |
| 1535 | - MicroAPI::LoadAlign(srcVreg, resUb + i * srcK + j * stride); | 1535 | + Reg::LoadAlign(srcVreg, resUb + i * srcK + j * stride); |
| 1536 | - MicroAPI::Select(dstVreg, tmpVreg, srcVreg, cmpMaskReg); | 1536 | + Reg::Select(dstVreg, tmpVreg, srcVreg, cmpMaskReg); |
| 1537 | - MicroAPI::StoreAlign(resUb + i * srcK + j * stride, dstVreg, maskReg); | 1537 | + Reg::StoreAlign(resUb + i * srcK + j * stride, dstVreg, maskReg); |
| 1538 | } | 1538 | } |
| 1539 | } | 1539 | } |
| 1540 | - MicroAPI::StoreAlign((__ubuf__ uint8_t*)maskUb, dstMask); | 1540 | + Reg::StoreAlign((__ubuf__ uint8_t*)maskUb, dstMask); |
| 1541 | } | 1541 | } |
| 1542 | 1542 | ||
| 1543 | template <typename T1, typename T2, bool isDataFormatNZ = false, uint8_t stepSizeMode = 0> | 1543 | template <typename T1, typename T2, bool isDataFormatNZ = false, uint8_t stepSizeMode = 0> |