378
csrc/attention/compressor/op_kernel/arch35/vf/vf_add.h
Normal file
378
csrc/attention/compressor/op_kernel/arch35/vf/vf_add.h
Normal 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> ®List, 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> ®List, 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
|
||||
318
csrc/attention/compressor/op_kernel/arch35/vf/vf_mul.h
Normal file
318
csrc/attention/compressor/op_kernel/arch35/vf/vf_mul.h
Normal 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> ®List, 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
|
||||
95
csrc/attention/compressor/op_kernel/arch35/vf/vf_rms_norm.h
Normal file
95
csrc/attention/compressor/op_kernel/arch35/vf/vf_rms_norm.h
Normal 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
|
||||
158
csrc/attention/compressor/op_kernel/arch35/vf/vf_rope.h
Normal file
158
csrc/attention/compressor/op_kernel/arch35/vf/vf_rope.h
Normal 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
|
||||
1592
csrc/attention/compressor/op_kernel/arch35/vf/vf_softmax.h
Normal file
1592
csrc/attention/compressor/op_kernel/arch35/vf/vf_softmax.h
Normal file
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user