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,378 @@
/**
 * 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_add.h
* \brief
*/
#ifndef VF_ADD_H
#define VF_ADD_H
#include "kernel_operator.h"
#include <cstdint>
using namespace AscendC;
constexpr uint32_t FLOAT_REP_SIZE = 64;
constexpr uint32_t BTYEALIGNSIZE = 32;
constexpr uint32_t REGSIZE = 256;
constexpr uint32_t HALFCORED = 128;
template <typename T>
struct AddRegList {
MicroAPI::RegTensor<T> vreg;
MicroAPI::RegTensor<T> vregape;
};
template <typename T>
__simd_callee__ void AddVFImpl(__ubuf__ T *inputAddr, __ubuf__ T *apeAddr, AddRegList<T> &regList, uint32_t row,
uint32_t col, uint64_t offset0, uint64_t offset1)
{
uint32_t maskValue = col;
MicroAPI::MaskReg mask = MicroAPI::UpdateMask<T>(maskValue);
MicroAPI::LoadAlign(regList.vreg, inputAddr + offset0);
MicroAPI::LoadAlign(regList.vregape, apeAddr + offset1);
MicroAPI::Add(regList.vreg, regList.vreg, regList.vregape, mask);
MicroAPI::StoreAlign(inputAddr + offset0, regList.vreg, mask);
}
template <bool IS_FIRST, typename T>
__simd_callee__ void MultiAddVFImpl(__ubuf__ T *outputAddr, __ubuf__ T *inputAddr, AddRegList<T> &regList, uint32_t row,
uint32_t col, uint64_t offset, uint32_t repeatNum, uint64_t repeatOffset)
{
uint32_t maskValue = col;
uint32_t initialRepeatIdx = IS_FIRST ? 1 : 0;
__ubuf__ T *initialAddr = IS_FIRST ? inputAddr : outputAddr;
MicroAPI::MaskReg mask = MicroAPI::UpdateMask<T>(maskValue);
MicroAPI::LoadAlign(regList.vreg, initialAddr + offset);
for (uint32_t repeatIdx = initialRepeatIdx; repeatIdx < repeatNum; repeatIdx++) {
uint64_t addOffset = offset + repeatIdx * repeatOffset;
MicroAPI::LoadAlign(regList.vregape, inputAddr + addOffset);
MicroAPI::Add(regList.vreg, regList.vreg, regList.vregape, mask);
}
MicroAPI::StoreAlign(outputAddr + offset, regList.vreg, mask);
}
template <typename T>
__simd_vf__ void Add64VFImpl(__ubuf__ T *inputAddr, __ubuf__ T *apeAddr, uint32_t row, uint32_t col, uint32_t actualCol0, uint32_t actualCol1)
{
AddRegList<T> regList[4];
uint32_t loopTimes = row / 4;
for (uint32_t idx = 0; idx < loopTimes; idx++) {
uint64_t offset0 = idx * 4 * actualCol0;
uint64_t offset1 = idx * 4 * actualCol1;
AddVFImpl(inputAddr, apeAddr, regList[0], row, col, offset0, offset1);
AddVFImpl(inputAddr, apeAddr, regList[1], row, col, offset0 + actualCol0, offset1 + actualCol1);
AddVFImpl(inputAddr, apeAddr, regList[2], row, col, offset0 + 2 * actualCol0, offset1 + 2 * actualCol1);
AddVFImpl(inputAddr, apeAddr, regList[3], row, col, offset0 + 3 * actualCol0, offset1 + 3 * actualCol1);
}
if (row % 4 > 0) {
AddVFImpl(inputAddr, apeAddr, regList[0], row, col, loopTimes * 4 * actualCol0, loopTimes * 4 * actualCol1);
}
if (row % 4 > 1) {
AddVFImpl(inputAddr, apeAddr, regList[1], row, col, (loopTimes * 4 + 1) * actualCol0, (loopTimes * 4 + 1) * actualCol1);
}
if (row % 4 > 2) {
AddVFImpl(inputAddr, apeAddr, regList[2], row, col, (loopTimes * 4 + 2) * actualCol0, (loopTimes * 4 + 2) * actualCol1);
}
}
template <typename T>
__simd_vf__ void Add128VFImpl(__ubuf__ T *inputAddr, __ubuf__ T *apeAddr, uint32_t row, uint32_t actualCol0, uint32_t actualCol1)
{
AddRegList<T> regList[4];
uint32_t loopTimes = row / 2;
for (uint32_t idx = 0; idx < loopTimes; idx++) {
uint64_t offset0 = idx * 2 * actualCol0;
uint64_t offset1 = idx * 2 * actualCol1;
AddVFImpl(inputAddr, apeAddr, regList[0], row, FLOAT_REP_SIZE, offset0, offset1);
AddVFImpl(inputAddr, apeAddr, regList[1], row, FLOAT_REP_SIZE, offset0 + FLOAT_REP_SIZE, offset1 + FLOAT_REP_SIZE);
AddVFImpl(inputAddr, apeAddr, regList[2], row, FLOAT_REP_SIZE, offset0 + actualCol0, offset1 + actualCol1);
AddVFImpl(inputAddr, apeAddr, regList[3], row, FLOAT_REP_SIZE, offset0 + actualCol0 + FLOAT_REP_SIZE, offset1 + actualCol1 + FLOAT_REP_SIZE);
}
if (row % 2 > 0) {
AddVFImpl(inputAddr, apeAddr, regList[0], row, FLOAT_REP_SIZE, loopTimes * 2 * actualCol0, loopTimes * 2 * actualCol1);
AddVFImpl(inputAddr, apeAddr, regList[1], row, FLOAT_REP_SIZE, loopTimes * 2 * actualCol0 + FLOAT_REP_SIZE, loopTimes * 2 * actualCol1 + FLOAT_REP_SIZE);
}
}
template <typename T>
__simd_vf__ void Add256VFImpl(__ubuf__ T *inputAddr, __ubuf__ T *apeAddr, uint32_t row, uint32_t actualCol0, uint32_t actualCol1)
{
AddRegList<T> regList[4];
MicroAPI::MaskReg mask = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
for (uint32_t idx = 0; idx < row; idx++) {
uint64_t offset0 = idx * actualCol0;
uint64_t offset1 = idx * actualCol1;
AddVFImpl(inputAddr, apeAddr, regList[0], row, FLOAT_REP_SIZE, offset0, offset1);
AddVFImpl(inputAddr, apeAddr, regList[1], row, FLOAT_REP_SIZE, offset0 + FLOAT_REP_SIZE, offset1 + FLOAT_REP_SIZE);
AddVFImpl(inputAddr, apeAddr, regList[2], row, FLOAT_REP_SIZE, offset0 + 2 * FLOAT_REP_SIZE, offset1 + 2 * FLOAT_REP_SIZE);
AddVFImpl(inputAddr, apeAddr, regList[3], row, FLOAT_REP_SIZE, offset0 + 3 * FLOAT_REP_SIZE, offset1 + 3 * FLOAT_REP_SIZE);
}
}
template <typename T>
__simd_vf__ void Add512VFImpl(__ubuf__ T *inputAddr, __ubuf__ T *apeAddr, uint32_t row, uint32_t actualCol0, uint32_t actualCol1)
{
AddRegList<T> regList[8];
MicroAPI::MaskReg mask = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
for (uint32_t idx = 0; idx < row; idx++) {
uint64_t offset0 = idx * actualCol0;
uint64_t offset1 = idx * actualCol1;
AddVFImpl(inputAddr, apeAddr, regList[0], row, FLOAT_REP_SIZE, offset0, offset1);
AddVFImpl(inputAddr, apeAddr, regList[1], row, FLOAT_REP_SIZE, offset0 + FLOAT_REP_SIZE, offset1 + FLOAT_REP_SIZE);
AddVFImpl(inputAddr, apeAddr, regList[2], row, FLOAT_REP_SIZE, offset0 + 2 * FLOAT_REP_SIZE, offset1 + 2 * FLOAT_REP_SIZE);
AddVFImpl(inputAddr, apeAddr, regList[3], row, FLOAT_REP_SIZE, offset0 + 3 * FLOAT_REP_SIZE, offset1 + 3 * FLOAT_REP_SIZE);
AddVFImpl(inputAddr, apeAddr, regList[4], row, FLOAT_REP_SIZE, offset0 + 4 * FLOAT_REP_SIZE, offset1 + 4 * FLOAT_REP_SIZE);
AddVFImpl(inputAddr, apeAddr, regList[5], row, FLOAT_REP_SIZE, offset0 + 5 * FLOAT_REP_SIZE, offset1 + 5 * FLOAT_REP_SIZE);
AddVFImpl(inputAddr, apeAddr, regList[6], row, FLOAT_REP_SIZE, offset0 + 6 * FLOAT_REP_SIZE, offset1 + 6 * FLOAT_REP_SIZE);
AddVFImpl(inputAddr, apeAddr, regList[7], row, FLOAT_REP_SIZE, offset0 + 7 * FLOAT_REP_SIZE, offset1 + 7 * FLOAT_REP_SIZE);
}
}
template <bool IS_FIRST, typename T>
__simd_vf__ void MultiAdd64VFImpl(__ubuf__ T *outputAddr, __ubuf__ T *inputAddr, uint32_t row, uint32_t col,
uint32_t actualCol, uint32_t repeatNum, uint64_t repeatOffset)
{
AddRegList<T> regList[4];
uint32_t loopTimes = row / 4;
uint32_t maskValue = col;
uint32_t initialRepeatIdx = IS_FIRST ? 1 : 0;
__ubuf__ T *initialAddr = IS_FIRST ? inputAddr : outputAddr;
MicroAPI::MaskReg mask = MicroAPI::UpdateMask<T>(maskValue);
for (uint32_t idx = 0; idx < loopTimes; idx++) {
uint64_t offset = idx * 4 * actualCol;
MicroAPI::LoadAlign(regList[0].vreg, initialAddr + offset);
MicroAPI::LoadAlign(regList[1].vreg, initialAddr + offset + actualCol);
MicroAPI::LoadAlign(regList[2].vreg, initialAddr + offset + 2 * actualCol);
MicroAPI::LoadAlign(regList[3].vreg, initialAddr + offset + 3 * actualCol);
for (uint32_t repeatIdx = initialRepeatIdx; repeatIdx < repeatNum; repeatIdx++) {
uint64_t addOffset = offset + repeatIdx * repeatOffset;
MicroAPI::LoadAlign(regList[0].vregape, inputAddr + addOffset);
MicroAPI::LoadAlign(regList[1].vregape, inputAddr + addOffset + actualCol);
MicroAPI::LoadAlign(regList[2].vregape, inputAddr + addOffset + 2 * actualCol);
MicroAPI::LoadAlign(regList[3].vregape, inputAddr + addOffset + 3 * actualCol);
MicroAPI::Add(regList[0].vreg, regList[0].vreg, regList[0].vregape, mask);
MicroAPI::Add(regList[1].vreg, regList[1].vreg, regList[1].vregape, mask);
MicroAPI::Add(regList[2].vreg, regList[2].vreg, regList[2].vregape, mask);
MicroAPI::Add(regList[3].vreg, regList[3].vreg, regList[3].vregape, mask);
}
MicroAPI::StoreAlign(outputAddr + offset, regList[0].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + actualCol, regList[1].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + 2 * actualCol, regList[2].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + 3 * actualCol, regList[3].vreg, mask);
}
if (row % 4 > 0) {
MultiAddVFImpl<IS_FIRST, T>(outputAddr, inputAddr, regList[0], row, col, loopTimes * 4 * actualCol, repeatNum,
repeatOffset);
}
if (row % 4 > 1) {
MultiAddVFImpl<IS_FIRST, T>(outputAddr, inputAddr, regList[1], row, col, (loopTimes * 4 + 1) * actualCol,
repeatNum, repeatOffset);
}
if (row % 4 > 2) {
MultiAddVFImpl<IS_FIRST, T>(outputAddr, inputAddr, regList[2], row, col, (loopTimes * 4 + 2) * actualCol,
repeatNum, repeatOffset);
}
}
template <bool IS_FIRST, typename T>
__simd_vf__ void MultiAdd128VFImpl(__ubuf__ T *outputAddr, __ubuf__ T *inputAddr, uint32_t row, uint32_t col,
uint32_t actualCol, uint32_t repeatNum, uint64_t repeatOffset)
{
AddRegList<T> regList[4];
uint32_t loopTimes = row / 2;
uint32_t initialRepeatIdx = IS_FIRST ? 1 : 0;
__ubuf__ T *initialAddr = IS_FIRST ? inputAddr : outputAddr;
MicroAPI::MaskReg mask = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
for (uint32_t idx = 0; idx < loopTimes; idx++) {
uint64_t offset = idx * actualCol * 2;
MicroAPI::LoadAlign(regList[0].vreg, initialAddr + offset);
MicroAPI::LoadAlign(regList[1].vreg, initialAddr + offset + FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[2].vreg, initialAddr + offset + actualCol);
MicroAPI::LoadAlign(regList[3].vreg, initialAddr + offset + actualCol + FLOAT_REP_SIZE);
for (uint32_t repeatIdx = initialRepeatIdx; repeatIdx < repeatNum; repeatIdx++) {
uint64_t addOffset = offset + repeatIdx * repeatOffset;
MicroAPI::LoadAlign(regList[0].vregape, inputAddr + addOffset);
MicroAPI::LoadAlign(regList[1].vregape, inputAddr + addOffset + FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[2].vregape, inputAddr + addOffset + actualCol);
MicroAPI::LoadAlign(regList[3].vregape, inputAddr + addOffset + actualCol + FLOAT_REP_SIZE);
MicroAPI::Add(regList[0].vreg, regList[0].vreg, regList[0].vregape, mask);
MicroAPI::Add(regList[1].vreg, regList[1].vreg, regList[1].vregape, mask);
MicroAPI::Add(regList[2].vreg, regList[2].vreg, regList[2].vregape, mask);
MicroAPI::Add(regList[3].vreg, regList[3].vreg, regList[3].vregape, mask);
}
MicroAPI::StoreAlign(outputAddr + offset, regList[0].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + FLOAT_REP_SIZE, regList[1].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + actualCol, regList[2].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + actualCol + FLOAT_REP_SIZE, regList[3].vreg, mask);
}
if (row % 2 > 0) {
MultiAddVFImpl<IS_FIRST, T>(outputAddr, inputAddr, regList[0], row, col, loopTimes * 2 * actualCol, repeatNum,
repeatOffset);
MultiAddVFImpl<IS_FIRST, T>(outputAddr, inputAddr, regList[1], row, col,
loopTimes * 2 * actualCol + FLOAT_REP_SIZE, repeatNum, repeatOffset);
}
}
template <bool IS_FIRST, typename T>
__simd_vf__ void MultiAdd256VFImpl(__ubuf__ T *outputAddr, __ubuf__ T *inputAddr, uint32_t row,
uint32_t actualCol, uint32_t repeatNum, uint64_t repeatOffset)
{
AddRegList<T> regList[4];
uint32_t loopTimes = row;
uint32_t initialRepeatIdx = IS_FIRST ? 1 : 0;
__ubuf__ T *initialAddr = IS_FIRST ? inputAddr : outputAddr;
MicroAPI::MaskReg mask = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
for (uint32_t idx = 0; idx < loopTimes; idx++) {
uint64_t offset = idx * actualCol;
MicroAPI::LoadAlign(regList[0].vreg, initialAddr + offset);
MicroAPI::LoadAlign(regList[1].vreg, initialAddr + offset + FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[2].vreg, initialAddr + offset + 2 * FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[3].vreg, initialAddr + offset + 3 * FLOAT_REP_SIZE);
for (uint32_t repeatIdx = initialRepeatIdx; repeatIdx < repeatNum; repeatIdx++) {
uint64_t addOffset = offset + repeatIdx * repeatOffset;
MicroAPI::LoadAlign(regList[0].vregape, inputAddr + addOffset);
MicroAPI::LoadAlign(regList[1].vregape, inputAddr + addOffset + FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[2].vregape, inputAddr + addOffset + 2 * FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[3].vregape, inputAddr + addOffset + 3 * FLOAT_REP_SIZE);
MicroAPI::Add(regList[0].vreg, regList[0].vreg, regList[0].vregape, mask);
MicroAPI::Add(regList[1].vreg, regList[1].vreg, regList[1].vregape, mask);
MicroAPI::Add(regList[2].vreg, regList[2].vreg, regList[2].vregape, mask);
MicroAPI::Add(regList[3].vreg, regList[3].vreg, regList[3].vregape, mask);
}
MicroAPI::StoreAlign(outputAddr + offset, regList[0].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + FLOAT_REP_SIZE, regList[1].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + 2 * FLOAT_REP_SIZE, regList[2].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + 3 * FLOAT_REP_SIZE, regList[3].vreg, mask);
}
}
template <bool IS_FIRST, typename T>
__simd_vf__ void MultiAdd512VFImpl(__ubuf__ T *outputAddr, __ubuf__ T *inputAddr, uint32_t row,
uint32_t actualCol, uint32_t repeatNum, uint64_t repeatOffset)
{
AddRegList<T> regList[8];
uint32_t loopTimes = row;
uint32_t initialRepeatIdx = IS_FIRST ? 1 : 0;
__ubuf__ T *initialAddr = IS_FIRST ? inputAddr : outputAddr;
MicroAPI::MaskReg mask = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
for (uint32_t idx = 0; idx < loopTimes; idx++) {
uint64_t offset = idx * actualCol;
MicroAPI::LoadAlign(regList[0].vreg, initialAddr + offset);
MicroAPI::LoadAlign(regList[1].vreg, initialAddr + offset + FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[2].vreg, initialAddr + offset + 2 * FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[3].vreg, initialAddr + offset + 3 * FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[4].vreg, initialAddr + offset + 4 * FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[5].vreg, initialAddr + offset + 5 * FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[6].vreg, initialAddr + offset + 6 * FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[7].vreg, initialAddr + offset + 7 * FLOAT_REP_SIZE);
for (uint32_t repeatIdx = initialRepeatIdx; repeatIdx < repeatNum; repeatIdx++) {
uint64_t addOffset = offset + repeatIdx * row * actualCol;
MicroAPI::LoadAlign(regList[0].vregape, inputAddr + addOffset);
MicroAPI::LoadAlign(regList[1].vregape, inputAddr + addOffset + FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[2].vregape, inputAddr + addOffset + 2 * FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[3].vregape, inputAddr + addOffset + 3 * FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[4].vregape, inputAddr + addOffset + 4 * FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[5].vregape, inputAddr + addOffset + 5 * FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[6].vregape, inputAddr + addOffset + 6 * FLOAT_REP_SIZE);
MicroAPI::LoadAlign(regList[7].vregape, inputAddr + addOffset + 7 * FLOAT_REP_SIZE);
MicroAPI::Add(regList[0].vreg, regList[0].vreg, regList[0].vregape, mask);
MicroAPI::Add(regList[1].vreg, regList[1].vreg, regList[1].vregape, mask);
MicroAPI::Add(regList[2].vreg, regList[2].vreg, regList[2].vregape, mask);
MicroAPI::Add(regList[3].vreg, regList[3].vreg, regList[3].vregape, mask);
MicroAPI::Add(regList[4].vreg, regList[4].vreg, regList[4].vregape, mask);
MicroAPI::Add(regList[5].vreg, regList[5].vreg, regList[5].vregape, mask);
MicroAPI::Add(regList[6].vreg, regList[6].vreg, regList[6].vregape, mask);
MicroAPI::Add(regList[7].vreg, regList[7].vreg, regList[7].vregape, mask);
}
MicroAPI::StoreAlign(outputAddr + offset, regList[0].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + FLOAT_REP_SIZE, regList[1].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + 2 * FLOAT_REP_SIZE, regList[2].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + 3 * FLOAT_REP_SIZE, regList[3].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + 4 * FLOAT_REP_SIZE, regList[4].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + 5 * FLOAT_REP_SIZE, regList[5].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + 6 * FLOAT_REP_SIZE, regList[6].vreg, mask);
MicroAPI::StoreAlign(outputAddr + offset + 7 * FLOAT_REP_SIZE, regList[7].vreg, mask);
}
}
/**
* @brief AddVF 输入与apt相加
* @param rightLocal 输出tensor []
* @param leftLocal 输入tensor [row, col]
* @param aptLocal apt输入tensor [r]
* @param apeIdx ape起始位置
* @param d coff*d为ape的D轴大小
* @param coreSplitD scoreleft大小,coff*coreSplitD为总大小
* @param coreSplitS 核间d轴切分大小
*/
template <typename T>
__aicore__ inline void AddVF(const LocalTensor<T> &scoreLocal, const LocalTensor<T> &apeLocal, uint32_t row,
uint32_t col, uint32_t actualCol0, uint32_t actualCol1)
{
__ubuf__ T *scoreAddr = (__ubuf__ T *)scoreLocal.GetPhyAddr();
__ubuf__ T *apeAddr = (__ubuf__ T *)apeLocal.GetPhyAddr();
if (col <= 64) {
Add64VFImpl<T>(scoreAddr, apeAddr, row, col, actualCol0, actualCol1);
} else if (col == 128) {
Add128VFImpl<T>(scoreAddr, apeAddr, row, actualCol0, actualCol1);
} else if (col == 256) {
Add256VFImpl<T>(scoreAddr, apeAddr, row, actualCol0, actualCol1);
} else if (col == 512) {
Add512VFImpl<T>(scoreAddr, apeAddr, row, actualCol0, actualCol1);
}
}
template <typename T>
__aicore__ inline void AddVF(const LocalTensor<T> &scoreLocal, const LocalTensor<T> &apeLocal, uint32_t row,
uint32_t col, uint32_t actualCol)
{
__ubuf__ T *scoreAddr = (__ubuf__ T *)scoreLocal.GetPhyAddr();
__ubuf__ T *apeAddr = (__ubuf__ T *)apeLocal.GetPhyAddr();
if (col <= 64) {
Add64VFImpl<T>(scoreAddr, apeAddr, row, col, actualCol, actualCol);
} else if (col == 128) {
Add128VFImpl<T>(scoreAddr, apeAddr, row, actualCol, actualCol);
} else if (col == 256) {
Add256VFImpl<T>(scoreAddr, apeAddr, row, actualCol, actualCol);
} else if (col == 512) {
Add512VFImpl<T>(scoreAddr, apeAddr, row, actualCol, actualCol);
}
}
template <bool IS_FIRST, typename T>
__aicore__ inline void MultiAddVF(const LocalTensor<T> &outputLocal, const LocalTensor<T> &inputLocal, uint32_t row,
uint32_t col, uint32_t actualCol, uint32_t repeatNum, uint64_t repeatOffset)
{
__ubuf__ T *outputAddr = (__ubuf__ T *)outputLocal.GetPhyAddr();
__ubuf__ T *inputAddr = (__ubuf__ T *)inputLocal.GetPhyAddr();
if (col <= 64) {
MultiAdd64VFImpl<IS_FIRST, T>(outputAddr, inputAddr, row, col, actualCol, repeatNum, repeatOffset);
} else if (col == 128) {
MultiAdd128VFImpl<IS_FIRST, T>(outputAddr, inputAddr, row, col, actualCol, repeatNum, repeatOffset);
} else if (col == 256) {
MultiAdd256VFImpl<IS_FIRST, T>(outputAddr, inputAddr, row, actualCol, repeatNum, repeatOffset);
} else if (col == 512) {
MultiAdd512VFImpl<IS_FIRST, T>(outputAddr, inputAddr, row, actualCol, repeatNum, repeatOffset);
}
}
#endif

