init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View 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

View 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

View File

@@ -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