@@ -0,0 +1,139 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_aligned128_no_update_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_ALIGNED128_NO_UPDATE_SFA_H
|
||||
#define VF_BASIC_BLOCK_ALIGNED128_NO_UPDATE_SFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
// no update, originN == 128
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__simd_vf__ void ProcessVec1NoUpdateImpl128VF(
|
||||
__ubuf__ T2 * expUb, __ubuf__ T * expSumUb, __ubuf__ T * maxUb, __ubuf__ T * maxUbStart,
|
||||
__ubuf__ T * srcUb, const uint32_t blockStride, const uint32_t repeatStride,
|
||||
const uint16_t m, const T scale, const T minValue)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_even;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_odd;
|
||||
|
||||
// bfloat16_t
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_even_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_odd_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
|
||||
// half
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_even_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_odd_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
|
||||
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_max;
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
|
||||
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<T, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x_unroll, srcUb + floatRepSize + i * s2BaseSize);
|
||||
|
||||
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_all); // Muls(scale)
|
||||
AscendC::MicroAPI::Muls(vreg_input_x_unroll, vreg_input_x_unroll, scale, preg_all);
|
||||
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + floatRepSize + i * s2BaseSize, vreg_input_x_unroll, preg_all);
|
||||
AscendC::MicroAPI::Max(vreg_max_tmp, vreg_input_x, vreg_input_x_unroll, preg_all);
|
||||
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_input_max, vreg_max_tmp, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)maxUb), vreg_input_max, ureg_max, 1);
|
||||
}
|
||||
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)maxUb), ureg_max, 0);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
// maxUb is [S1, 1], BRC_B32 is reading one fp32 element and broadcast it to all 64 vreg element
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(
|
||||
vreg_max_brc, maxUbStart + i);
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_DINTLV_B32>(
|
||||
vreg_input_x, vreg_input_x_unroll, srcUb + i * s2BaseSize);
|
||||
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_even, vreg_input_x, vreg_max_brc, preg_all);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_odd, vreg_input_x_unroll, vreg_max_brc, preg_all);
|
||||
|
||||
// x_sum = sum(x_exp, axis=-1, keepdims=True)
|
||||
AscendC::MicroAPI::Add(vreg_exp_sum, vreg_exp_even, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_exp_sum, vreg_exp_sum, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)expSumUb), vreg_exp_sum, ureg_exp_sum, 1);
|
||||
|
||||
if constexpr (IsSameType<T2, bfloat16_t>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_bf16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_bf16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_bf16, (RegTensor<uint16_t>&)vreg_exp_even_bf16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_bf16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_bf16, blockStride, repeatStride, preg_all_b16);
|
||||
} else if constexpr (IsSameType<T2, half>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_fp16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_fp16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_fp16, (RegTensor<uint16_t>&)vreg_exp_even_fp16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_fp16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_fp16, blockStride, repeatStride, preg_all_b16);
|
||||
}
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)expSumUb), ureg_exp_sum, 0);
|
||||
}
|
||||
|
||||
// no update, originN == 128
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__aicore__ inline void ProcessVec1NoUpdateImpl128(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
|
||||
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
|
||||
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
|
||||
const uint32_t repeatStride = 1;
|
||||
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUbStart = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
|
||||
ProcessVec1NoUpdateImpl128VF<T, T2, s1BaseSize, s2BaseSize>(
|
||||
expUb, expSumUb, maxUb, maxUbStart, srcUb, blockStride, repeatStride, m, scale, minValue);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_ALIGNED128_NO_UPDATE_SFA_H
|
||||
@@ -0,0 +1,143 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_aligned128_update_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_ALIGNED128_UPDATE_SFA_H
|
||||
#define VF_BASIC_BLOCK_ALIGNED128_UPDATE_SFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
// update, originN == 128
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 128, uint32_t s2BaseSize = 128>
|
||||
__simd_vf__ void ProcessVec1UpdateImpl128VF(
|
||||
__ubuf__ T2 * expUb, __ubuf__ T * srcUb, __ubuf__ T * inMaxUb,
|
||||
__ubuf__ T * tmpExpSumUb, __ubuf__ T * tmpMaxUb, __ubuf__ T * tmpMaxUb2, const uint32_t blockStride,
|
||||
const uint32_t repeatStride, const uint16_t m, const T scale, const T minValue)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_in_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_new;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_cur_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_in_exp_sum;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_even;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_odd;
|
||||
|
||||
// bfloat16_t
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_even_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_odd_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
|
||||
// half
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_even_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_odd_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
|
||||
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_max;
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
// x_max = max(src, axis=-1, keepdims=True); x_max = Max(x_max, inMax)
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x_unroll, srcUb + floatRepSize + i * s2BaseSize);
|
||||
|
||||
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_all); // Muls(scale)
|
||||
AscendC::MicroAPI::Muls(vreg_input_x_unroll, vreg_input_x_unroll, scale, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + floatRepSize + i * s2BaseSize, vreg_input_x_unroll, preg_all);
|
||||
AscendC::MicroAPI::Max(vreg_max_tmp, vreg_input_x, vreg_input_x_unroll, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_max_tmp, vreg_max_tmp, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpMaxUb), vreg_max_tmp, ureg_max, 1);
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpMaxUb), ureg_max, 0);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_in_max, inMaxUb);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
AscendC::MicroAPI::LoadAlign(vreg_cur_max, tmpMaxUb2); // 获取新的max[s1, 1]
|
||||
AscendC::MicroAPI::Max(vreg_max_new, vreg_cur_max, vreg_in_max, preg_all); // 计算新、旧max的最大值
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)tmpMaxUb2, vreg_max_new, preg_all);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_max_brc, tmpMaxUb2 + i);
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_DINTLV_B32>(
|
||||
vreg_input_x, vreg_input_x_unroll, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_even, vreg_input_x, vreg_max_brc, preg_all);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_odd, vreg_input_x_unroll, vreg_max_brc, preg_all);
|
||||
|
||||
// x_sum = sum(x_exp, axis=-1, keepdims=True)
|
||||
AscendC::MicroAPI::Add(vreg_exp_sum, vreg_exp_even, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_exp_sum, vreg_exp_sum, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpExpSumUb), vreg_exp_sum, ureg_exp_sum, 1);
|
||||
|
||||
if constexpr (IsSameType<T2, bfloat16_t>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_bf16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_bf16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_bf16, (RegTensor<uint16_t>&)vreg_exp_even_bf16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_bf16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_bf16, blockStride, repeatStride, preg_all_b16);
|
||||
} else if constexpr (IsSameType<T2, half>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_fp16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_fp16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_fp16, (RegTensor<uint16_t>&)vreg_exp_even_fp16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_fp16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_fp16, blockStride, repeatStride, preg_all_b16);
|
||||
}
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpExpSumUb), ureg_exp_sum, 0);
|
||||
}
|
||||
|
||||
// update, originN == 128
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__aicore__ inline void ProcessVec1UpdateImpl128(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
|
||||
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
|
||||
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
|
||||
const uint32_t repeatStride = 1;
|
||||
|
||||
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
__ubuf__ T * inMaxUb = (__ubuf__ T*)inMaxTensor.GetPhyAddr();
|
||||
__ubuf__ T * tmpExpSumUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr();
|
||||
__ubuf__ T * tmpMaxUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
__ubuf__ T * tmpMaxUb2 = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
|
||||
ProcessVec1UpdateImpl128VF <T, T2, s1BaseSize, s2BaseSize>(
|
||||
expUb, srcUb, inMaxUb, tmpExpSumUb, tmpMaxUb, tmpMaxUb2, blockStride, repeatStride, m, scale, minValue);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_ALIGNED128_UPDATE_SFA_H
|
||||
@@ -0,0 +1,149 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_unaligned128_no_update_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_UNALIGNED128_NO_UPDATE_SFA_H
|
||||
#define VF_BASIC_BLOCK_UNALIGNED128_NO_UPDATE_SFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__simd_vf__ void ProcessVec1NoUpdateGeneralImpl128VF(
|
||||
__ubuf__ T2 * expUb, __ubuf__ T * expSumUb, __ubuf__ T * maxUb, __ubuf__ T * maxUbStart,
|
||||
__ubuf__ T * srcUb, const uint32_t blockStride, const uint32_t repeatStride,
|
||||
const uint16_t m, const T scale, const T minValue, uint32_t pltOriTailN, uint32_t pltTailN)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_min;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll_new;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_even;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_odd;
|
||||
|
||||
// bfloat16_t
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_even_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_odd_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
|
||||
// half
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_even_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_odd_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
|
||||
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_max;
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
|
||||
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b8 = AscendC::MicroAPI::CreateMask<T2, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_tail_n = AscendC::MicroAPI::UpdateMask<float>(pltTailN);
|
||||
AscendC::MicroAPI::MaskReg preg_ori_tail_n = AscendC::MicroAPI::UpdateMask<float>(pltOriTailN);
|
||||
AscendC::MicroAPI::MaskReg preg_reduce_n =
|
||||
AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::VL8>();
|
||||
|
||||
AscendC::MicroAPI::Duplicate(vreg_min, minValue);
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x_unroll, srcUb + floatRepSize + i * s2BaseSize);
|
||||
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_all); // Muls(scale)
|
||||
AscendC::MicroAPI::Muls(vreg_input_x_unroll, vreg_input_x_unroll, scale, preg_ori_tail_n);
|
||||
AscendC::MicroAPI::Select(vreg_input_x_unroll_new, vreg_input_x_unroll, vreg_min, preg_ori_tail_n);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + floatRepSize + i * s2BaseSize, vreg_input_x_unroll_new, preg_tail_n);
|
||||
|
||||
AscendC::MicroAPI::Max(vreg_max_tmp, vreg_input_x, vreg_input_x_unroll_new, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_input_max, vreg_max_tmp, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)maxUb), vreg_input_max, ureg_max, 1);
|
||||
}
|
||||
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)maxUb), ureg_max, 0);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_max_brc, maxUbStart + i);
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_DINTLV_B32>(
|
||||
vreg_input_x, vreg_input_x_unroll, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_even, vreg_input_x, vreg_max_brc, preg_all);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_odd, vreg_input_x_unroll, vreg_max_brc, preg_all);
|
||||
|
||||
// x_sum = sum(x_exp, axis=-1, keepdims=True)
|
||||
AscendC::MicroAPI::Add(vreg_exp_sum, vreg_exp_even, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_exp_sum, vreg_exp_sum, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)expSumUb), vreg_exp_sum, ureg_exp_sum, 1);
|
||||
|
||||
if constexpr (IsSameType<T2, bfloat16_t>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_bf16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_bf16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_bf16, (RegTensor<uint16_t>&)vreg_exp_even_bf16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_bf16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_bf16, blockStride, repeatStride, preg_all_b16);
|
||||
} else if constexpr (IsSameType<T2, half>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_fp16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_fp16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_fp16, (RegTensor<uint16_t>&)vreg_exp_even_fp16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_fp16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_fp16, blockStride, repeatStride, preg_all_b16);
|
||||
}
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)expSumUb), ureg_exp_sum, 0);
|
||||
}
|
||||
|
||||
// no update, 64 < originN <= 128
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__aicore__ inline void ProcessVec1NoUpdateGeneralImpl128(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
|
||||
// stride, high 16bits: blockStride (65*16*2/32),单位block, low 16bits: repeatStride (1)
|
||||
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
|
||||
const uint32_t repeatStride = 1;
|
||||
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUbStart = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
|
||||
const uint32_t oriTailN = originN - floatRepSize;
|
||||
const uint32_t tailN = s2BaseSize - floatRepSize;
|
||||
uint32_t pltOriTailN = oriTailN;
|
||||
uint32_t pltTailN = tailN;
|
||||
|
||||
ProcessVec1NoUpdateGeneralImpl128VF<T, T2, s1BaseSize, s2BaseSize>(
|
||||
expUb, expSumUb, maxUb, maxUbStart, srcUb, blockStride, repeatStride, m, scale, minValue,
|
||||
pltOriTailN, pltTailN);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_UNALIGNED128_NO_UPDATE_SFA_H
|
||||
@@ -0,0 +1,159 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_unaligned128_update_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_UNALIGNED128_UPDATE_SFA_H
|
||||
#define VF_BASIC_BLOCK_UNALIGNED128_UPDATE_SFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__simd_vf__ void ProcessVec1UpdateGeneralImpl128VF(
|
||||
__ubuf__ T2 * expUb, __ubuf__ T * srcUb, __ubuf__ T * inMaxUb,
|
||||
__ubuf__ T * tmpExpSumUb, __ubuf__ T * tmpMaxUb, __ubuf__ T * tmpMaxUb2, const uint32_t blockStride,
|
||||
const uint32_t repeatStride, const uint16_t m, const T scale, const T minValue, uint32_t pltOriTailN,
|
||||
uint32_t pltTailN, uint32_t pltN)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_min;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll_new;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_cur_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_new;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_in_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_even;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_odd;
|
||||
|
||||
// bfloat16_t
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_even_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_odd_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
|
||||
// half
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_even_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_odd_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
|
||||
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_max;
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
|
||||
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b16 = AscendC::MicroAPI::CreateMask<uint16_t,
|
||||
AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_n_b16 = AscendC::MicroAPI::UpdateMask<uint16_t>(pltN);
|
||||
AscendC::MicroAPI::MaskReg preg_tail_n = AscendC::MicroAPI::UpdateMask<T>(pltTailN);
|
||||
AscendC::MicroAPI::MaskReg preg_ori_tail_n = AscendC::MicroAPI::UpdateMask<T>(pltOriTailN);
|
||||
|
||||
AscendC::MicroAPI::Duplicate(vreg_min, minValue);
|
||||
// x_max = max(src, axis=-1, keepdims=True); x_max = Max(x_max, inMax)
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x_unroll, srcUb + floatRepSize + i * s2BaseSize);
|
||||
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_all); // Muls(scale)
|
||||
AscendC::MicroAPI::Muls(vreg_input_x_unroll, vreg_input_x_unroll, scale, preg_ori_tail_n);
|
||||
AscendC::MicroAPI::Select(vreg_input_x_unroll_new, vreg_input_x_unroll, vreg_min, preg_ori_tail_n);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + floatRepSize + i * s2BaseSize, vreg_input_x_unroll_new, preg_tail_n);
|
||||
AscendC::MicroAPI::Max(vreg_max_tmp, vreg_input_x, vreg_input_x_unroll_new, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_cur_max, vreg_max_tmp, preg_all);
|
||||
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpMaxUb), vreg_cur_max, ureg_max, 1);
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpMaxUb), ureg_max, 0);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_in_max, inMaxUb);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
AscendC::MicroAPI::LoadAlign(vreg_cur_max, tmpMaxUb2); // 获取新的max[s1, 1]
|
||||
AscendC::MicroAPI::Max(vreg_max_new, vreg_cur_max, vreg_in_max, preg_all); // 计算新、旧max的最大值
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)tmpMaxUb2, vreg_max_new, preg_all);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(
|
||||
vreg_max_brc, tmpMaxUb2 + i);
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_DINTLV_B32>(
|
||||
vreg_input_x, vreg_input_x_unroll, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_even, vreg_input_x, vreg_max_brc, preg_all);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_odd, vreg_input_x_unroll, vreg_max_brc, preg_all);
|
||||
|
||||
// x_sum = sum(x_exp, axis=-1, keepdims=True)
|
||||
AscendC::MicroAPI::Add(vreg_exp_sum, vreg_exp_even, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_exp_sum, vreg_exp_sum, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpExpSumUb), vreg_exp_sum, ureg_exp_sum, 1);
|
||||
|
||||
if constexpr (IsSameType<T2, bfloat16_t>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_bf16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_bf16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_bf16, (RegTensor<uint16_t>&)vreg_exp_even_bf16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_bf16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_bf16, blockStride, repeatStride, preg_n_b16);
|
||||
} else if constexpr (IsSameType<T2, half>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_fp16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_fp16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_fp16, (RegTensor<uint16_t>&)vreg_exp_even_fp16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_fp16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_fp16, blockStride, repeatStride, preg_n_b16);
|
||||
}
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpExpSumUb), ureg_exp_sum, 0);
|
||||
}
|
||||
|
||||
|
||||
// update, 64 < originN <= 128
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__aicore__ inline void ProcessVec1UpdateGeneralImpl128(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
|
||||
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
|
||||
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
|
||||
const uint32_t repeatStride = 1;
|
||||
const uint32_t oriTailN = originN - floatRepSize;
|
||||
const uint32_t tailN = s2BaseSize - floatRepSize;
|
||||
uint32_t pltOriTailN = oriTailN;
|
||||
uint32_t pltTailN = tailN;
|
||||
uint32_t pltN = s2BaseSize;
|
||||
|
||||
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
__ubuf__ T * inMaxUb = (__ubuf__ T*)inMaxTensor.GetPhyAddr();
|
||||
__ubuf__ T * tmpExpSumUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr();
|
||||
__ubuf__ T * tmpMaxUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
__ubuf__ T * tmpMaxUb2 = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
|
||||
ProcessVec1UpdateGeneralImpl128VF<T, T2, s1BaseSize, s2BaseSize>(
|
||||
expUb, srcUb, inMaxUb, tmpExpSumUb, tmpMaxUb, tmpMaxUb2, blockStride, repeatStride,
|
||||
m, scale, minValue, pltOriTailN, pltTailN, pltN);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_UNALIGNED128_UPDATE_SFA_H
|
||||
@@ -0,0 +1,129 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_unaligned64_no_update_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_UNALIGNED64_NO_UPDATE_SFA_H
|
||||
#define VF_BASIC_BLOCK_UNALIGNED64_NO_UPDATE_SFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__simd_vf__ void ProcessVec1NoUpdateImpl64VF(
|
||||
__ubuf__ T2 * expUb, __ubuf__ T * expSumUb, __ubuf__ T * maxUb, __ubuf__ T * maxUbStart,
|
||||
__ubuf__ T * srcUb, const uint32_t blockStride, const uint32_t repeatStride,
|
||||
const uint16_t m, const T scale, const T minValue, uint32_t pltOriginalN, uint32_t pltSrcN)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_min;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
|
||||
// bfloat16_t
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_dst_even_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_dst_odd_bf16;
|
||||
// half
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_dst_even_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_dst_odd_fp16;
|
||||
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_max;
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
|
||||
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_src_n = AscendC::MicroAPI::UpdateMask<float>(pltSrcN);
|
||||
AscendC::MicroAPI::MaskReg preg_src_n_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::H>();
|
||||
AscendC::MicroAPI::MaskReg preg_ori_src_n = AscendC::MicroAPI::UpdateMask<T>(pltOriginalN);
|
||||
|
||||
// x_max = max(src, axis=-1, keepdims=True)
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_ori_src_n); // Muls(scale)
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_src_n);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_input_max, vreg_input_x, preg_ori_src_n);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)maxUb), vreg_input_max, ureg_max, 1);
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)maxUb), ureg_max, 0);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(
|
||||
vreg_max_brc, maxUbStart + i);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp, vreg_input_x, vreg_max_brc, preg_ori_src_n);
|
||||
|
||||
// x_sum = sum(x_exp, axis=-1, keepdims=True)
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_exp_sum, vreg_exp, preg_ori_src_n);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)expSumUb), vreg_exp_sum, ureg_exp_sum, 1);
|
||||
|
||||
if constexpr (IsSameType<T2, bfloat16_t>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_bf16, vreg_exp, preg_all_b16);
|
||||
AscendC::MicroAPI::DeInterleave(vreg_dst_even_bf16, vreg_dst_odd_bf16,
|
||||
vreg_exp_bf16, vreg_exp_bf16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_dst_even_bf16, blockStride, repeatStride, preg_src_n_b16);
|
||||
} else if constexpr (IsSameType<T2, half>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_fp16, vreg_exp, preg_all_b16);
|
||||
AscendC::MicroAPI::DeInterleave(vreg_dst_even_fp16, vreg_dst_odd_fp16,
|
||||
vreg_exp_fp16, vreg_exp_fp16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_dst_even_fp16, blockStride, repeatStride, preg_src_n_b16);
|
||||
}
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)expSumUb), ureg_exp_sum, 0);
|
||||
}
|
||||
|
||||
// no update, originN <= 64
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__aicore__ inline void ProcessVec1NoUpdateImpl64(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUbStart = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
|
||||
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
|
||||
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
|
||||
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
|
||||
const uint32_t repeatStride = 1;
|
||||
uint32_t pltOriginalN = originN;
|
||||
uint32_t pltSrcN = s2BaseSize;
|
||||
|
||||
ProcessVec1NoUpdateImpl64VF<T, T2, s1BaseSize, s2BaseSize>(
|
||||
expUb, expSumUb, maxUb, maxUbStart, srcUb, blockStride, repeatStride, m, scale, minValue,
|
||||
pltOriginalN, pltSrcN);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_UNALIGNED64_NO_UPDATE_SFA_H
|
||||
@@ -0,0 +1,141 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_aligned64_update_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_ALIGNED64_UPDATE_SFA_H
|
||||
#define VF_BASIC_BLOCK_ALIGNED64_UPDATE_SFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
// update, originN <= 64
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 128, uint32_t s2BaseSize = 128>
|
||||
__simd_vf__ void ProcessVec1UpdateImpl64VF(
|
||||
__ubuf__ T2 * expUb, __ubuf__ T * srcUb, __ubuf__ T * inMaxUb,
|
||||
__ubuf__ T * tmpExpSumUb, __ubuf__ T * tmpMaxUb, __ubuf__ T * tmpMaxUb2, const uint32_t blockStride,
|
||||
const uint32_t repeatStride, const uint16_t m, const T scale, const T minValue, uint32_t pltOriginalN,
|
||||
uint32_t pltSrcN)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_in_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_new;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_cur_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
|
||||
// bfloat16_t
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_dst_even_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_dst_odd_bf16;
|
||||
// half
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_dst_even_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_dst_odd_fp16;
|
||||
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_max;
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
|
||||
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_ori_src_n = AscendC::MicroAPI::UpdateMask<T>(pltOriginalN);
|
||||
AscendC::MicroAPI::MaskReg preg_src_n = AscendC::MicroAPI::UpdateMask<T>(pltSrcN);
|
||||
AscendC::MicroAPI::MaskReg preg_src_n_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::H>();
|
||||
|
||||
// x_max = max(src, axis=-1, keepdims=True)
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_ori_src_n);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_src_n);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_cur_max, vreg_input_x, preg_ori_src_n);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpMaxUb), vreg_cur_max, ureg_max, 1);
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpMaxUb), ureg_max, 0);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_in_max, inMaxUb);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
AscendC::MicroAPI::LoadAlign(vreg_cur_max, tmpMaxUb2);
|
||||
AscendC::MicroAPI::Max(vreg_max_new, vreg_cur_max, vreg_in_max, preg_all); // 计算新、旧的最大值
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)tmpMaxUb2, vreg_max_new, preg_all);
|
||||
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(
|
||||
vreg_max_brc, tmpMaxUb2 + i);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp, vreg_input_x, vreg_max_brc, preg_ori_src_n);
|
||||
|
||||
// x_sum = sum(x_exp, axis=-1, keepdims=True)
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_exp_sum, vreg_exp, preg_ori_src_n);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpExpSumUb), vreg_exp_sum, ureg_exp_sum, 1);
|
||||
|
||||
if constexpr (IsSameType<T2, bfloat16_t>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_bf16, vreg_exp, preg_all_b16);
|
||||
AscendC::MicroAPI::DeInterleave(vreg_dst_even_bf16, vreg_dst_odd_bf16,
|
||||
vreg_exp_bf16, vreg_exp_bf16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_dst_even_bf16, blockStride, repeatStride, preg_src_n_b16);
|
||||
} else if constexpr (IsSameType<T2, half>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_fp16, vreg_exp, preg_all_b16);
|
||||
AscendC::MicroAPI::DeInterleave(vreg_dst_even_fp16, vreg_dst_odd_fp16,
|
||||
vreg_exp_fp16, vreg_exp_fp16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_dst_even_fp16, blockStride, repeatStride, preg_src_n_b16);
|
||||
}
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpExpSumUb), ureg_exp_sum, 0);
|
||||
}
|
||||
|
||||
|
||||
// update, originN <= 64
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__aicore__ inline void ProcessVec1UpdateImpl64(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
|
||||
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
|
||||
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
|
||||
const uint32_t repeatStride = 1;
|
||||
uint32_t pltOriginalN = originN;
|
||||
uint32_t pltSrcN = s2BaseSize;
|
||||
|
||||
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
__ubuf__ T * inMaxUb = (__ubuf__ T*)inMaxTensor.GetPhyAddr();
|
||||
__ubuf__ T * tmpExpSumUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr();
|
||||
__ubuf__ T * tmpMaxUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
__ubuf__ T * tmpMaxUb2 = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
|
||||
ProcessVec1UpdateImpl64VF <T, T2, s1BaseSize, s2BaseSize>(
|
||||
expUb, srcUb, inMaxUb, tmpExpSumUb, tmpMaxUb, tmpMaxUb2, blockStride, repeatStride, m, scale, minValue,
|
||||
pltOriginalN, pltSrcN);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_ALIGNED64_UPDATE_SFA_H
|
||||
112
csrc/attention/common/op_kernel/arch35/vf/vf_basic_block_utils.h
Normal file
112
csrc/attention/common/op_kernel/arch35/vf/vf_basic_block_utils.h
Normal file
@@ -0,0 +1,112 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_utils.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_UTILS_H
|
||||
#define VF_BASIC_BLOCK_UTILS_H
|
||||
|
||||
#if ASC_DEVKIT_MAJOR >= 9
|
||||
#include "kernel_basic_intf.h"
|
||||
#else
|
||||
#include "kernel_operator.h"
|
||||
#endif
|
||||
|
||||
namespace FaVectorApi {
|
||||
constexpr uint32_t floatRepSize = 64;
|
||||
constexpr uint32_t halfRepSize = 128;
|
||||
constexpr uint32_t blockBytesU8 = 32;
|
||||
constexpr float fp8e4m3MaxValue = 448.0f;
|
||||
constexpr float int8MaxValue = 127.0f;
|
||||
constexpr float hifp8MaxValue = 32768.0f;
|
||||
constexpr float floatEps = 2.220446049250313e-16;
|
||||
/* **************************************************************************************************
|
||||
* Muls + Select(optional) + SoftmaxFlashV2 + Cast(fp32->fp16/bf16) + ND2NZ
|
||||
* ************************************************************************************************* */
|
||||
using namespace MicroAPI;
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitZero = {
|
||||
AscendC::MicroAPI::RegLayout::ZERO,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_ROUND,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitOne = {
|
||||
AscendC::MicroAPI::RegLayout::ONE,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_ROUND,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitTwo = {
|
||||
AscendC::MicroAPI::RegLayout::TWO,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_ROUND,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitThree = {
|
||||
AscendC::MicroAPI::RegLayout::THREE,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_ROUND,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitRintZero = {
|
||||
AscendC::MicroAPI::RegLayout::ZERO,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_RINT,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitRintOne = {
|
||||
AscendC::MicroAPI::RegLayout::ONE,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_RINT,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitRintTwo = {
|
||||
AscendC::MicroAPI::RegLayout::TWO,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_RINT,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitRintThree = {
|
||||
AscendC::MicroAPI::RegLayout::THREE,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_RINT,
|
||||
};
|
||||
|
||||
#define USE_MLA_FULLQUANT_V1_P(vreg_exp, vreg_rowmax_p, MaskReg) \
|
||||
do { \
|
||||
Muls(vreg_exp, vreg_exp, fp8e4m3MaxValue, MaskReg); \
|
||||
Div(vreg_exp, vreg_exp, vreg_rowmax_p, MaskReg); \
|
||||
} while (0)
|
||||
|
||||
#define USE_MLA_FULLQUANT_V1_P_INT8(vreg_exp, vreg_rowmax_p, MaskReg) \
|
||||
do { \
|
||||
Muls(vreg_exp, vreg_exp, int8MaxValue, MaskReg); \
|
||||
Div(vreg_exp, vreg_exp, vreg_rowmax_p, MaskReg); \
|
||||
} while (0)
|
||||
|
||||
#define USE_MLA_FULLQUANT_V1_P_HIFP8(vreg_exp, vreg_rowmax_p, MaskReg) \
|
||||
do { \
|
||||
Muls(vreg_exp, vreg_exp, hifp8MaxValue, MaskReg); \
|
||||
Div(vreg_exp, vreg_exp, vreg_rowmax_p, MaskReg); \
|
||||
} while (0)
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_UTILS_H
|
||||
727
csrc/attention/common/op_kernel/arch35/vf/vf_flashupdate_new.h
Normal file
727
csrc/attention/common/op_kernel/arch35/vf/vf_flashupdate_new.h
Normal file
@@ -0,0 +1,727 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_flashupdate_new.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef MY_FLASH_UPDATE_NEW_INTERFACE_H
|
||||
#define MY_FLASH_UPDATE_NEW_INTERFACE_H
|
||||
|
||||
#include "kernel_tensor.h"
|
||||
|
||||
namespace FaVectorApi {
|
||||
// bf16->fp32
|
||||
static constexpr MicroAPI::CastTrait castTraitFp16_32_update = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN,
|
||||
MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
constexpr uint16_t REDUCE_SIZE = 1;
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize, bool isUpdatePre, bool isMlaFullQuant>
|
||||
__simd_vf__ inline void FlashUpdateBasicVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * preUb,
|
||||
__ubuf__ float * expMaxUb, __ubuf__ float * rowMaxUb, const uint16_t m, const uint16_t d,
|
||||
const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
constexpr uint16_t dLoops = srcD / floatRepSize;
|
||||
RegTensor<float> vreg_exp_max;
|
||||
RegTensor<float> vreg_row_max;
|
||||
RegTensor<float> vreg_input_pre;
|
||||
RegTensor<float> vreg_input_cur;
|
||||
RegTensor<float> vreg_mul;
|
||||
RegTensor<float> vreg_add;
|
||||
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
|
||||
// dstTensor = preTensor * expMaxTensor + curTensor
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize); // [m,8]
|
||||
if constexpr (isMlaFullQuant) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_row_max, rowMaxUb + i * reduceSize);
|
||||
}
|
||||
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
LoadAlign(vreg_input_pre, preUb + i * d + j * floatRepSize);
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
|
||||
if constexpr (isMlaFullQuant) {
|
||||
Mul(vreg_input_cur, vreg_input_cur, vreg_row_max, preg_all);
|
||||
}
|
||||
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_all);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
if constexpr (isUpdatePre) {
|
||||
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
|
||||
}
|
||||
}
|
||||
Add(vreg_add, vreg_mul, vreg_input_cur, preg_all);
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_add, preg_all);
|
||||
}
|
||||
}
|
||||
}
|
||||
/* **************************************************************************************************
|
||||
* FlashUpdate, fp32
|
||||
* ************************************************************************************************* */
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize, bool isUpdatePre, bool isMlaFullQuant>
|
||||
__aicore__ inline void FlashUpdateBasic(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
|
||||
const LocalTensor<T>& preTensor, const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& rowMaxTensor,
|
||||
const uint16_t m, const uint16_t d, const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
|
||||
__ubuf__ float * preUb = (__ubuf__ T*)preTensor.GetPhyAddr();
|
||||
__ubuf__ float * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
|
||||
__ubuf__ float * rowMaxUb = (__ubuf__ T*)rowMaxTensor.GetPhyAddr();
|
||||
|
||||
FlashUpdateBasicVF<T, INPUT_T, OUTPUT_T, srcD, reduceSize, isUpdatePre, isMlaFullQuant>(
|
||||
dstUb, curUb, preUb, expMaxUb, rowMaxUb, m, d, deScaleV, deScaleVPre);
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t reduceSize, bool isUpdatePre>
|
||||
__simd_vf__ inline void FlashUpdateGeneralVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * preUb,
|
||||
__ubuf__ float * expMaxUb, const uint16_t m, const uint16_t d,
|
||||
const float deScaleV, const float deScaleVPre, const uint32_t pltTailD, const uint16_t hasTail)
|
||||
{
|
||||
RegTensor<float> vreg_exp_max;
|
||||
RegTensor<float> vreg_input_pre;
|
||||
RegTensor<float> vreg_input_cur;
|
||||
RegTensor<float> vreg_mul;
|
||||
RegTensor<float> vreg_add;
|
||||
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
uint32_t tmpTailD = pltTailD;
|
||||
MaskReg preg_tail_d = UpdateMask<float>(tmpTailD);
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
const uint16_t dLoops = d / floatRepSize;
|
||||
|
||||
// dstTensor = preTensor * expMaxTensor + curTensor
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize); // [m,8]
|
||||
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
LoadAlign(vreg_input_pre, preUb + i * d + j * floatRepSize);
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
|
||||
|
||||
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_all);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
if constexpr (isUpdatePre) {
|
||||
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
|
||||
}
|
||||
}
|
||||
Add(vreg_add, vreg_mul, vreg_input_cur, preg_all);
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_add, preg_all);
|
||||
}
|
||||
for (uint16_t t = 0; t < hasTail; ++t) {
|
||||
LoadAlign(vreg_input_pre, preUb + i * d + dLoops * floatRepSize);
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + dLoops * floatRepSize);
|
||||
|
||||
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_tail_d);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
if constexpr (isUpdatePre) {
|
||||
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
|
||||
}
|
||||
}
|
||||
Add(vreg_add, vreg_mul, vreg_input_cur, preg_tail_d);
|
||||
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + dLoops * floatRepSize, vreg_add, preg_tail_d);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t reduceSize, bool isUpdatePre>
|
||||
__aicore__ inline void FlashUpdateGeneral(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
|
||||
const LocalTensor<T>& preTensor, const LocalTensor<T>& expMaxTensor, const uint16_t m, const uint16_t d,
|
||||
const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
|
||||
__ubuf__ float * preUb = (__ubuf__ T*)preTensor.GetPhyAddr();
|
||||
__ubuf__ float * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
|
||||
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
const uint16_t tailD = d % floatRepSize;
|
||||
uint32_t pltTailD = static_cast<uint32_t>(tailD);
|
||||
|
||||
uint16_t hasTail = 0;
|
||||
if (tailD > 0) {
|
||||
hasTail = 1;
|
||||
}
|
||||
|
||||
FlashUpdateGeneralVF<T, INPUT_T, OUTPUT_T, reduceSize, isUpdatePre>(
|
||||
dstUb, curUb, preUb, expMaxUb, m, d, deScaleV, deScaleVPre, pltTailD, hasTail);
|
||||
}
|
||||
|
||||
/*
|
||||
* @ingroup FlashUpdate
|
||||
* @brief compute, dstTensor = preTensor * expMaxTensor + curTensor
|
||||
* @param [out] dstTensor, output LocalTensor
|
||||
* @param [in] curTensor, input LocalTensor
|
||||
* @param [in] preTensor, input LocalTensor
|
||||
* @param [in] expMaxTensor, input LocalTensor
|
||||
* @param [in] m, input rows
|
||||
* @param [in] d, input columns, should be 32 bytes aligned
|
||||
*/
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, bool isUpdatePre, bool isMlaFullQuant>
|
||||
__aicore__ inline void FlashUpdateNew(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
|
||||
const LocalTensor<T>& preTensor, const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& rowMaxTensor, const uint16_t m, const uint16_t d,
|
||||
const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
static_assert(IsSameType<T, float>::value, "VF FlashUpdate, T must be float");
|
||||
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
if constexpr(srcD % floatRepSize == 0) {
|
||||
FlashUpdateBasic<T, INPUT_T, OUTPUT_T, srcD, REDUCE_SIZE, isUpdatePre, isMlaFullQuant>(dstTensor, curTensor, preTensor, expMaxTensor, rowMaxTensor,
|
||||
m, d, deScaleV, deScaleVPre);
|
||||
} else {
|
||||
|
||||
FlashUpdateGeneral<T, INPUT_T, OUTPUT_T, REDUCE_SIZE, isUpdatePre>(dstTensor, curTensor, preTensor, expMaxTensor, m, d,
|
||||
deScaleV, deScaleVPre);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize, bool isUpdatePre, bool isMlaFullQuant>
|
||||
__simd_vf__ inline void FlashUpdateLastBasicVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * preUb,
|
||||
__ubuf__ float * expMaxUb, __ubuf__ float * expSumUb, __ubuf__ float * rowMaxUb, const uint16_t m, const uint16_t d,
|
||||
const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
RegTensor<float> vreg_exp_max;
|
||||
RegTensor<float> vreg_row_max;
|
||||
RegTensor<float> vreg_input_pre;
|
||||
RegTensor<float> vreg_input_cur;
|
||||
RegTensor<float> vreg_mul;
|
||||
RegTensor<float> vreg_add;
|
||||
RegTensor<float> vreg_div;
|
||||
RegTensor<half> vreg_cast;
|
||||
RegTensor<float> vreg_exp_sum;
|
||||
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
constexpr uint16_t dLoops = srcD / floatRepSize;
|
||||
constexpr float fp8e4m3MaxValueRec = 1 / 448.0f;
|
||||
constexpr float int8MaxValueRec = 1 / 127.0f;
|
||||
constexpr float hifp8MaxValueRec = 1 / 32768.0f;
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize);
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_sum, expSumUb + i * reduceSize);
|
||||
if constexpr (isMlaFullQuant) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_row_max, rowMaxUb + i * reduceSize);
|
||||
}
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
LoadAlign(vreg_input_pre, preUb + i * d + j * floatRepSize);
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
|
||||
if constexpr (isMlaFullQuant) {
|
||||
Mul(vreg_input_cur, vreg_input_cur, vreg_row_max, preg_all);
|
||||
}
|
||||
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_all);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
if constexpr (isUpdatePre) {
|
||||
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
|
||||
}
|
||||
}
|
||||
Add(vreg_add, vreg_mul, vreg_input_cur, preg_all);
|
||||
Div(vreg_div, vreg_add, vreg_exp_sum, preg_all);
|
||||
if constexpr (isMlaFullQuant) {
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e4m3fn_t>::value) {
|
||||
Muls(vreg_div, vreg_div, fp8e4m3MaxValueRec, preg_all);
|
||||
} else if constexpr (IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_div, vreg_div, int8MaxValueRec, preg_all);
|
||||
} else {
|
||||
Muls(vreg_div, vreg_div, hifp8MaxValueRec, preg_all);
|
||||
}
|
||||
}
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_div, preg_all);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize, bool isUpdatePre, bool isMlaFullQuant>
|
||||
__aicore__ inline void FlashUpdateLastBasic(const LocalTensor<T>& dstTensor,
|
||||
const LocalTensor<T>& curTensor, const LocalTensor<T>& preTensor,
|
||||
const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& rowMaxTensor, const LocalTensor<T>& expSumTensor,
|
||||
const uint16_t m, const uint16_t d, const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
|
||||
__ubuf__ float * preUb = (__ubuf__ T*)preTensor.GetPhyAddr();
|
||||
__ubuf__ float * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
|
||||
__ubuf__ float * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
__ubuf__ float * rowMaxUb = (__ubuf__ T*)rowMaxTensor.GetPhyAddr();
|
||||
|
||||
FlashUpdateLastBasicVF<T, INPUT_T, OUTPUT_T, srcD, reduceSize, isUpdatePre, isMlaFullQuant>(
|
||||
dstUb, curUb, preUb, expMaxUb, expSumUb, rowMaxUb, m, d, deScaleV, deScaleVPre);
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t reduceSize, bool isUpdatePre>
|
||||
__simd_vf__ inline void FlashUpdateLastGeneralVF(__ubuf__ float * dstUb, __ubuf__ float * curUb,
|
||||
__ubuf__ float * preUb, __ubuf__ float * expMaxUb, __ubuf__ float * expSumUb, const uint16_t m, const uint16_t d,
|
||||
const float deScaleV, const float deScaleVPre, const uint32_t pltTailD, const uint16_t hasTail)
|
||||
{
|
||||
RegTensor<float> vreg_exp_max;
|
||||
RegTensor<float> vreg_input_pre;
|
||||
RegTensor<float> vreg_input_cur;
|
||||
RegTensor<float> vreg_mul;
|
||||
RegTensor<float> vreg_add;
|
||||
RegTensor<float> vreg_div;
|
||||
RegTensor<half> vreg_cast;
|
||||
RegTensor<float> vreg_exp_sum;
|
||||
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
uint32_t tmpTailD = pltTailD;
|
||||
MaskReg preg_tail_d = UpdateMask<float>(tmpTailD);
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
uint16_t dLoops = d / floatRepSize;
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize);
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_sum, expSumUb + i * reduceSize);
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
LoadAlign(vreg_input_pre, preUb + i * d + j * floatRepSize);
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
|
||||
|
||||
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_all);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
if constexpr (isUpdatePre) {
|
||||
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
|
||||
}
|
||||
}
|
||||
Add(vreg_add, vreg_mul, vreg_input_cur, preg_all);
|
||||
Div(vreg_div, vreg_add, vreg_exp_sum, preg_all);
|
||||
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_div, preg_all);
|
||||
}
|
||||
|
||||
for (uint16_t t = 0; t < hasTail; ++t) {
|
||||
LoadAlign(vreg_input_pre, preUb + i * d + dLoops * floatRepSize);
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + dLoops * floatRepSize);
|
||||
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_tail_d);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
if constexpr (isUpdatePre) {
|
||||
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
|
||||
}
|
||||
}
|
||||
Add(vreg_add, vreg_mul, vreg_input_cur, preg_tail_d);
|
||||
Div(vreg_div, vreg_add, vreg_exp_sum, preg_tail_d);
|
||||
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + dLoops * floatRepSize, vreg_div, preg_tail_d);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t reduceSize, bool isUpdatePre>
|
||||
__aicore__ inline void FlashUpdateLastGeneral(const LocalTensor<T>& dstTensor,
|
||||
const LocalTensor<T>& curTensor, const LocalTensor<T>& preTensor,
|
||||
const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& expSumTensor,
|
||||
const uint16_t m, const uint16_t d, const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
|
||||
__ubuf__ float * preUb = (__ubuf__ T*)preTensor.GetPhyAddr();
|
||||
__ubuf__ float * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
|
||||
__ubuf__ float * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
uint16_t tailD = d % floatRepSize;
|
||||
uint32_t pltTailD = tailD;
|
||||
|
||||
uint16_t hasTail = 0;
|
||||
if (tailD > 0) {
|
||||
hasTail = 1;
|
||||
}
|
||||
|
||||
FlashUpdateLastGeneralVF<T, INPUT_T, OUTPUT_T, reduceSize, isUpdatePre>(
|
||||
dstUb, curUb, preUb, expMaxUb, expSumUb, m, d, deScaleV, deScaleVPre, pltTailD, hasTail);
|
||||
}
|
||||
|
||||
/*
|
||||
* @ingroup FlashUpdateLast
|
||||
* @brief compute, dstTensor = (preTensor * expMaxTensor + curTensor) / expSumTensor
|
||||
* @param [out] dstTensor, output LocalTensor
|
||||
* @param [in] curTensor, input LocalTensor
|
||||
* @param [in] preTensor, input LocalTensor
|
||||
* @param [in] expMaxTensor, input LocalTensor
|
||||
* @param [in] expSumTensor, input LocalTensor
|
||||
* @param [in] m, input rows
|
||||
* @param [in] d, input columns, 32 bytes align
|
||||
*/
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, bool isUpdatePre, bool isMlaFullQuant>
|
||||
__aicore__ inline void FlashUpdateLastNew(const LocalTensor<T>& dstTensor,
|
||||
const LocalTensor<T>& curTensor, const LocalTensor<T>& preTensor,
|
||||
const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& rowMaxTensor, const LocalTensor<T>& expSumTensor,
|
||||
uint16_t m, uint16_t d, const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
static_assert(IsSameType<T, float>::value, "VF FlashUpdateLast, T must be float");
|
||||
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
if constexpr(srcD % floatRepSize == 0) {
|
||||
FlashUpdateLastBasic<T, INPUT_T, OUTPUT_T, srcD, REDUCE_SIZE, isUpdatePre, isMlaFullQuant>(
|
||||
dstTensor, curTensor, preTensor, expMaxTensor, rowMaxTensor, expSumTensor, m, d, deScaleV, deScaleVPre);
|
||||
} else {
|
||||
FlashUpdateLastGeneral<T, INPUT_T, OUTPUT_T, REDUCE_SIZE, isUpdatePre>(
|
||||
dstTensor, curTensor, preTensor, expMaxTensor, expSumTensor, m, d, deScaleV, deScaleVPre);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint32_t srcD, bool isMlaFullQuant>
|
||||
__simd_vf__ inline void LastDivNewVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * expSumUb,
|
||||
const uint16_t m, const uint16_t d, const float deScaleV)
|
||||
{
|
||||
RegTensor<float> vreg_input_cur;
|
||||
RegTensor<float> vreg_div;
|
||||
RegTensor<float> vreg_exp_sum;
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
const uint16_t dLoops = d >> 6;
|
||||
constexpr float fp8e4m3MaxValueRec = 1 / 448.0f;
|
||||
constexpr float int8MaxValueRec = 1 / 127.0f;
|
||||
constexpr float hifp8MaxValueRec = 1 / 32768.0f;
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
uint32_t sreg_init = d;
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_sum, expSumUb + i * REDUCE_SIZE);
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
MaskReg preg_update = UpdateMask<float>(sreg_init);
|
||||
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
}
|
||||
Div(vreg_div, vreg_input_cur, vreg_exp_sum, preg_update);
|
||||
if constexpr (isMlaFullQuant) {
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e4m3fn_t>::value) {
|
||||
Muls(vreg_div, vreg_div, fp8e4m3MaxValueRec, preg_all);
|
||||
} else if constexpr (IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_div, vreg_div, int8MaxValueRec, preg_all);
|
||||
} else {
|
||||
Muls(vreg_div, vreg_div, hifp8MaxValueRec, preg_all);
|
||||
}
|
||||
}
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_div, preg_update);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// dstTensor = curTensor / expSumTensor, curTensor: [64,128], expSumTensor: [64,8]
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint32_t srcD, bool isMlaFullQuant>
|
||||
__aicore__ inline void LastDivNew(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
|
||||
const LocalTensor<T>& expSumTensor, const uint16_t m, const uint16_t d, const float deScaleV)
|
||||
{
|
||||
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
|
||||
__ubuf__ float * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
|
||||
LastDivNewVF<T, INPUT_T, OUTPUT_T, srcD, isMlaFullQuant>(dstUb, curUb, expSumUb, m, d, deScaleV);
|
||||
}
|
||||
|
||||
template <typename T, uint32_t srcD>
|
||||
__simd_vf__ inline void InvalidLineUpdateVF(__ubuf__ T * dstUb, __ubuf__ T * srcUb, __ubuf__ T * maxUb,
|
||||
const uint16_t m, const uint16_t d, const T minValue, const T invalidValue)
|
||||
{
|
||||
RegTensor<float> vreg_invalid_value;
|
||||
RegTensor<float> vreg_max;
|
||||
RegTensor<float> vreg_input;
|
||||
RegTensor<float> vreg_input_brc;
|
||||
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
MaskReg preg_compare;
|
||||
const uint16_t dLoops = d >> 6;
|
||||
|
||||
Duplicate(vreg_invalid_value, invalidValue);
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_max, maxUb + i);
|
||||
Compares<T, CMPMODE::EQ>(preg_compare, vreg_max, minValue, preg_all);
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
LoadAlign(vreg_input, srcUb + i * d + j * floatRepSize);
|
||||
Select(vreg_input_brc, vreg_invalid_value, vreg_input, preg_compare);
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_input_brc, preg_all);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, uint32_t srcD>
|
||||
__aicore__ inline void InvalidLineUpdate(const LocalTensor<T>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& maxTensor, const uint16_t m, const uint16_t d, const T minValue, const T invalidValue)
|
||||
{
|
||||
__ubuf__ T * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
uint16_t dLoops = d >> 6;
|
||||
|
||||
InvalidLineUpdateVF<T, srcD>(dstUb, srcUb, maxUb, m, d, minValue, invalidValue);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__simd_vf__ inline void ComputeLseOutputVF(__ubuf__ T *srcSumUb, __ubuf__ T *srcMaxUb, __ubuf__ T *dstUb, const uint32_t dealCount)
|
||||
{
|
||||
MicroAPI::RegTensor<T> vregSum;
|
||||
MicroAPI::RegTensor<T> vregMax;
|
||||
MicroAPI::RegTensor<T> vregRes;
|
||||
MicroAPI::RegTensor<T> vregResFinal;
|
||||
MicroAPI::RegTensor<float> vregMinValue;
|
||||
MicroAPI::RegTensor<float> vregInfValue;
|
||||
MicroAPI::MaskReg pregCompare;
|
||||
constexpr uint32_t dealRows = 8;
|
||||
constexpr uint32_t floatRepSize = 64; // 64: 一个寄存器存64个float
|
||||
constexpr float infValue = 3e+99; // 3e+99 for float inf
|
||||
constexpr uint32_t tmpMin = 0xFF7FFFFF;
|
||||
float minValue = *((float*)&tmpMin);
|
||||
uint16_t updateLoops = dealCount / dealRows;
|
||||
uint16_t tailLSize = dealCount % dealRows * 8;
|
||||
uint32_t pltTail = static_cast<uint32_t>(tailLSize);
|
||||
|
||||
MicroAPI::MaskReg pregAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregTail = MicroAPI::UpdateMask<T>(pltTail);
|
||||
MicroAPI::Duplicate<float, float>(vregMinValue, minValue);
|
||||
MicroAPI::Duplicate<float, float>(vregInfValue, infValue);
|
||||
|
||||
for (uint16_t i = 0; i < updateLoops; ++i) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(vregSum, srcSumUb + (i * dealRows));
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(vregMax, srcMaxUb + (i * dealRows));
|
||||
|
||||
MicroAPI::Log<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSum, pregAll);
|
||||
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, vregMax, pregAll);
|
||||
|
||||
MicroAPI::Compare<float, CMPMODE::EQ>(pregCompare, vregMax, vregMinValue, pregAll);
|
||||
MicroAPI::Select<T>(vregResFinal, vregInfValue, vregRes, pregCompare);
|
||||
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(dstUb + (i * floatRepSize), vregResFinal, pregAll);
|
||||
}
|
||||
|
||||
if (tailLSize != 0) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(vregSum, srcSumUb + dealRows * updateLoops);
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(vregMax, srcMaxUb + dealRows * updateLoops);
|
||||
|
||||
MicroAPI::Log<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSum, pregTail);
|
||||
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, vregMax, pregTail);
|
||||
|
||||
MicroAPI::Compare<float, CMPMODE::EQ>(pregCompare, vregMax, vregMinValue, pregTail);
|
||||
MicroAPI::Select<T>(vregResFinal, vregInfValue, vregRes, pregCompare);
|
||||
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(dstUb + floatRepSize * updateLoops, vregResFinal, pregTail);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void ComputeLseOutputVF(const LocalTensor<T>& dstTensor, const LocalTensor<T>& softmaxSumTensor,
|
||||
const LocalTensor<T>& softmaxMaxTensor, uint32_t dealCount)
|
||||
{
|
||||
__ubuf__ T * srcSumUb = (__ubuf__ T *)softmaxSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcMaxUb = (__ubuf__ T *)softmaxMaxTensor.GetPhyAddr();
|
||||
__ubuf__ T * dstUb = (__ubuf__ T *)dstTensor.GetPhyAddr();
|
||||
|
||||
ComputeLseOutputVF<T>(srcSumUb, srcMaxUb, dstUb, dealCount);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__simd_vf__ inline void SinkSubExpAddVF(__ubuf__ T *srcSumUb, __ubuf__ T *srcMaxUb, const T sinkValue, const uint32_t dealCount)
|
||||
{
|
||||
MicroAPI::RegTensor<T> vregSum;
|
||||
MicroAPI::RegTensor<T> vregMax;
|
||||
MicroAPI::RegTensor<T> vregRes;
|
||||
MicroAPI::RegTensor<T> vregSink;
|
||||
|
||||
constexpr uint32_t floatRepSize = 64;
|
||||
|
||||
uint16_t updateLoops = dealCount / floatRepSize;
|
||||
uint16_t tailSize = dealCount % floatRepSize;
|
||||
uint32_t pltTail = static_cast<uint32_t>(tailSize);
|
||||
|
||||
//mask
|
||||
MicroAPI::MaskReg pregAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregTail = MicroAPI::UpdateMask<T>(pltTail);
|
||||
|
||||
Duplicate(vregSink, sinkValue);
|
||||
|
||||
for (uint16_t i = 0; i < updateLoops; ++i) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregSum, srcSumUb + (i * floatRepSize));
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregMax, srcMaxUb + (i * floatRepSize));
|
||||
|
||||
MicroAPI::Sub<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSink, vregMax, pregAll);
|
||||
MicroAPI::Exp<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, pregAll);
|
||||
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregSum, vregSum, vregRes, pregAll);
|
||||
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(srcSumUb + (i * floatRepSize), vregSum, pregAll);
|
||||
}
|
||||
|
||||
for (uint16_t i = 0; i < tailSize; i = i + tailSize) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregSum, srcSumUb + (updateLoops * floatRepSize));
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregMax, srcMaxUb + (updateLoops * floatRepSize));
|
||||
|
||||
MicroAPI::Sub<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSink, vregMax, pregTail);
|
||||
MicroAPI::Exp<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, pregTail);
|
||||
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregSum, vregSum, vregRes, pregTail);
|
||||
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(srcSumUb + (updateLoops * floatRepSize), vregSum, pregTail);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void SinkSubExpAddVF(const LocalTensor<T>& softmaxSumTensor, const LocalTensor<T>& softmaxMaxTensor,
|
||||
const T sinkValue, uint32_t dealCount)
|
||||
{
|
||||
__ubuf__ T * srcSumUb = (__ubuf__ T *)softmaxSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcMaxUb = (__ubuf__ T *)softmaxMaxTensor.GetPhyAddr();
|
||||
|
||||
SinkSubExpAddVF<T>(srcSumUb, srcMaxUb, sinkValue, dealCount);
|
||||
}
|
||||
|
||||
template <typename T, typename SINK_T>
|
||||
__simd_vf__ inline void SinkSubExpAddGSFusedVF(__ubuf__ T *srcSumUb, __ubuf__ T *srcMaxUb, __ubuf__ uint16_t *sinkUb, const uint32_t dealCount)
|
||||
{
|
||||
MicroAPI::RegTensor<T> vregSum;
|
||||
MicroAPI::RegTensor<T> vregMax;
|
||||
MicroAPI::RegTensor<T> vregRes;
|
||||
MicroAPI::RegTensor<SINK_T> vregSink;
|
||||
MicroAPI::RegTensor<T> vregSinkCast;
|
||||
|
||||
constexpr uint32_t floatRepSize = 64;
|
||||
|
||||
uint16_t updateLoops = dealCount / floatRepSize;
|
||||
uint16_t tailSize = dealCount % floatRepSize;
|
||||
uint32_t pltTail = static_cast<uint32_t>(tailSize);
|
||||
|
||||
//mask
|
||||
MicroAPI::MaskReg pregAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregTail = MicroAPI::UpdateMask<T>(pltTail);
|
||||
MicroAPI::MaskReg pregSinkAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::LoadAlign<uint16_t, MicroAPI::LoadDist::DIST_UNPACK_B16>((MicroAPI::RegTensor<uint16_t>&)vregSink, sinkUb);
|
||||
MicroAPI::Cast<T, SINK_T, castTraitFp16_32_update>(vregSinkCast, vregSink, pregSinkAll);
|
||||
|
||||
for (uint16_t i = 0; i < updateLoops; ++i) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregSum, srcSumUb + (i * floatRepSize));
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregMax, srcMaxUb + (i * floatRepSize));
|
||||
|
||||
MicroAPI::Sub<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSinkCast, vregMax, pregAll);
|
||||
MicroAPI::Exp<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, pregAll);
|
||||
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregSum, vregSum, vregRes, pregAll);
|
||||
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(srcSumUb + (i * floatRepSize), vregSum, pregAll);
|
||||
}
|
||||
|
||||
if (tailSize != 0) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregSum, srcSumUb + (updateLoops * floatRepSize));
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregMax, srcMaxUb + (updateLoops * floatRepSize));
|
||||
|
||||
MicroAPI::Sub<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSinkCast, vregMax, pregTail);
|
||||
MicroAPI::Exp<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, pregTail);
|
||||
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregSum, vregSum, vregRes, pregTail);
|
||||
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(srcSumUb + (updateLoops * floatRepSize), vregSum, pregTail);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename SINK_T>
|
||||
__aicore__ inline void SinkSubExpAddGSFusedVF(const LocalTensor<SINK_T>& dstTensor, const LocalTensor<T>& softmaxSumTensor,
|
||||
const LocalTensor<T>& softmaxMaxTensor, uint32_t dealCount)
|
||||
{
|
||||
__ubuf__ T * srcSumUb = (__ubuf__ T *)softmaxSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcMaxUb = (__ubuf__ T *)softmaxMaxTensor.GetPhyAddr();
|
||||
__ubuf__ uint16_t * dstUb = (__ubuf__ uint16_t *)dstTensor.GetPhyAddr();
|
||||
|
||||
SinkSubExpAddGSFusedVF<T, SINK_T>(srcSumUb, srcMaxUb, dstUb, dealCount);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__simd_vf__ inline void RowInvalidUpdateVF(__ubuf__ T *finalUb, __ubuf__ float *maxUb, const uint16_t m,
|
||||
const uint16_t d, int64_t dSize, const uint32_t pltTailD, const uint16_t hasTail)
|
||||
{
|
||||
constexpr uint16_t floatRepSize = 64; // 64: 一个寄存器可以存储64个float类型数据
|
||||
const uint16_t dLoops = d / floatRepSize;
|
||||
|
||||
|
||||
constexpr uint32_t tmpZero = 0x00000000; // zero value of fp16 and fp32
|
||||
const T zeroValue = *((T*)&tmpZero);
|
||||
constexpr uint32_t tmpMin = 0xFF7FFFFF; // min value of float
|
||||
const float minValue = *((float*)&tmpMin);
|
||||
MicroAPI::RegTensor<float> vregMinValue;
|
||||
MicroAPI::RegTensor<T> vregZeroValue;
|
||||
MicroAPI::RegTensor<float> vregMax;
|
||||
MicroAPI::RegTensor<T> vregFinal;
|
||||
MicroAPI::RegTensor<T> vregFinalNew;
|
||||
|
||||
MicroAPI::MaskReg pregAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
|
||||
uint32_t tmpTailD = pltTailD;
|
||||
MicroAPI::MaskReg pregTailD = MicroAPI::UpdateMask<T>(tmpTailD);
|
||||
MicroAPI::MaskReg pregCompare;
|
||||
|
||||
MicroAPI::Duplicate<float, float>(vregMinValue, minValue);
|
||||
MicroAPI::Duplicate<T, T>(vregZeroValue, zeroValue);
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(vregMax, maxUb + i);
|
||||
MicroAPI::Compare<float, CMPMODE::EQ>(pregCompare, vregMax, vregMinValue, pregAll);
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregFinal, finalUb + i * dSize + j * floatRepSize);
|
||||
MicroAPI::Select<T>(vregFinalNew, vregZeroValue, vregFinal, pregCompare);
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(finalUb + i * dSize + j * floatRepSize,
|
||||
vregFinalNew, pregAll);
|
||||
}
|
||||
for (uint16_t t = 0; t < hasTail; ++t) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregFinal, finalUb + i * dSize + dLoops * floatRepSize);
|
||||
MicroAPI::Select<T>(vregFinalNew, vregZeroValue, vregFinal, pregCompare);
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(finalUb + i * dSize + dLoops * floatRepSize,
|
||||
vregFinalNew, pregTailD);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void RowInvalidUpdateVF(const LocalTensor<T>& finalTensor, const LocalTensor<float>& maxTensor,
|
||||
const uint16_t m, const uint16_t d, int64_t dSize)
|
||||
{
|
||||
__ubuf__ T * finalUb = (__ubuf__ T*)finalTensor.GetPhyAddr();
|
||||
__ubuf__ float * maxUb = (__ubuf__ float*)maxTensor.GetPhyAddr();
|
||||
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
const uint16_t tailD = d % floatRepSize;
|
||||
uint32_t pltTailD = static_cast<uint32_t>(tailD);
|
||||
uint16_t hasTail = 0;
|
||||
if (tailD > 0) {
|
||||
hasTail = 1;
|
||||
}
|
||||
|
||||
RowInvalidUpdateVF<T>(finalUb, maxUb, m, d, dSize, pltTailD, hasTail);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // MY_FLASH_UPDATE_INTERFACE_H
|
||||
@@ -0,0 +1,164 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_mul_sel_softmaxflashv2_cast_nz_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef MUL_SEL_SOFTMAX_FLASH_V2_CAST_NZ_SFA_INTERFACE_H
|
||||
#define MUL_SEL_SOFTMAX_FLASH_V2_CAST_NZ_SFA_INTERFACE_H
|
||||
|
||||
#include "vf_basic_block_aligned128_no_update_sfa.h"
|
||||
#include "vf_basic_block_aligned128_update_sfa.h"
|
||||
#include "vf_basic_block_unaligned64_update_sfa.h"
|
||||
#include "vf_basic_block_unaligned64_no_update_sfa.h"
|
||||
#include "vf_basic_block_unaligned128_no_update_sfa.h"
|
||||
#include "vf_basic_block_unaligned128_update_sfa.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
/* **************************************************************************************************
|
||||
* Muls + Select(optional) + SoftmaxFlashV2 + Cast(fp32->fp16/bf16) + ND2NZ
|
||||
* ************************************************************************************************* */
|
||||
using AscendC::LocalTensor;
|
||||
|
||||
enum class OriginNRange {
|
||||
EQ_128_SFA = 0, // originN == 128, better performance than GT_64_AND_LTE_128 (s2BaseSize=128)
|
||||
GT_0_AND_LTE_64_SFA, // 0 < originN <= 64 (s2BaseSize <= 64 or tail s2)
|
||||
GT_64_AND_LTE_128_SFA, // 64 < originN <= 128, support for non-alignment (s2BaseSize=128)
|
||||
N_INVALID_SFA
|
||||
};
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128,
|
||||
OriginNRange oriNRange = OriginNRange::EQ_128_SFA>
|
||||
__aicore__ inline void ProcessVec1NoUpdate(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
if constexpr (oriNRange == OriginNRange::EQ_128_SFA) {
|
||||
ProcessVec1NoUpdateImpl128<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
} else if constexpr (oriNRange == OriginNRange::GT_0_AND_LTE_64_SFA) {
|
||||
ProcessVec1NoUpdateImpl64<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
} else if constexpr (oriNRange == OriginNRange::GT_64_AND_LTE_128_SFA) {
|
||||
ProcessVec1NoUpdateGeneralImpl128<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128,
|
||||
OriginNRange oriNRange = OriginNRange::EQ_128_SFA>
|
||||
__aicore__ inline void ProcessVec1Update(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
if constexpr (oriNRange == OriginNRange::EQ_128_SFA) {
|
||||
ProcessVec1UpdateImpl128<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
} else if constexpr (oriNRange == OriginNRange::GT_0_AND_LTE_64_SFA) {
|
||||
ProcessVec1UpdateImpl64<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
} else if constexpr (oriNRange == OriginNRange::GT_64_AND_LTE_128_SFA) {
|
||||
ProcessVec1UpdateGeneralImpl128<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename T2, bool isUpdate = false, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128,
|
||||
OriginNRange oriNRange = OriginNRange::EQ_128_SFA>
|
||||
__aicore__ inline void ProcessVec1Vf(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
static_assert(IsSameType<T, float>::value, "VF mul_sel_softmaxflashv2_cast_nz, T must be float");
|
||||
static_assert((IsSameType<T2, half>::value || IsSameType<T2, bfloat16_t>::value),
|
||||
"VF mul_sel_softmaxflashv2_cast_nz, T2 must be half or bfloat16");
|
||||
|
||||
if constexpr (!isUpdate) {
|
||||
ProcessVec1NoUpdate<T, T2, s1BaseSize, s2BaseSize, oriNRange>(
|
||||
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
} else {
|
||||
ProcessVec1Update<T, T2, s1BaseSize, s2BaseSize, oriNRange>(
|
||||
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__simd_vf__ inline void UpdateExpSumAndExpMaxVF(__ubuf__ T * maxUb, __ubuf__ T * inMaxUb, __ubuf__ T * expMaxUb,
|
||||
__ubuf__ T * expSumUb, __ubuf__ T * inExpSumUb, __ubuf__ T * tmpExpSumUb, __ubuf__ T * tmpMaxUb, const uint32_t m)
|
||||
{
|
||||
RegTensor<float> vreg_input_x;
|
||||
RegTensor<float> vreg_input_x_unroll;
|
||||
RegTensor<float> vreg_max;
|
||||
RegTensor<float> vreg_in_max;
|
||||
RegTensor<float> vreg_exp_sum;
|
||||
RegTensor<float> vreg_in_exp_sum;
|
||||
RegTensor<float> vreg_exp_max;
|
||||
RegTensor<float> vreg_exp_sum_brc;
|
||||
RegTensor<float> vreg_exp_sum_update;
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
// 注意:当m大于64的时候需要开启循环
|
||||
LoadAlign(vreg_max, tmpMaxUb);
|
||||
LoadAlign(vreg_in_max, inMaxUb);
|
||||
FusedExpSub(vreg_exp_max, vreg_in_max, vreg_max, preg_all);
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)expMaxUb, vreg_exp_max, preg_all);
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)maxUb, vreg_max, preg_all);
|
||||
LoadAlign(vreg_in_exp_sum, inExpSumUb);
|
||||
|
||||
// x_sum = exp_max * insum + x_sum
|
||||
LoadAlign(vreg_exp_sum_brc, tmpExpSumUb);
|
||||
Mul(vreg_exp_sum_update, vreg_exp_max, vreg_in_exp_sum, preg_all);
|
||||
Add(vreg_exp_sum_update, vreg_exp_sum_update, vreg_exp_sum_brc, preg_all);
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)expSumUb, vreg_exp_sum_update, preg_all);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void SFAUpdateExpSumAndExpMax(
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor,
|
||||
const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& inExpSumTensor,
|
||||
const LocalTensor<T>& inMaxTensor, const LocalTensor<T>& sharedTmpBuffer, const uint32_t m)
|
||||
{
|
||||
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * inMaxUb = (__ubuf__ T*)inMaxTensor.GetPhyAddr();
|
||||
|
||||
__ubuf__ T * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
|
||||
__ubuf__ T * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * inExpSumUb = (__ubuf__ T*)inExpSumTensor.GetPhyAddr();
|
||||
|
||||
__ubuf__ T * tmpExpSumUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr();
|
||||
__ubuf__ T * tmpMaxUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
|
||||
UpdateExpSumAndExpMaxVF<T>(maxUb, inMaxUb, expMaxUb, expSumUb, inExpSumUb, tmpExpSumUb, tmpMaxUb, m);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__simd_vf__ inline void DuplicateSumWithR0VF(__ubuf__ T * sumUb, const T R0, uint32_t m) {
|
||||
AscendC::MicroAPI::RegTensor<T> vreg_sum;
|
||||
AscendC::MicroAPI::MaskReg preg_m = AscendC::MicroAPI::UpdateMask<T>(m);
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg;
|
||||
AscendC::MicroAPI::Duplicate<T, MicroAPI::MaskMergeMode::ZEROING, T>(vreg_sum, R0, preg_m);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(sumUb, vreg_sum, preg_m);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void DuplicateSumWithR0(const LocalTensor<T>& sumTensor, const T R0, uint32_t m)
|
||||
{
|
||||
__ubuf__ T * sumUb = (__ubuf__ T*)sumTensor.GetPhyAddr();
|
||||
DuplicateSumWithR0VF<T>(sumUb, R0, m);
|
||||
}
|
||||
} // namespace
|
||||
#endif // MUL_SEL_SOFTMAX_FLASH_V2_CAST_NZ_SFA_INTERFACE_H
|
||||
Reference in New Issue
Block a user