View File

@@ -0,0 +1,318 @@
/**
 * 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.h
* \brief
*/
#ifndef VF_MUL_H
#define VF_MUL_H
#include "kernel_operator.h"
#include <cstdint>
using namespace AscendC;
constexpr uint32_t FLOATBYTE = 4;
constexpr uint32_t baseD8 = 8;
constexpr uint32_t baseD16 = 16;
constexpr uint32_t baseD32 = 32;
constexpr uint32_t baseD64 = 64;
constexpr uint32_t baseD128 = 128;
constexpr uint32_t baseD256 = 256;
constexpr uint32_t baseD512 = 512;
template <typename T>
__simd_callee__ inline T SimdCeilDivT(T num1, T num2)
{
if (num2 == 0) {
return static_cast<T>(0);
}
return (num1 + num2 - 1) / num2;
}
template <typename T>
struct ReduceMulRegList {
MicroAPI::RegTensor<T> vreg0;
MicroAPI::RegTensor<T> vreg1;
MicroAPI::RegTensor<T> vregMul;
MicroAPI::RegTensor<T> vregSum;
};
template <typename T>
__simd_callee__ void LoadMulAddVFImpl(__ubuf__ T *kvAddr, __ubuf__ T *scoreAddr, ReduceMulRegList<T> &regList, uint64_t offset, uint32_t maskValue)
{
MicroAPI::MaskReg mask = MicroAPI::UpdateMask<T>(maskValue);
MicroAPI::LoadAlign(regList.vreg0, kvAddr + offset);
MicroAPI::LoadAlign(regList.vreg1, scoreAddr + offset);
MicroAPI::Mul(regList.vregMul, regList.vreg0, regList.vreg1, mask);
MicroAPI::Add(regList.vregSum, regList.vregSum, regList.vregMul, mask);
}
template <typename T>
__simd_vf__ void MulReduceSumbase8VFImpl(__ubuf__ T *kvAddr, __ubuf__ T *scoreAddr, __ubuf__ T *outputAddr,
const uint32_t coff, const uint32_t cmpRatio, const uint32_t scLoopCnt,
const uint32_t baseD)
{
ReduceMulRegList<T> regList;
MicroAPI::RegTensor<T> vregSum0;
MicroAPI::MaskReg mask = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
MicroAPI::MaskReg maskL32 = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::VL32>();
MicroAPI::MaskReg maskL16 = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::VL16>();
MicroAPI::MaskReg maskL8 = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::VL8>();
MicroAPI::MaskReg maskH32;
MicroAPI::MaskReg maskH48;
MicroAPI::MaskReg maskH56;
MicroAPI::Not(maskH48, maskL16, mask);
MicroAPI::Not(maskH32, maskL32, mask);
MicroAPI::Not(maskH56, maskL8, mask);
uint32_t offset = 0;
uint32_t rCnt = coff * cmpRatio;
for (uint32_t scLoop = 0; scLoop < scLoopCnt; scLoop++) {
MicroAPI::Duplicate(regList.vregSum, 0, mask);
// 当前仅支持coff * cmpRatio为2的幂的情况
for (uint32_t rLoop = 0; rLoop < SimdCeilDivT(rCnt, 8U); rLoop++) {
uint32_t dealLen = min((rCnt - rLoop * 8) * baseD, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList, offset, dealLen);
offset += dealLen;
}
// 64 -> 32
MicroAPI::Squeeze<T, AscendC::MicroAPI::GatherMaskMode::NO_STORE_REG>(vregSum0, regList.vregSum, maskH32);
MicroAPI::Add(regList.vregSum, regList.vregSum, vregSum0, maskL32);
// 32 -> 16
MicroAPI::Squeeze<T, AscendC::MicroAPI::GatherMaskMode::NO_STORE_REG>(vregSum0, regList.vregSum, maskH48);
MicroAPI::Add(regList.vregSum, regList.vregSum, vregSum0, maskL16);
// 16 -> 8
MicroAPI::Squeeze<T, AscendC::MicroAPI::GatherMaskMode::NO_STORE_REG>(vregSum0, regList.vregSum, maskH56);
MicroAPI::Add(regList.vregSum, regList.vregSum, vregSum0, maskL8);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD, regList.vregSum, maskL8);
}
}
template <typename T>
__simd_vf__ void MulReduceSumbase16VFImpl(__ubuf__ T *kvAddr, __ubuf__ T *scoreAddr, __ubuf__ T *outputAddr,
const uint32_t coff, const uint32_t cmpRatio, const uint32_t scLoopCnt,
const uint32_t baseD)
{
ReduceMulRegList<T> regList;
MicroAPI::RegTensor<T> vregSum0;
MicroAPI::MaskReg mask = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
MicroAPI::MaskReg maskL32 = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::VL32>();
MicroAPI::MaskReg maskL16 = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::VL16>();
MicroAPI::MaskReg maskH32;
MicroAPI::MaskReg maskH48;
MicroAPI::Not(maskH48, maskL16, mask);
MicroAPI::Not(maskH32, maskL32, mask);
uint32_t offset = 0;
uint32_t rCnt = coff * cmpRatio;
for (uint32_t scLoop = 0; scLoop < scLoopCnt; scLoop++) {
MicroAPI::Duplicate(regList.vregSum, 0, mask);
// 当前仅支持coff * cmpRatio为2的幂的情况
for (uint32_t rLoop = 0; rLoop < SimdCeilDivT(rCnt, 4U); rLoop++) {
uint32_t dealLen = min((rCnt - rLoop * 4) * baseD, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList, offset, dealLen);
offset += dealLen;
}
// 64 -> 32
MicroAPI::Squeeze<T, AscendC::MicroAPI::GatherMaskMode::NO_STORE_REG>(vregSum0, regList.vregSum, maskH32);
MicroAPI::Add(regList.vregSum, regList.vregSum, vregSum0, maskL32);
// 32 -> 16
MicroAPI::Squeeze<T, AscendC::MicroAPI::GatherMaskMode::NO_STORE_REG>(vregSum0, regList.vregSum, maskH48);
MicroAPI::Add(regList.vregSum, regList.vregSum, vregSum0, maskL16);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD, regList.vregSum, maskL16);
}
}
template <typename T>
__simd_vf__ void MulReduceSumbase32VFImpl(__ubuf__ T *kvAddr, __ubuf__ T *scoreAddr, __ubuf__ T *outputAddr,
const uint32_t coff, const uint32_t cmpRatio, const uint32_t scLoopCnt,
const uint32_t baseD)
{
ReduceMulRegList<T> regList;
MicroAPI::RegTensor<T> vregSum0;
MicroAPI::RegTensor<T> vregSum1;
MicroAPI::MaskReg mask = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
MicroAPI::MaskReg maskL32 = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::VL32>();
MicroAPI::MaskReg maskH32;
MicroAPI::Not(maskH32, maskL32, mask);
uint32_t offset = 0;
uint32_t rCnt = coff * cmpRatio;
for (uint32_t scLoop = 0; scLoop < scLoopCnt; scLoop++) {
MicroAPI::Duplicate(regList.vregSum, 0, mask);
// 当前仅支持coff * cmpRatio为2的幂的情况
for (uint32_t rLoop = 0; rLoop < SimdCeilDivT(rCnt, 2U); rLoop++) {
uint32_t dealLen = min((rCnt - rLoop * 2) * baseD, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList, offset, dealLen);
offset += dealLen;
}
// 64 -> 32
MicroAPI::Squeeze<T, AscendC::MicroAPI::GatherMaskMode::NO_STORE_REG>(vregSum0, regList.vregSum, maskH32);
MicroAPI::Add(regList.vregSum, regList.vregSum, vregSum0, maskL32);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD, regList.vregSum, maskL32);
}
}
template <typename T>
__simd_vf__ void MulReduceSumbase64VFImpl(__ubuf__ T *kvAddr, __ubuf__ T *scoreAddr, __ubuf__ T *outputAddr,
const uint32_t coff, const uint32_t cmpRatio, const uint32_t scLoopCnt,
const uint32_t baseD)
{
ReduceMulRegList<T> regList;
MicroAPI::MaskReg mask = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
uint32_t offset = 0;
uint32_t rCnt = coff * cmpRatio;
for (uint32_t scLoop = 0; scLoop < scLoopCnt; scLoop++) {
MicroAPI::Duplicate(regList.vregSum, 0, mask);
for (uint32_t rLoop = 0; rLoop < rCnt; rLoop++) {
LoadMulAddVFImpl(kvAddr, scoreAddr, regList, offset, baseD64);
offset += baseD;
}
MicroAPI::StoreAlign(outputAddr + scLoop * baseD, regList.vregSum, mask);
}
}
template <typename T>
__simd_vf__ void MulReduceSumbase128VFImpl(__ubuf__ T *kvAddr, __ubuf__ T *scoreAddr, __ubuf__ T *outputAddr,
const uint32_t coff, const uint32_t cmpRatio, const uint32_t scLoopCnt,
const uint32_t baseD)
{
ReduceMulRegList<T> regList[2];
MicroAPI::MaskReg mask = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
uint32_t offset = 0;
uint32_t rCnt = coff * cmpRatio;
for (uint32_t scLoop = 0; scLoop < scLoopCnt; scLoop++) {
MicroAPI::Duplicate(regList[0].vregSum, 0, mask);
MicroAPI::Duplicate(regList[1].vregSum, 0, mask);
for (uint32_t rLoop = 0; rLoop < rCnt; rLoop++) {
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[0], offset, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[1], offset + baseD64, baseD64);
offset += baseD;
}
MicroAPI::StoreAlign(outputAddr + scLoop * baseD, regList[0].vregSum, mask);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD + baseD64, regList[1].vregSum, mask);
}
}
template <typename T>
__simd_vf__ void MulReduceSumbase256VFImpl(__ubuf__ T *kvAddr, __ubuf__ T *scoreAddr, __ubuf__ T *outputAddr,
const uint32_t coff, const uint32_t cmpRatio, const uint32_t scLoopCnt,
const uint32_t baseD)
{
ReduceMulRegList<T> regList[4];
MicroAPI::MaskReg mask = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
uint32_t offset = 0;
uint32_t rCnt = coff * cmpRatio;
for (uint32_t scLoop = 0; scLoop < scLoopCnt; scLoop++) {
MicroAPI::Duplicate(regList[0].vregSum, 0, mask);
MicroAPI::Duplicate(regList[1].vregSum, 0, mask);
MicroAPI::Duplicate(regList[2].vregSum, 0, mask);
MicroAPI::Duplicate(regList[3].vregSum, 0, mask);
for (uint32_t rLoop = 0; rLoop < rCnt; rLoop++) {
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[0], offset, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[1], offset + baseD64, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[2], offset + 2 * baseD64, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[3], offset + 3 * baseD64, baseD64);
offset += baseD;
}
MicroAPI::StoreAlign(outputAddr + scLoop * baseD, regList[0].vregSum, mask);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD + baseD64, regList[1].vregSum, mask);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD + 2 * baseD64, regList[2].vregSum, mask);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD + 3 * baseD64, regList[3].vregSum, mask);
}
}
template <typename T>
__simd_vf__ void MulReduceSumbase512VFImpl(__ubuf__ T *kvAddr, __ubuf__ T *scoreAddr, __ubuf__ T *outputAddr,
const uint32_t coff, const uint32_t cmpRatio, const uint32_t scLoopCnt,
const uint32_t baseD)
{
ReduceMulRegList<T> regList[8];
MicroAPI::MaskReg mask = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
uint32_t offset = 0;
uint32_t rCnt = coff * cmpRatio;
for (uint32_t scLoop = 0; scLoop < scLoopCnt; scLoop++) {
MicroAPI::Duplicate(regList[0].vregSum, 0, mask);
MicroAPI::Duplicate(regList[1].vregSum, 0, mask);
MicroAPI::Duplicate(regList[2].vregSum, 0, mask);
MicroAPI::Duplicate(regList[3].vregSum, 0, mask);
MicroAPI::Duplicate(regList[4].vregSum, 0, mask);
MicroAPI::Duplicate(regList[5].vregSum, 0, mask);
MicroAPI::Duplicate(regList[6].vregSum, 0, mask);
MicroAPI::Duplicate(regList[7].vregSum, 0, mask);
for (uint32_t rLoop = 0; rLoop < rCnt; rLoop++) {
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[0], offset, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[1], offset + baseD64, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[2], offset + 2 * baseD64, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[3], offset + 3 * baseD64, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[4], offset + 4 * baseD64, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[5], offset + 5 * baseD64, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[6], offset + 6 * baseD64, baseD64);
LoadMulAddVFImpl(kvAddr, scoreAddr, regList[7], offset + 7 * baseD64, baseD64);
offset += baseD;
}
MicroAPI::StoreAlign(outputAddr + scLoop * baseD, regList[0].vregSum, mask);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD + baseD64, regList[1].vregSum, mask);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD + 2 * baseD64, regList[2].vregSum, mask);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD + 3 * baseD64, regList[3].vregSum, mask);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD + 4 * baseD64, regList[4].vregSum, mask);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD + 5 * baseD64, regList[5].vregSum, mask);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD + 6 * baseD64, regList[6].vregSum, mask);
MicroAPI::StoreAlign(outputAddr + scLoop * baseD + 7 * baseD64, regList[7].vregSum, mask);
}
}
/**
* @brief MulReduceSumbaseVF 包含mul和reducesum
* @param outputLocal 输出tensor []
* @param coff
* @param cmpRatio 压缩块大小
* @param baseD 核内d轴切分大小
* @param scLoopCnt sc数,
*/
// 当前仅支持coff * cmpRatio为2的幂的情况
template <typename T>
__aicore__ inline void MulReduceSumbaseVF(const LocalTensor<T> &kvLocal, const LocalTensor<T> &scoreLocal,
const LocalTensor<T> &outputLocal, const uint32_t coff, const uint32_t cmpRatio,
const uint32_t baseD, const uint32_t scLoopCnt)
{
__ubuf__ T *kvAddr = (__ubuf__ T *)kvLocal.GetPhyAddr();
__ubuf__ T *scoreAddr = (__ubuf__ T *)scoreLocal.GetPhyAddr();
__ubuf__ T *outputAddr = (__ubuf__ T *)outputLocal.GetPhyAddr();
if (baseD == baseD8) {
MulReduceSumbase8VFImpl(kvAddr, scoreAddr, outputAddr, coff, cmpRatio, scLoopCnt, baseD);
} else if (baseD == baseD16) {
MulReduceSumbase16VFImpl(kvAddr, scoreAddr, outputAddr, coff, cmpRatio, scLoopCnt, baseD);
} else if (baseD == baseD32) {
MulReduceSumbase32VFImpl(kvAddr, scoreAddr, outputAddr, coff, cmpRatio, scLoopCnt, baseD);
} else if (baseD == baseD64) {
MulReduceSumbase64VFImpl(kvAddr, scoreAddr, outputAddr, coff, cmpRatio, scLoopCnt, baseD);
} else if (baseD == baseD128) {
MulReduceSumbase128VFImpl(kvAddr, scoreAddr, outputAddr, coff, cmpRatio, scLoopCnt, baseD);
} else if (baseD == baseD256) {
MulReduceSumbase256VFImpl(kvAddr, scoreAddr, outputAddr, coff, cmpRatio, scLoopCnt, baseD);
} else if (baseD == baseD512) {
MulReduceSumbase512VFImpl(kvAddr, scoreAddr, outputAddr, coff, cmpRatio, scLoopCnt, baseD);
}
}
#endif

View File

@@ -0,0 +1,95 @@
/**
 * 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_rms_norm.h
* \brief
*/
#ifndef VF_RMS_NORM_H
#define VF_RMS_NORM_H
#include "kernel_tensor.h"
//repeatTimes——D轴的分块数
template <typename T, typename GammaType>
__simd_vf__ void RmsNormVFImpl(__ubuf__ T * inputBuf, __ubuf__ GammaType * gammaBuf, __ubuf__ T * outputBuf,
uint32_t repeatTimes, float reciprocal, float epsilon)
{
MicroAPI::RegTensor<T> vregSum;
MicroAPI::RegTensor<T> vregSumReduce;
MicroAPI::RegTensor<T> vregDiv;
MicroAPI::RegTensor<T> vregSquareRoot;
MicroAPI::MaskReg maskAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
MicroAPI::MaskReg maskFirst = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::VL1>();
static constexpr MicroAPI::CastTrait castTraitB162B32 = {MicroAPI::RegLayout::ZERO,
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
MicroAPI::Duplicate<T,T>(vregSum, 0.0f);
for(uint32_t i = 0; i < repeatTimes; ++i){
MicroAPI::RegTensor<T> vregX;
MicroAPI::RegTensor<T> vregXSquare;
uint64_t loopOffset = i * FLOAT_REP_SIZE;
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregX, inputBuf + loopOffset);
MicroAPI::Mul(vregXSquare, vregX, vregX, maskAll);
MicroAPI::Add(vregSum, vregXSquare, vregSum, maskAll);
}
MicroAPI::Reduce<MicroAPI::ReduceType::SUM, T, T, MicroAPI::MaskMergeMode::ZEROING>(vregSumReduce, vregSum, maskAll);
MicroAPI::Muls<T, T, MicroAPI::MaskMergeMode::ZEROING>(vregSumReduce, vregSumReduce, reciprocal, maskFirst);
MicroAPI::Adds<T, T, MicroAPI::MaskMergeMode::ZEROING>(vregSumReduce, vregSumReduce, epsilon, maskFirst);
MicroAPI::Sqrt(vregSquareRoot, vregSumReduce, maskFirst);
MicroAPI::Duplicate<T, MicroAPI::HighLowPart::LOWEST, MicroAPI::MaskMergeMode::ZEROING>(vregDiv, vregSquareRoot, maskAll);
for(uint32_t i = 0; i < repeatTimes; ++i){
MicroAPI::RegTensor<T> vregX;
MicroAPI::RegTensor<T> vregGammaCast;
uint16_t loopOffset = i * FLOAT_REP_SIZE;
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregX, inputBuf + loopOffset);
MicroAPI::LoadAlign<GammaType, MicroAPI::LoadDist::DIST_NORM>(vregGammaCast, gammaBuf + loopOffset);
MicroAPI::Div(vregX, vregX, vregDiv, maskAll);
MicroAPI::Mul(vregX, vregX, vregGammaCast, maskAll);
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM>(outputBuf + loopOffset, vregX, maskAll);
}
}
/**
* @brief RmsNormVF 对一行进行rmsnorm
* @param outputLocal 输出tensor [row, col],row目前均为1
* @param inputLocal 输入tensor [row, col]
* @param gammaLocal gamma参数tensor [row, col]
* @param rmsNormParams rmsNrom计算所需系数,包括
row 行数 1
col 列数,对应headSizeCq或headSizeCkv
reciprocal ,1/N
epsilon,防止除零极小数
*/
template <typename T, typename GammaType>
__aicore__ inline void RmsNormVF(const LocalTensor<T> outputLocal, const LocalTensor<T> inputLocal, const LocalTensor<GammaType> gammaLocal,
float reciprocal, float epsilon, uint32_t row, uint32_t col)
{
uint32_t cnt = row * col;
uint32_t repeatTimes = (cnt + FLOAT_REP_SIZE - 1) / FLOAT_REP_SIZE;
__ubuf__ T * inputBuf = (__ubuf__ T *)inputLocal.GetPhyAddr();
__ubuf__ GammaType * gammaBuf = (__ubuf__ GammaType *)gammaLocal.GetPhyAddr();
__ubuf__ T * outputBuf = (__ubuf__ T *)outputLocal.GetPhyAddr();
RmsNormVFImpl<T, GammaType>(inputBuf, gammaBuf, outputBuf, repeatTimes, reciprocal, epsilon);
}
#endif

View File

@@ -0,0 +1,158 @@
/**
 * 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_rope.h
* \brief
*/
#ifndef VF_ROPE_H
#define VF_ROPE_H
#include "kernel_operator.h"
#include "../compressor_comm.h"
using namespace AscendC;
constexpr MicroAPI::CastTrait castTraitB162B32 = {
MicroAPI::RegLayout::ZERO,
MicroAPI::SatMode::UNKNOWN,
MicroAPI::MaskMergeMode::ZEROING,
RoundMode::UNKNOWN,
};
constexpr MicroAPI::CastTrait castTraitB322B16 = {
MicroAPI::RegLayout::ZERO,
MicroAPI::SatMode::NO_SAT,
MicroAPI::MaskMergeMode::ZEROING,
RoundMode::CAST_RINT,
};
template <typename T, typename ROPET>
__simd_vf__ void HalfModeRopeVF(__ubuf__ T *sinUb, __ubuf__ T *cosUb, __ubuf__ T *inUb, __ubuf__ ROPET *outUb,
uint32_t row, uint32_t col, uint32_t actualCol, uint64_t baseAddr)
{
MicroAPI::RegTensor<T> vregCos;
MicroAPI::RegTensor<T> vregHalfCos;
MicroAPI::RegTensor<T> vregSin;
MicroAPI::RegTensor<T> vregHalfSin;
MicroAPI::RegTensor<T> vregIn;
MicroAPI::RegTensor<T> vregHalfIn;
MicroAPI::RegTensor<T> vregOut;
MicroAPI::RegTensor<T> vregHalfOut;
MicroAPI::RegTensor<T> vregCastIn;
MicroAPI::RegTensor<ROPET> vregOutBf16;
MicroAPI::RegTensor<ROPET> vregOutHalfBf16;
MicroAPI::RegTensor<ROPET> vregCastOut;
uint32_t maskValue = col / 2;
MicroAPI::MaskReg mask = MicroAPI::UpdateMask<T>(maskValue);
uint32_t halfCol = col / 2;
for (uint32_t rIdx = 0; rIdx < row; rIdx++) {
__ubuf__ T *curSinUb = sinUb + rIdx * col;
__ubuf__ T *curCosUb = cosUb + rIdx * col;
__ubuf__ T *curInUb = inUb + rIdx * actualCol;
__ubuf__ ROPET *curOutUb = outUb + rIdx * actualCol;
MicroAPI::DataCopy(vregIn, curInUb + baseAddr);
MicroAPI::DataCopy(vregHalfIn, curInUb + baseAddr + halfCol);
MicroAPI::DataCopy(vregCos, curCosUb);
MicroAPI::DataCopy(vregHalfCos, curCosUb + halfCol);
MicroAPI::DataCopy(vregSin, curSinUb);
MicroAPI::DataCopy(vregHalfSin, curSinUb + halfCol);
MicroAPI::Mul(vregSin, vregSin, vregHalfIn, mask);
MicroAPI::Mul(vregHalfSin, vregHalfSin, vregIn, mask);
MicroAPI::Mul(vregCos, vregCos, vregIn, mask);
MicroAPI::Sub(vregOut, vregCos, vregSin, mask);
MicroAPI::Mul(vregHalfCos, vregHalfCos, vregHalfIn, mask);
MicroAPI::Add(vregHalfOut, vregHalfSin, vregHalfCos, mask);
MicroAPI::Cast<ROPET, T, castTraitB322B16>(vregOutBf16, vregOut, mask);
MicroAPI::DataCopy<ROPET, MicroAPI::StoreDist::DIST_PACK_B32>(curOutUb + baseAddr, vregOutBf16, mask);
MicroAPI::Cast<ROPET, T, castTraitB322B16>(vregOutHalfBf16, vregHalfOut, mask);
MicroAPI::DataCopy<ROPET, MicroAPI::StoreDist::DIST_PACK_B32>(curOutUb + baseAddr + halfCol, vregOutHalfBf16,
mask);
for (uint64_t dOffset = 0; dOffset < baseAddr; dOffset += 64) {
uint32_t castMaskValue = min(baseAddr - dOffset, static_cast<uint64_t>(64));
MicroAPI::MaskReg castMask = MicroAPI::UpdateMask<T>(castMaskValue);
MicroAPI::DataCopy(vregCastIn, curInUb + dOffset);
MicroAPI::Cast<ROPET, T, castTraitB322B16>(vregCastOut, vregCastIn, castMask);
MicroAPI::DataCopy<ROPET, MicroAPI::StoreDist::DIST_PACK_B32>(curOutUb + dOffset, vregCastOut, castMask);
}
}
}
template <typename T, typename ROPET>
__simd_vf__ void InterleaveModeRopeVF(__ubuf__ T *sinUb, __ubuf__ T *cosUb, __ubuf__ T *inUb, __ubuf__ ROPET *outUb,
uint32_t row, uint32_t col, uint32_t actualCol, uint64_t baseAddr)
{
MicroAPI::RegTensor<T> vregCos;
MicroAPI::RegTensor<T> vregSin;
MicroAPI::RegTensor<T> vregIn;
MicroAPI::RegTensor<T> vregOdd;
MicroAPI::RegTensor<T> vregEven;
MicroAPI::RegTensor<T> vregOut;
MicroAPI::RegTensor<T> vregTemp;
MicroAPI::RegTensor<T> vregCastIn;
MicroAPI::RegTensor<ROPET> vregOutBf16;
MicroAPI::RegTensor<ROPET> vregCastOut;
uint32_t maskValue = col;
MicroAPI::MaskReg mask = MicroAPI::UpdateMask<T>(maskValue);
for (uint32_t rIdx = 0; rIdx < row; rIdx++) {
__ubuf__ T *curSinUb = sinUb + rIdx * col;
__ubuf__ T *curCosUb = cosUb + rIdx * col;
__ubuf__ T *curInUb = inUb + rIdx * actualCol;
__ubuf__ ROPET *curOutUb = outUb + rIdx * actualCol;
MicroAPI::DataCopy(vregIn, curInUb + baseAddr);
MicroAPI::DataCopy(vregCos, curCosUb);
MicroAPI::DataCopy(vregSin, curSinUb);
MicroAPI::Mul(vregCos, vregCos, vregIn, mask);
MicroAPI::DeInterleave<T>(vregEven, vregOdd, vregIn, vregTemp);
MicroAPI::Muls(vregOdd, vregOdd, static_cast<T>(-1.0), mask);
MicroAPI::Interleave<T>(vregIn, vregTemp, vregOdd, vregEven);
MicroAPI::Mul(vregSin, vregSin, vregIn, mask);
MicroAPI::Add(vregOut, vregCos, vregSin, mask);
MicroAPI::Cast<ROPET, T, castTraitB322B16>(vregOutBf16, vregOut, mask);
MicroAPI::DataCopy<ROPET, MicroAPI::StoreDist::DIST_PACK_B32>(curOutUb + baseAddr, vregOutBf16, mask);
for (uint64_t dOffset = 0; dOffset < baseAddr; dOffset += 64) {
uint32_t castMaskValue = min(baseAddr - dOffset, static_cast<uint64_t>(64));
MicroAPI::MaskReg castMask = MicroAPI::UpdateMask<T>(castMaskValue);
MicroAPI::DataCopy(vregCastIn, curInUb + dOffset);
MicroAPI::Cast<ROPET, T, castTraitB322B16>(vregCastOut, vregCastIn, castMask);
MicroAPI::DataCopy<ROPET, MicroAPI::StoreDist::DIST_PACK_B32>(curOutUb + dOffset, vregCastOut, castMask);
}
}
}
template <Compressor::ROTARY_MODE MODE, typename T, typename ROPET>
__aicore__ inline void RopeVF(const LocalTensor<T> &sinTensor, const LocalTensor<T> &cosTensor,
const LocalTensor<T> &inTensor, const LocalTensor<ROPET> &outTensor, uint32_t row,
uint32_t col, uint32_t actualCol, uint64_t baseAddr)
{
__ubuf__ T *sinUb = (__ubuf__ T *)sinTensor.GetPhyAddr();
__ubuf__ T *cosUb = (__ubuf__ T *)cosTensor.GetPhyAddr();
__ubuf__ T *inUb = (__ubuf__ T *)inTensor.GetPhyAddr();
__ubuf__ ROPET *outUb = (__ubuf__ ROPET *)outTensor.GetPhyAddr();
if constexpr (MODE == Compressor::ROTARY_MODE::HALF) {
HalfModeRopeVF(sinUb, cosUb, inUb, outUb, row, col, actualCol, baseAddr);
} else {
InterleaveModeRopeVF(sinUb, cosUb, inUb, outUb, row, col, actualCol, baseAddr);
}
}
#endif

File diff suppressed because it is too large Load Diff