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,514 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file CopyInL1.h
* \brief
*/
#ifndef COPYINL1_H
#define COPYINL1_H
enum class KVLAYOUT
{
BNBD, // [blockNums, headNum, blockSize, headDim]
BBH, // [blockNums, blockSize, headNum * headDim]
NZ // [blockNums, headNum, d1, blockSize, d0], d1 = headDim / d0, d0 = 32 (block byte) / sizeof(KV_T)
};
struct CopyParam{
uint32_t width;
uint32_t height;
uint32_t orgWidth;
};
struct PAShape{
uint32_t blockNum;
uint32_t blockSize;
uint32_t headNum; // 一般为kv的head num
uint32_t headDim; // mla下rope为64, 非rope为512
uint32_t maxblockNumPerBatch; // block table 每一行的最大个数
uint32_t actHeadDim; // 实际拷贝col大小,考虑到N切块 s*d, 对应d
uint32_t copyRowNum;
uint32_t copyRowNumAlign;
uint32_t pageStride;
};
struct Position{
uint32_t bIdx;
uint32_t n2Idx;
uint32_t s2Offset;
uint32_t dIdx; // N轴被切,对应D轴被切
};
template<typename L1Type>
__aicore__ inline void GmCopyInToL1(LocalTensor<L1Type>& L1Tensor, GlobalTensor<L1Type>& GmTensor, const CopyParam& mmCopyParam)
{
Nd2NzParams Gm2L1Nd2NzParams;
Gm2L1Nd2NzParams.ndNum = 1; // ND矩阵的个数
Gm2L1Nd2NzParams.nValue = mmCopyParam.height; // 单个ND矩阵的实际行数,单位为元素个数
Gm2L1Nd2NzParams.dValue = mmCopyParam.width; // 单个ND矩阵的实际列数(vD),单位为元素个数
Gm2L1Nd2NzParams.srcNdMatrixStride = 0; // 相邻ND矩阵起始地址之间的偏移, 单位为元素个数
Gm2L1Nd2NzParams.srcDValue = mmCopyParam.orgWidth; // 同一个ND矩阵中相邻行起始地址之间的偏移, 单位为元素个数
Gm2L1Nd2NzParams.dstNzC0Stride = (Gm2L1Nd2NzParams.nValue + 15) >> 4 << 4; // 转换为NZ矩阵后,相邻Block起始地址之间的偏移, 单位为Block个数
Gm2L1Nd2NzParams.dstNzNStride = 1; // 转换为NZ矩阵后,ND之间相邻两行在NZ矩阵中起始地址之间的偏移, 单位为Block个数
Gm2L1Nd2NzParams.dstNzMatrixStride = 0; // 两个NZ矩阵,起始地址之间的偏移, 单位为元素数量
DataCopy(L1Tensor, GmTensor, Gm2L1Nd2NzParams);
}
// 场景:key、value GM to L1
// GM按ND格式存储
// L1按NZ格式存储
// GM的行、列、列的stride(D or ND)BNSD 和 BSH的区别
template<typename L1Type>
__aicore__ inline void DataCopyGmNDToL1(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
uint32_t rowAct,
uint32_t rowAlign,
uint32_t col, // D
uint32_t colStride) // D or N*D
{
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = rowAct; // 行数
nd2nzPara.dValue = col;
nd2nzPara.srcDValue = colStride;
nd2nzPara.dstNzC0Stride = rowAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(l1Tensor, gmTensor, nd2nzPara);
}
template<typename L1Type>
__aicore__ inline void DataCopyGmScaleNDToL1(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
uint32_t rowAct,
uint32_t rowAlign,
uint32_t col, // D
uint32_t colStride) // D or N*D
{
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = rowAct;
nd2nzPara.dValue = col;
nd2nzPara.srcDValue = colStride;
nd2nzPara.dstNzC0Stride = rowAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = nd2nzPara.nValue;
LocalTensor<bfloat16_t> l1TensorCast = l1Tensor.template ReinterpretCast<bfloat16_t>();
GlobalTensor<bfloat16_t> gmTensorCast;
gmTensorCast.SetGlobalBuffer(((__gm__ bfloat16_t*)(gmTensor.GetPhyAddr())));
DataCopy(l1TensorCast, gmTensorCast, nd2nzPara);
}
template<typename L1Type>
__aicore__ inline void DataCopyGmScaleDNToL1(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
uint32_t rowAct,
uint32_t rowAlign,
uint32_t col,
uint32_t colStride)
{
Dn2NzParams dn2nzPara;
dn2nzPara.dnNum = 1;
dn2nzPara.nValue = col / 2;
dn2nzPara.dValue = rowAct;
dn2nzPara.srcDValue = colStride / 2;
dn2nzPara.dstNzC0Stride = dn2nzPara.nValue;
dn2nzPara.dstNzNStride = 1;
dn2nzPara.srcDnMatrixStride = 0;
dn2nzPara.dstNzMatrixStride = dn2nzPara.nValue;
LocalTensor<bfloat16_t> l1TensorCast = l1Tensor.template ReinterpretCast<bfloat16_t>();
GlobalTensor<bfloat16_t> gmTensorCast;
gmTensorCast.SetGlobalBuffer(((__gm__ bfloat16_t*)(gmTensor.GetPhyAddr())));
DataCopy(l1TensorCast, gmTensorCast, dn2nzPara);
}
template<typename L1Type>
__aicore__ inline void DataCopyGmNZToL1(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
uint32_t rowAct,
uint32_t dstRowStride,
uint32_t srcRowStride,
uint32_t col)
{
uint32_t blockElementCnt = 32U / sizeof(L1Type);
if constexpr (IsSameType<L1Type, int4b_t>::value) {
blockElementCnt = 64U;
}
DataCopyParams intriParams;
intriParams.blockCount = col / blockElementCnt;
intriParams.blockLen = rowAct;
intriParams.dstStride = dstRowStride;
intriParams.srcStride = srcRowStride;
DataCopy(l1Tensor, gmTensor, intriParams);
}
template<typename L1Type>
__aicore__ inline void GmCopyInToL1HasRopePANoContinue(LocalTensor<L1Type>& nopeTensor, LocalTensor<L1Type>& ropeTensor,
GlobalTensor<L1Type>& nopeGmTensor, GlobalTensor<L1Type>& ropeGmTensor,
GlobalTensor<int32_t>& blockTableGm, KVLAYOUT kvLayout,
const PAShape &shape,
const PAShape &ropeShape,
const Position &startPos)
{
uint32_t copyFinishRowCnt = 0;
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch; // 块表的基偏移量
uint32_t curS2Idx = startPos.s2Offset;
uint32_t blockElementCnt = 32U / sizeof(L1Type); // 每个块的元素数量
// ropeshape的M方向与nopeshape保持一样, 此处只判断nopeshape的
while(copyFinishRowCnt < shape.copyRowNum){
uint64_t blockIdOffset = curS2Idx / shape.blockSize; // 获取block table上的索引
uint64_t remainRowCnt = curS2Idx % shape.blockSize; // 获取在单个块上超出的行数
uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset); // 从block table上获取的编号
//计算可以拷贝行数
uint32_t copyRowCnt = shape.blockSize - remainRowCnt; // 一次只能处理一个Block
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum){
copyRowCnt = shape.copyRowNum - copyFinishRowCnt; // 一个block未拷满
}
uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim; // PA的偏移
if (shape.pageStride > 0) {
offset = idInBlockTable * shape.pageStride;
}
uint64_t keyRopeOffset = idInBlockTable * ropeShape.blockSize * ropeShape.headNum * ropeShape.headDim;
if (ropeShape.pageStride > 0) {
keyRopeOffset = idInBlockTable * ropeShape.pageStride;
}
if (kvLayout == KVLAYOUT::NZ) {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.blockSize * shape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * shape.blockSize;
keyRopeOffset += static_cast<uint64_t>(startPos.n2Idx * ropeShape.blockSize * ropeShape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * ropeShape.blockSize;
LocalTensor<L1Type> tmpNopeDstTensor = nopeTensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpNopeSrcTensor = nopeGmTensor[offset];
DataCopyGmNZToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, (shape.copyRowNumAlign - copyRowCnt), (shape.blockSize - copyRowCnt), shape.actHeadDim);
LocalTensor<L1Type> tmpRopeDstTensor = ropeTensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpRopeSrcTensor = ropeGmTensor[keyRopeOffset];
DataCopyGmNZToL1(tmpRopeDstTensor, tmpRopeSrcTensor, copyRowCnt, (ropeShape.copyRowNumAlign - copyRowCnt), (ropeShape.blockSize - copyRowCnt), ropeShape.actHeadDim);
} else {
uint64_t dStride = shape.headDim;
uint64_t dRopeStride = ropeShape.headDim;
if (kvLayout == KVLAYOUT::BBH) {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim) + remainRowCnt * shape.headDim * shape.headNum + startPos.dIdx;
keyRopeOffset += static_cast<uint64_t>(startPos.n2Idx * ropeShape.headDim) + remainRowCnt * ropeShape.headDim * ropeShape.headNum;
dStride = shape.headDim * shape.headNum;
dRopeStride = ropeShape.headDim * ropeShape.headNum;
} else{
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim * shape.blockSize) + remainRowCnt * shape.headDim + startPos.dIdx;
keyRopeOffset += static_cast<uint64_t>(startPos.n2Idx * ropeShape.headDim * ropeShape.blockSize) + remainRowCnt * ropeShape.headDim;
}
uint32_t dValue = shape.actHeadDim;
uint32_t srcDValue = dStride;
uint32_t dRopeValue = ropeShape.actHeadDim;
uint32_t srcRopeDValue = dRopeStride;
LocalTensor<L1Type> tmpNopeDstTensor = nopeTensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpNopeSrcTensor = nopeGmTensor[offset];
DataCopyGmNDToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, shape.copyRowNumAlign, dValue, srcDValue);
LocalTensor<L1Type> tmpRopeDstTensor = ropeTensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpRopeSrcTensor = ropeGmTensor[keyRopeOffset];
DataCopyGmNDToL1(tmpRopeDstTensor, tmpRopeSrcTensor, copyRowCnt, shape.copyRowNumAlign, dRopeValue, srcRopeDValue);
}
copyFinishRowCnt += copyRowCnt;
curS2Idx += copyRowCnt;
}
}
template<typename L1Type>
__aicore__ inline void GmCopyInToL1HasRopePA(LocalTensor<L1Type>& nopeTensor, LocalTensor<L1Type>& ropeTensor,
GlobalTensor<L1Type>& nopeGmTensor, GlobalTensor<L1Type>& ropeGmTensor,
GlobalTensor<int32_t>& blockTableGm, KVLAYOUT kvLayout,
const PAShape &shape,
const PAShape &ropeShape,
const Position &startPos)
{
uint32_t copyFinishRowCnt = 0;
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch; // 块表的基偏移量
uint32_t curS2Idx = startPos.s2Offset;
uint32_t blockElementCnt = 32U / sizeof(L1Type); // 每个块的元素数量
// ropeshape的M方向与nopeshape保持一样, 此处只判断nopeshape的
while(copyFinishRowCnt < shape.copyRowNum){
uint64_t blockIdOffset = curS2Idx / shape.blockSize; // 获取block table上的索引
uint64_t remainRowCnt = curS2Idx % shape.blockSize; // 获取在单个块上超出的行数
uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset); // 从block table上获取的编号
//计算可以拷贝行数
uint32_t copyRowCnt = shape.blockSize - remainRowCnt; // 一次只能处理一个Block
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum){
copyRowCnt = shape.copyRowNum - copyFinishRowCnt; // 一个block未拷满
}
uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim; // PA的偏移
uint64_t keyRopeOffset = idInBlockTable * ropeShape.blockSize * ropeShape.headNum * ropeShape.headDim;
if (kvLayout == KVLAYOUT::NZ) {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.blockSize * shape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * shape.blockSize;
keyRopeOffset += static_cast<uint64_t>(startPos.n2Idx * ropeShape.blockSize * ropeShape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * ropeShape.blockSize;
LocalTensor<L1Type> tmpNopeDstTensor = nopeTensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpNopeSrcTensor = nopeGmTensor[offset];
DataCopyGmNZToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, (shape.copyRowNumAlign - copyRowCnt), (shape.blockSize - copyRowCnt), shape.actHeadDim);
LocalTensor<L1Type> tmpRopeDstTensor = ropeTensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpRopeSrcTensor = ropeGmTensor[keyRopeOffset];
DataCopyGmNZToL1(tmpRopeDstTensor, tmpRopeSrcTensor, copyRowCnt, (ropeShape.copyRowNumAlign - copyRowCnt), (ropeShape.blockSize - copyRowCnt), ropeShape.actHeadDim);
} else {
uint64_t dStride = shape.headDim;
uint64_t dRopeStride = ropeShape.headDim;
if (kvLayout == KVLAYOUT::BBH) {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim) + remainRowCnt * shape.headDim * shape.headNum + startPos.dIdx;
keyRopeOffset += static_cast<uint64_t>(startPos.n2Idx * ropeShape.headDim) + remainRowCnt * ropeShape.headDim * ropeShape.headNum;
dStride = shape.headDim * shape.headNum;
dRopeStride = ropeShape.headDim * ropeShape.headNum;
} else{
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim * shape.blockSize) + remainRowCnt * shape.headDim + startPos.dIdx;
keyRopeOffset += static_cast<uint64_t>(startPos.n2Idx * ropeShape.headDim * ropeShape.blockSize) + remainRowCnt * ropeShape.headDim;
}
uint32_t dValue = shape.actHeadDim;
uint32_t srcDValue = dStride;
uint32_t dRopeValue = ropeShape.actHeadDim;
uint32_t srcRopeDValue = dRopeStride;
LocalTensor<L1Type> tmpNopeDstTensor = nopeTensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpNopeSrcTensor = nopeGmTensor[offset];
DataCopyGmNDToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, shape.copyRowNumAlign, dValue, srcDValue);
LocalTensor<L1Type> tmpRopeDstTensor = ropeTensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpRopeSrcTensor = ropeGmTensor[keyRopeOffset];
DataCopyGmNDToL1(tmpRopeDstTensor, tmpRopeSrcTensor, copyRowCnt, shape.copyRowNumAlign, dRopeValue, srcRopeDValue);
}
copyFinishRowCnt += copyRowCnt;
curS2Idx += copyRowCnt;
}
}
template<typename L1Type>
__aicore__ inline void GmCopyInToL1PA(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
GlobalTensor<int32_t>& blockTableGm, KVLAYOUT kvLayout,
const PAShape &shape, const Position &startPos)
{
uint32_t copyFinishRowCnt = 0;
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch; // 块表的基偏移量
uint32_t curS2Idx = startPos.s2Offset;
uint32_t blockElementCnt = 32U / sizeof(L1Type); // 每个块的元素数量
while(copyFinishRowCnt < shape.copyRowNum){
uint64_t blockIdOffset = curS2Idx / shape.blockSize; // 获取block table上的索引
uint64_t remainRowCnt = curS2Idx % shape.blockSize; // 获取在单个块上超出的行数
uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset); // 从block table上获取的编号
//计算可以拷贝行数
uint32_t copyRowCnt = shape.blockSize - remainRowCnt; // 一次只能处理一个Block
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum){
copyRowCnt = shape.copyRowNum - copyFinishRowCnt; // 一个block未拷满
}
uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim; // PA的偏移
if (kvLayout == KVLAYOUT::NZ) {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.blockSize * shape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * shape.blockSize;
LocalTensor<L1Type> tmpNopeDstTensor = l1Tensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpNopeSrcTensor = gmTensor[offset];
DataCopyGmNZToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, (shape.copyRowNumAlign - copyRowCnt), (shape.blockSize - copyRowCnt), shape.actHeadDim);
} else {
uint64_t dStride = shape.headDim;
if (kvLayout == KVLAYOUT::BBH) {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim) + remainRowCnt * shape.headDim * shape.headNum + startPos.dIdx;
dStride = shape.headDim * shape.headNum;
} else {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim * shape.blockSize) + remainRowCnt * shape.headDim + startPos.dIdx;
}
uint32_t dValue = shape.actHeadDim;
uint32_t srcDValue = dStride;
LocalTensor<L1Type> tmpNopeDstTensor = l1Tensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpNopeSrcTensor = gmTensor[offset];
DataCopyGmNDToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, shape.copyRowNumAlign, dValue, srcDValue);
}
copyFinishRowCnt += copyRowCnt;
curS2Idx += copyRowCnt;
}
}
template<typename L1Type>
__aicore__ inline void GmScaleCopyInToL1PAForND(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
GlobalTensor<int32_t>& blockTableGm, KVLAYOUT kvLayout,
const PAShape &shape, const Position &startPos)
{
uint32_t copyFinishRowCnt = 0;
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch;
uint32_t curS2Idx = startPos.s2Offset;
constexpr uint32_t blockElementCnt = 32U / sizeof(L1Type);
while(copyFinishRowCnt < shape.copyRowNum) {
uint64_t blockIdOffset = curS2Idx / shape.blockSize;
uint64_t remainRowCnt = curS2Idx % shape.blockSize;
uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset);
uint32_t copyRowCnt = shape.blockSize - remainRowCnt;
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum) {
copyRowCnt = shape.copyRowNum - copyFinishRowCnt;
}
uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim;
if (kvLayout == KVLAYOUT::NZ) {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.blockSize * shape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * shape.blockSize;
LocalTensor<L1Type> tmpNopeDstTensor = l1Tensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpNopeSrcTensor = gmTensor[offset];
DataCopyGmNZToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, (shape.copyRowNumAlign - copyRowCnt), (shape.blockSize - copyRowCnt), shape.actHeadDim);
} else {
uint64_t dStride = shape.headDim;
if (kvLayout == KVLAYOUT::BBH) {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim) + remainRowCnt * shape.headDim * shape.headNum + startPos.dIdx;
dStride = shape.headDim * shape.headNum;
} else {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim * shape.blockSize) + remainRowCnt * shape.headDim + startPos.dIdx;
}
uint32_t dValue = shape.actHeadDim;
uint32_t srcDValue = dStride;
LocalTensor<L1Type> tmpNopeDstTensor = l1Tensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpNopeSrcTensor = gmTensor[offset * 2];
DataCopyGmScaleNDToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, copyRowCnt, dValue, srcDValue);
}
copyFinishRowCnt += copyRowCnt;
curS2Idx += copyRowCnt;
}
}
template<typename L1Type>
__aicore__ inline void GmScaleCopyInToL1PAForDN(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
GlobalTensor<int32_t>& blockTableGm, KVLAYOUT kvLayout,
const PAShape &shape, const Position &startPos)
{
uint32_t copyFinishRowCnt = 0;
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch;
uint32_t curS2Idx = startPos.s2Offset;
constexpr uint32_t blockElementCnt = 32U / sizeof(L1Type);
while(copyFinishRowCnt < shape.copyRowNum) {
uint64_t blockIdOffset = curS2Idx / shape.blockSize;
uint64_t remainRowCnt = curS2Idx % shape.blockSize;
uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset);
uint32_t copyRowCnt = shape.blockSize - remainRowCnt;
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum) {
copyRowCnt = shape.copyRowNum - copyFinishRowCnt;
}
uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim;
if (kvLayout == KVLAYOUT::NZ) {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.blockSize * shape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * shape.blockSize;
LocalTensor<L1Type> tmpNopeDstTensor = l1Tensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpNopeSrcTensor = gmTensor[offset];
DataCopyGmNZToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, (shape.copyRowNumAlign - copyRowCnt), (shape.blockSize - copyRowCnt), shape.actHeadDim);
} else {
uint64_t dStride = shape.headDim;
if (kvLayout == KVLAYOUT::BBH) {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim) + remainRowCnt * shape.headDim * shape.headNum + startPos.dIdx;
dStride = shape.headDim * shape.headNum;
} else {
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim * shape.blockSize) + remainRowCnt * shape.headDim + startPos.dIdx;
}
uint32_t dValue = shape.actHeadDim;
uint32_t srcDValue = dStride;
LocalTensor<L1Type> tmpNopeDstTensor = l1Tensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<L1Type> tmpNopeSrcTensor = gmTensor[offset];
DataCopyGmScaleDNToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, copyRowCnt, dValue, srcDValue);
}
copyFinishRowCnt += copyRowCnt;
curS2Idx += copyRowCnt;
}
}
template<typename INPUT_T>
__aicore__ inline void CopyToL1Nd2Nz(const LocalTensor<INPUT_T> &l1Tensor, const GlobalTensor<INPUT_T> &gmTensor,
uint32_t nValue, uint32_t dValue, uint32_t srcDValue)
{
Nd2NzParams gm2L1Nd2NzParams;
gm2L1Nd2NzParams.ndNum = 1; // ND矩阵的个数
gm2L1Nd2NzParams.nValue = nValue; // 单个ND矩阵的实际行数,单位为元素个数
gm2L1Nd2NzParams.dValue = dValue; // 单个ND矩阵的实际列数,单位为元素个数
gm2L1Nd2NzParams.srcNdMatrixStride = 0; // 相邻ND矩阵起始地址之间的偏移, 单位为元素个数
gm2L1Nd2NzParams.srcDValue = srcDValue; // 同一个ND矩阵中相邻行起始地址之间的偏移, 单位为元素个数
#if (__CCE_AICORE__ == 310) || (defined __DAV_310R6__) || (__NPU_ARCH__ == 5102)
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value || IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
IsSameType<INPUT_T, hifloat8_t>::value || IsSameType<INPUT_T, int8_t>::value) {
gm2L1Nd2NzParams.dstNzC0Stride = (nValue + 31) >> 5 << 5;
} else {
gm2L1Nd2NzParams.dstNzC0Stride = (nValue + 15) >> 4 << 4;
}
#else
gm2L1Nd2NzParams.dstNzC0Stride = (nValue + 15) >> 4 << 4; // NZ矩阵相邻Block起始地址之间的偏移, 单位为Block个数
#endif
gm2L1Nd2NzParams.dstNzNStride = 1; // 转换为NZ矩阵后,ND之间相邻两行在NZ矩阵中起始地址之间的偏移, 单位为Block个数
gm2L1Nd2NzParams.dstNzMatrixStride = 0; // 两个NZ矩阵,起始地址之间的偏移, 单位为元素数量
DataCopy(l1Tensor, gmTensor, gm2L1Nd2NzParams);
}
template<typename INPUT_T>
__aicore__ inline void CopyScaleToL1Nd2Nz(const LocalTensor<INPUT_T> &l1Tensor, const GlobalTensor<INPUT_T> &gmTensor,
uint32_t nValue, uint32_t dValue, uint32_t srcDValue)
{
Nd2NzParams gm2L1Nd2NzParams;
gm2L1Nd2NzParams.ndNum = 1; // ND矩阵的个数
gm2L1Nd2NzParams.nValue = nValue / 2; // 单个ND矩阵的实际行数,单位为元素个数
gm2L1Nd2NzParams.dValue = dValue; // 单个ND矩阵的实际列数,单位为元素个数
gm2L1Nd2NzParams.srcNdMatrixStride = 0; // 相邻ND矩阵起始地址之间的偏移, 单位为元素个数
gm2L1Nd2NzParams.srcDValue = srcDValue; // 同一个ND矩阵中相邻行起始地址之间的偏移, 单位为元素个数
gm2L1Nd2NzParams.dstNzC0Stride = nValue / 2; // NZ矩阵相邻Block起始地址之间的偏移, 单位为Block个数
gm2L1Nd2NzParams.dstNzNStride = 1; // 转换为NZ矩阵后,ND之间相邻两行在NZ矩阵中起始地址之间的偏移, 单位为Block个数
gm2L1Nd2NzParams.dstNzMatrixStride = gm2L1Nd2NzParams.nValue; // 两个NZ矩阵,起始地址之间的偏移, 单位为元素数量
LocalTensor<bfloat16_t> l1TensorCast = l1Tensor.template ReinterpretCast<bfloat16_t>();
GlobalTensor<bfloat16_t> gmTensorCast;
gmTensorCast.SetGlobalBuffer(((__gm__ bfloat16_t*)(gmTensor.GetPhyAddr())));
DataCopy(l1TensorCast, gmTensorCast, gm2L1Nd2NzParams);
}
template<typename INPUT_T>
__aicore__ inline void CopyScaleToL1Dn2Nz(const LocalTensor<INPUT_T> &l1Tensor, const GlobalTensor<INPUT_T> &gmTensor,
uint32_t nValue, uint32_t dValue, uint32_t srcDValue)
{
Dn2NzParams gm2L1Dn2NzParams;
gm2L1Dn2NzParams.dnNum = 1; // ND矩阵的个数
gm2L1Dn2NzParams.nValue = nValue / 2; // 单个DN矩阵的实际列数,单位为元素个数
gm2L1Dn2NzParams.dValue = dValue; // 单个DN矩阵的实际行数,单位为元素个数
gm2L1Dn2NzParams.srcDnMatrixStride = 0; // 相邻Dn矩阵起始地址之间的偏移, 单位为元素个数
gm2L1Dn2NzParams.srcDValue = srcDValue / 2; // 同一个Dn矩阵中相邻行起始地址之间的偏移, 单位为元素个数
gm2L1Dn2NzParams.dstNzC0Stride = nValue / 2;
gm2L1Dn2NzParams.dstNzNStride = 1; // 转换为NZ矩阵后,ND之间相邻两行在NZ矩阵中起始地址之间的偏移, 单位为Block个数
gm2L1Dn2NzParams.dstNzMatrixStride = gm2L1Dn2NzParams.nValue; // 两个NZ矩阵,起始地址之间的偏移, 单位为元素数量
LocalTensor<bfloat16_t> l1TensorCast = l1Tensor.template ReinterpretCast<bfloat16_t>();
GlobalTensor<bfloat16_t> gmTensorCast;
gmTensorCast.SetGlobalBuffer(((__gm__ bfloat16_t*)(gmTensor.GetPhyAddr())));
DataCopy(l1TensorCast, gmTensorCast, gm2L1Dn2NzParams);
}
template<typename INPUT_T>
__aicore__ inline void CopyToL1Nd2NzGS1Merge(const LocalTensor<INPUT_T> &l1Tensor, const GlobalTensor<INPUT_T> &gmTensor,
uint32_t ndNum, uint32_t nValue, uint32_t dValue, uint32_t srcNdMatrixStride, uint32_t srcDValue, uint32_t dstNzC0Stride) // BSNGD 合轴拷贝
{
Nd2NzParams gm2L1Nd2NzParams;
gm2L1Nd2NzParams.ndNum = ndNum; // ND矩阵的个数
gm2L1Nd2NzParams.nValue = nValue; // 单个ND矩阵的实际行数,单位为元素个数
gm2L1Nd2NzParams.dValue = dValue; // 单个ND矩阵的实际列数,单位为元素个数
gm2L1Nd2NzParams.srcNdMatrixStride = srcNdMatrixStride; // 相邻ND矩阵起始地址之间的偏移, 单位为元素个数
gm2L1Nd2NzParams.srcDValue = srcDValue; // 同一个ND矩阵中相邻行起始地址之间的偏移, 单位为元素个数
#if (__CCE_AICORE__ == 310) || (defined __DAV_310R6__) || (__NPU_ARCH__ == 5102)
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value || IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
IsSameType<INPUT_T, hifloat8_t>::value || IsSameType<INPUT_T, int8_t>::value) {
gm2L1Nd2NzParams.dstNzC0Stride = (dstNzC0Stride + 31) >> 5 << 5; // NZ矩阵相邻Block起始地址之间的偏移,单位为Block个数,32对齐
} else {
gm2L1Nd2NzParams.dstNzC0Stride = (dstNzC0Stride + 15) >> 4 << 4; // NZ矩阵相邻Block起始地址之间的偏移,单位为Block个数,16对齐
}
#else
gm2L1Nd2NzParams.dstNzC0Stride = (dstNzC0Stride + 15) >> 4 << 4; // NZ矩阵相邻Block起始地址之间的偏移,单位为Block个数,16对齐
#endif
gm2L1Nd2NzParams.dstNzNStride = 1; // 转换为NZ矩阵后,ND之间相邻两行在NZ矩阵中起始地址之间的偏移, 单位为Block个数
gm2L1Nd2NzParams.dstNzMatrixStride = nValue * 32 / sizeof(INPUT_T); // 两个NZ矩阵,起始地址之间的偏移, 单位为元素数量
DataCopy(l1Tensor, gmTensor, gm2L1Nd2NzParams);
}
#endif

View File

@@ -0,0 +1,56 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file FixpipeOut.h
* \brief
*/
#ifndef FIXPIPEOUT_H
#define FIXPIPEOUT_H
constexpr FixpipeConfig PFA_CFG_ROW_MAJOR_UB = {CO2Layout::ROW_MAJOR, true}; // ROW_MAJOR: 使能NZ2ND,输出数据格式为ND格式; true: 用于用户指定目的地址的位置是否是UB
constexpr FixpipeConfig PFA_CFG_ROW_MAJOR_GM = {CO2Layout::ROW_MAJOR, false}; // ROW_MAJOR: 使能NZ2ND,输出数据格式为ND格式; true: 用于用户指定目的地址的位置是否是UB
constexpr FixpipeConfig FA_CFG_NZ_UB = {CO2Layout::NZ, true}; // 不使能NZ2ND,输出数据格式为NZ格式; true: 用于用户指定目的地址的位置是否是UB
struct fixpipeOutParams {
uint32_t fixpOutMSize;
uint32_t fixpOutNSize;
};
template<typename mmOutputType, typename computeType, typename l0cType>
__aicore__ inline void FixpipeMmCopyOutToUB(LocalTensor<mmOutputType>& mmResUb, LocalTensor<l0cType>& L0CTensor, const fixpipeOutParams& fixpOutParam)
{
FixpipeParamsC310<CO2Layout::ROW_MAJOR> L0C2UbFixpParams; // L0C->UB
L0C2UbFixpParams.nSize = (fixpOutParam.fixpOutNSize + 7) >> 3 << 3; // L0C上的bmm1结果矩阵N方向的size大小;同mmadParams.n;8个元素(32B)对齐
L0C2UbFixpParams.mSize = (fixpOutParam.fixpOutMSize + 1) >> 1 << 1; // 有效数据不足16行,只需输出部分行即可;L0C上的bmm1结果矩阵M方向的size大小必须是偶数
L0C2UbFixpParams.srcStride = ((L0C2UbFixpParams.mSize + 15) >> 4) << 4; // L0C上matmul结果相邻连续数据片断间隔(前面一个数据块的头与后面数据块的头的间隔),单位为16 *sizeof(T) //源NZ矩阵中相邻Z排布的起始地址偏移
L0C2UbFixpParams.dstStride = (L0C2UbFixpParams.nSize + 15) >> 4 << 4; // mmResUb上两行之间的间隔,单位:element。 // 128:根据比对dump文件得到,ND方案(S1 * S2)时脏数据用mask剔除
L0C2UbFixpParams.dualDstCtl = 1; // 双目标模式,按M维度拆分, M / 2 * N写入每个UB,M必须为2的倍数
L0C2UbFixpParams.params.ndNum = 1;
L0C2UbFixpParams.params.srcNdStride = 0;
L0C2UbFixpParams.params.dstNdStride = 0;
Fixpipe<mmOutputType, computeType, PFA_CFG_ROW_MAJOR_UB>(mmResUb, L0CTensor, L0C2UbFixpParams); // 将matmul结果从L0C搬运到UB
}
template<typename mmOutputType, typename computeType, typename l0cType>
__aicore__ inline void FixpipeMmCopyOutToGm(GlobalTensor<mmOutputType>& mmResGm,LocalTensor<l0cType>& L0CTensor, const fixpipeOutParams& fixpOutParam)
{
FixpipeParamsC310<CO2Layout::ROW_MAJOR> L0C2GmFixpParams; // L0C->Gm
L0C2GmFixpParams.nSize = (fixpOutParam.fixpOutNSize + 7) >> 3 << 3; // L0C上的bmm1结果矩阵N方向的size大小;同mmadParams.n;8个元素(32B)对齐;分档计算且vector1中通过mask筛选出实际有效值
L0C2GmFixpParams.mSize = (fixpOutParam.fixpOutMSize + 1) >> 1 << 1; // 有效数据不足16行,只需输出部分行即可;L0C上的bmm1结果矩阵M方向的size大小;同mmadParams.m
L0C2GmFixpParams.srcStride = ((L0C2GmFixpParams.mSize + 15) >> 4) << 4; // L0C上bmm1结果相邻连续数据片断间隔(前面一个数据块的头与后面数据块的头的间隔)
L0C2GmFixpParams.dstStride = (L0C2GmFixpParams.nSize + 15) >> 4 << 4; // mmResGm上两行之间的间隔
L0C2GmFixpParams.dualDstCtl = 1;
L0C2GmFixpParams.params.ndNum = 1;
L0C2GmFixpParams.params.srcNdStride = 0;
L0C2GmFixpParams.params.dstNdStride = 0;
Fixpipe<mmOutputType, computeType, PFA_CFG_ROW_MAJOR_GM>(mmResGm, L0CTensor, L0C2GmFixpParams); // 将matmul结果从L0C搬运到Gm
}
#endif

View File

@@ -0,0 +1,139 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file vf_basic_block_aligned128_no_update_sfa.h
* \brief
*/
#ifndef VF_BASIC_BLOCK_ALIGNED128_NO_UPDATE_SFA_H
#define VF_BASIC_BLOCK_ALIGNED128_NO_UPDATE_SFA_H
#include "vf_basic_block_utils.h"
using namespace regbaseutil;
namespace FaVectorApi {
// no update, originN == 128
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
__simd_vf__ void ProcessVec1NoUpdateImpl128VF(
__ubuf__ T2 * expUb, __ubuf__ T * expSumUb, __ubuf__ T * maxUb, __ubuf__ T * maxUbStart,
__ubuf__ T * srcUb, const uint32_t blockStride, const uint32_t repeatStride,
const uint16_t m, const T scale, const T minValue)
{
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll;
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
AscendC::MicroAPI::RegTensor<float> vreg_input_max;
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
AscendC::MicroAPI::RegTensor<float> vreg_exp_even;
AscendC::MicroAPI::RegTensor<float> vreg_exp_odd;
// bfloat16_t
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_even_bf16;
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_odd_bf16;
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
// half
AscendC::MicroAPI::RegTensor<half> vreg_exp_even_fp16;
AscendC::MicroAPI::RegTensor<half> vreg_exp_odd_fp16;
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
AscendC::MicroAPI::UnalignRegForStore ureg_max;
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<T, AscendC::MicroAPI::MaskPattern::ALL>();
AscendC::MicroAPI::MaskReg preg_all_b16 =
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
for (uint16_t i = 0; i < m; ++i) {
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
AscendC::MicroAPI::LoadAlign(vreg_input_x_unroll, srcUb + floatRepSize + i * s2BaseSize);
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_all); // Muls(scale)
AscendC::MicroAPI::Muls(vreg_input_x_unroll, vreg_input_x_unroll, scale, preg_all);
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_all);
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)srcUb + floatRepSize + i * s2BaseSize, vreg_input_x_unroll, preg_all);
AscendC::MicroAPI::Max(vreg_max_tmp, vreg_input_x, vreg_input_x_unroll, preg_all);
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
vreg_input_max, vreg_max_tmp, preg_all);
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)maxUb), vreg_input_max, ureg_max, 1);
}
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)maxUb), ureg_max, 0);
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
for (uint16_t i = 0; i < m; ++i) {
// maxUb is [S1, 1], BRC_B32 is reading one fp32 element and broadcast it to all 64 vreg element
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(
vreg_max_brc, maxUbStart + i);
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_DINTLV_B32>(
vreg_input_x, vreg_input_x_unroll, srcUb + i * s2BaseSize);
AscendC::MicroAPI::ExpSub(vreg_exp_even, vreg_input_x, vreg_max_brc, preg_all);
AscendC::MicroAPI::ExpSub(vreg_exp_odd, vreg_input_x_unroll, vreg_max_brc, preg_all);
// x_sum = sum(x_exp, axis=-1, keepdims=True)
AscendC::MicroAPI::Add(vreg_exp_sum, vreg_exp_even, vreg_exp_odd, preg_all);
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
vreg_exp_sum, vreg_exp_sum, preg_all);
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)expSumUb), vreg_exp_sum, ureg_exp_sum, 1);
if constexpr (IsSameType<T2, bfloat16_t>::value) {
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_bf16, vreg_exp_even, preg_all);
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_bf16, vreg_exp_odd, preg_all);
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_bf16, (RegTensor<uint16_t>&)vreg_exp_even_bf16,
(RegTensor<uint16_t>&)vreg_exp_odd_bf16, preg_all_b16);
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T2 *&)expUb), vreg_exp_bf16, blockStride, repeatStride, preg_all_b16);
} else if constexpr (IsSameType<T2, half>::value) {
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_fp16, vreg_exp_even, preg_all);
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_fp16, vreg_exp_odd, preg_all);
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_fp16, (RegTensor<uint16_t>&)vreg_exp_even_fp16,
(RegTensor<uint16_t>&)vreg_exp_odd_fp16, preg_all_b16);
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T2 *&)expUb), vreg_exp_fp16, blockStride, repeatStride, preg_all_b16);
}
}
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)expSumUb), ureg_exp_sum, 0);
}
// no update, originN == 128
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
__aicore__ inline void ProcessVec1NoUpdateImpl128(
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
{
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
const uint32_t repeatStride = 1;
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
__ubuf__ T * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
__ubuf__ T * maxUbStart = (__ubuf__ T*)maxTensor.GetPhyAddr();
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
ProcessVec1NoUpdateImpl128VF<T, T2, s1BaseSize, s2BaseSize>(
expUb, expSumUb, maxUb, maxUbStart, srcUb, blockStride, repeatStride, m, scale, minValue);
}
} // namespace
#endif // VF_BASIC_BLOCK_ALIGNED128_NO_UPDATE_SFA_H

View File

@@ -0,0 +1,143 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file vf_basic_block_aligned128_update_sfa.h
* \brief
*/
#ifndef VF_BASIC_BLOCK_ALIGNED128_UPDATE_SFA_H
#define VF_BASIC_BLOCK_ALIGNED128_UPDATE_SFA_H
#include "vf_basic_block_utils.h"
using namespace regbaseutil;
namespace FaVectorApi {
// update, originN == 128
template <typename T, typename T2, uint32_t s1BaseSize = 128, uint32_t s2BaseSize = 128>
__simd_vf__ void ProcessVec1UpdateImpl128VF(
__ubuf__ T2 * expUb, __ubuf__ T * srcUb, __ubuf__ T * inMaxUb,
__ubuf__ T * tmpExpSumUb, __ubuf__ T * tmpMaxUb, __ubuf__ T * tmpMaxUb2, const uint32_t blockStride,
const uint32_t repeatStride, const uint16_t m, const T scale, const T minValue)
{
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll;
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
AscendC::MicroAPI::RegTensor<float> vreg_in_max;
AscendC::MicroAPI::RegTensor<float> vreg_max_new;
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
AscendC::MicroAPI::RegTensor<float> vreg_cur_max;
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
AscendC::MicroAPI::RegTensor<float> vreg_in_exp_sum;
AscendC::MicroAPI::RegTensor<float> vreg_exp_even;
AscendC::MicroAPI::RegTensor<float> vreg_exp_odd;
// bfloat16_t
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_even_bf16;
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_odd_bf16;
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
// half
AscendC::MicroAPI::RegTensor<half> vreg_exp_even_fp16;
AscendC::MicroAPI::RegTensor<half> vreg_exp_odd_fp16;
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
AscendC::MicroAPI::UnalignRegForStore ureg_max;
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
AscendC::MicroAPI::MaskReg preg_all_b16 =
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
// x_max = max(src, axis=-1, keepdims=True); x_max = Max(x_max, inMax)
for (uint16_t i = 0; i < m; ++i) {
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
AscendC::MicroAPI::LoadAlign(vreg_input_x_unroll, srcUb + floatRepSize + i * s2BaseSize);
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_all); // Muls(scale)
AscendC::MicroAPI::Muls(vreg_input_x_unroll, vreg_input_x_unroll, scale, preg_all);
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_all);
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)srcUb + floatRepSize + i * s2BaseSize, vreg_input_x_unroll, preg_all);
AscendC::MicroAPI::Max(vreg_max_tmp, vreg_input_x, vreg_input_x_unroll, preg_all);
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
vreg_max_tmp, vreg_max_tmp, preg_all);
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)tmpMaxUb), vreg_max_tmp, ureg_max, 1);
}
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)tmpMaxUb), ureg_max, 0);
AscendC::MicroAPI::LoadAlign(vreg_in_max, inMaxUb);
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
AscendC::MicroAPI::LoadAlign(vreg_cur_max, tmpMaxUb2); // 获取新的max[s1, 1]
AscendC::MicroAPI::Max(vreg_max_new, vreg_cur_max, vreg_in_max, preg_all); // 计算新、旧max的最大值
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)tmpMaxUb2, vreg_max_new, preg_all);
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
for (uint16_t i = 0; i < m; ++i) {
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_max_brc, tmpMaxUb2 + i);
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_DINTLV_B32>(
vreg_input_x, vreg_input_x_unroll, srcUb + i * s2BaseSize);
AscendC::MicroAPI::ExpSub(vreg_exp_even, vreg_input_x, vreg_max_brc, preg_all);
AscendC::MicroAPI::ExpSub(vreg_exp_odd, vreg_input_x_unroll, vreg_max_brc, preg_all);
// x_sum = sum(x_exp, axis=-1, keepdims=True)
AscendC::MicroAPI::Add(vreg_exp_sum, vreg_exp_even, vreg_exp_odd, preg_all);
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
vreg_exp_sum, vreg_exp_sum, preg_all);
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)tmpExpSumUb), vreg_exp_sum, ureg_exp_sum, 1);
if constexpr (IsSameType<T2, bfloat16_t>::value) {
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_bf16, vreg_exp_even, preg_all);
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_bf16, vreg_exp_odd, preg_all);
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_bf16, (RegTensor<uint16_t>&)vreg_exp_even_bf16,
(RegTensor<uint16_t>&)vreg_exp_odd_bf16, preg_all_b16);
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T2 *&)expUb), vreg_exp_bf16, blockStride, repeatStride, preg_all_b16);
} else if constexpr (IsSameType<T2, half>::value) {
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_fp16, vreg_exp_even, preg_all);
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_fp16, vreg_exp_odd, preg_all);
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_fp16, (RegTensor<uint16_t>&)vreg_exp_even_fp16,
(RegTensor<uint16_t>&)vreg_exp_odd_fp16, preg_all_b16);
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T2 *&)expUb), vreg_exp_fp16, blockStride, repeatStride, preg_all_b16);
}
}
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)tmpExpSumUb), ureg_exp_sum, 0);
}
// update, originN == 128
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
__aicore__ inline void ProcessVec1UpdateImpl128(
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor, const LocalTensor<T>& inMaxTensor,
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
{
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
const uint32_t repeatStride = 1;
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
__ubuf__ T * inMaxUb = (__ubuf__ T*)inMaxTensor.GetPhyAddr();
__ubuf__ T * tmpExpSumUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr();
__ubuf__ T * tmpMaxUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
__ubuf__ T * tmpMaxUb2 = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
ProcessVec1UpdateImpl128VF <T, T2, s1BaseSize, s2BaseSize>(
expUb, srcUb, inMaxUb, tmpExpSumUb, tmpMaxUb, tmpMaxUb2, blockStride, repeatStride, m, scale, minValue);
}
} // namespace
#endif // VF_BASIC_BLOCK_ALIGNED128_UPDATE_SFA_H

View File

@@ -0,0 +1,149 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file vf_basic_block_unaligned128_no_update_sfa.h
* \brief
*/
#ifndef VF_BASIC_BLOCK_UNALIGNED128_NO_UPDATE_SFA_H
#define VF_BASIC_BLOCK_UNALIGNED128_NO_UPDATE_SFA_H
#include "vf_basic_block_utils.h"
using namespace regbaseutil;
namespace FaVectorApi {
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
__simd_vf__ void ProcessVec1NoUpdateGeneralImpl128VF(
__ubuf__ T2 * expUb, __ubuf__ T * expSumUb, __ubuf__ T * maxUb, __ubuf__ T * maxUbStart,
__ubuf__ T * srcUb, const uint32_t blockStride, const uint32_t repeatStride,
const uint16_t m, const T scale, const T minValue, uint32_t pltOriTailN, uint32_t pltTailN)
{
AscendC::MicroAPI::RegTensor<float> vreg_min;
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll;
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll_new;
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
AscendC::MicroAPI::RegTensor<float> vreg_input_max;
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
AscendC::MicroAPI::RegTensor<float> vreg_exp_even;
AscendC::MicroAPI::RegTensor<float> vreg_exp_odd;
// bfloat16_t
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_even_bf16;
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_odd_bf16;
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
// half
AscendC::MicroAPI::RegTensor<half> vreg_exp_even_fp16;
AscendC::MicroAPI::RegTensor<half> vreg_exp_odd_fp16;
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
AscendC::MicroAPI::UnalignRegForStore ureg_max;
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
AscendC::MicroAPI::MaskReg preg_all_b16 =
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
AscendC::MicroAPI::MaskReg preg_all_b8 = AscendC::MicroAPI::CreateMask<T2, AscendC::MicroAPI::MaskPattern::ALL>();
AscendC::MicroAPI::MaskReg preg_tail_n = AscendC::MicroAPI::UpdateMask<float>(pltTailN);
AscendC::MicroAPI::MaskReg preg_ori_tail_n = AscendC::MicroAPI::UpdateMask<float>(pltOriTailN);
AscendC::MicroAPI::MaskReg preg_reduce_n =
AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::VL8>();
AscendC::MicroAPI::Duplicate(vreg_min, minValue);
for (uint16_t i = 0; i < m; ++i) {
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
AscendC::MicroAPI::LoadAlign(vreg_input_x_unroll, srcUb + floatRepSize + i * s2BaseSize);
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_all); // Muls(scale)
AscendC::MicroAPI::Muls(vreg_input_x_unroll, vreg_input_x_unroll, scale, preg_ori_tail_n);
AscendC::MicroAPI::Select(vreg_input_x_unroll_new, vreg_input_x_unroll, vreg_min, preg_ori_tail_n);
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_all);
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)srcUb + floatRepSize + i * s2BaseSize, vreg_input_x_unroll_new, preg_tail_n);
AscendC::MicroAPI::Max(vreg_max_tmp, vreg_input_x, vreg_input_x_unroll_new, preg_all);
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
vreg_input_max, vreg_max_tmp, preg_all);
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)maxUb), vreg_input_max, ureg_max, 1);
}
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)maxUb), ureg_max, 0);
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
for (uint16_t i = 0; i < m; ++i) {
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_max_brc, maxUbStart + i);
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_DINTLV_B32>(
vreg_input_x, vreg_input_x_unroll, srcUb + i * s2BaseSize);
AscendC::MicroAPI::ExpSub(vreg_exp_even, vreg_input_x, vreg_max_brc, preg_all);
AscendC::MicroAPI::ExpSub(vreg_exp_odd, vreg_input_x_unroll, vreg_max_brc, preg_all);
// x_sum = sum(x_exp, axis=-1, keepdims=True)
AscendC::MicroAPI::Add(vreg_exp_sum, vreg_exp_even, vreg_exp_odd, preg_all);
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
vreg_exp_sum, vreg_exp_sum, preg_all);
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)expSumUb), vreg_exp_sum, ureg_exp_sum, 1);
if constexpr (IsSameType<T2, bfloat16_t>::value) {
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_bf16, vreg_exp_even, preg_all);
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_bf16, vreg_exp_odd, preg_all);
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_bf16, (RegTensor<uint16_t>&)vreg_exp_even_bf16,
(RegTensor<uint16_t>&)vreg_exp_odd_bf16, preg_all_b16);
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T2 *&)expUb), vreg_exp_bf16, blockStride, repeatStride, preg_all_b16);
} else if constexpr (IsSameType<T2, half>::value) {
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_fp16, vreg_exp_even, preg_all);
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_fp16, vreg_exp_odd, preg_all);
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_fp16, (RegTensor<uint16_t>&)vreg_exp_even_fp16,
(RegTensor<uint16_t>&)vreg_exp_odd_fp16, preg_all_b16);
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T2 *&)expUb), vreg_exp_fp16, blockStride, repeatStride, preg_all_b16);
}
}
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)expSumUb), ureg_exp_sum, 0);
}
// no update, 64 < originN <= 128
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
__aicore__ inline void ProcessVec1NoUpdateGeneralImpl128(
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
{
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
// stride, high 16bits: blockStride (65*16*2/32),单位block, low 16bits: repeatStride (1)
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
const uint32_t repeatStride = 1;
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
__ubuf__ T * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
__ubuf__ T * maxUbStart = (__ubuf__ T*)maxTensor.GetPhyAddr();
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
const uint32_t oriTailN = originN - floatRepSize;
const uint32_t tailN = s2BaseSize - floatRepSize;
uint32_t pltOriTailN = oriTailN;
uint32_t pltTailN = tailN;
ProcessVec1NoUpdateGeneralImpl128VF<T, T2, s1BaseSize, s2BaseSize>(
expUb, expSumUb, maxUb, maxUbStart, srcUb, blockStride, repeatStride, m, scale, minValue,
pltOriTailN, pltTailN);
}
} // namespace
#endif // VF_BASIC_BLOCK_UNALIGNED128_NO_UPDATE_SFA_H

View File

@@ -0,0 +1,159 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file vf_basic_block_unaligned128_update_sfa.h
* \brief
*/
#ifndef VF_BASIC_BLOCK_UNALIGNED128_UPDATE_SFA_H
#define VF_BASIC_BLOCK_UNALIGNED128_UPDATE_SFA_H
#include "vf_basic_block_utils.h"
using namespace regbaseutil;
namespace FaVectorApi {
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
__simd_vf__ void ProcessVec1UpdateGeneralImpl128VF(
__ubuf__ T2 * expUb, __ubuf__ T * srcUb, __ubuf__ T * inMaxUb,
__ubuf__ T * tmpExpSumUb, __ubuf__ T * tmpMaxUb, __ubuf__ T * tmpMaxUb2, const uint32_t blockStride,
const uint32_t repeatStride, const uint16_t m, const T scale, const T minValue, uint32_t pltOriTailN,
uint32_t pltTailN, uint32_t pltN)
{
AscendC::MicroAPI::RegTensor<float> vreg_min;
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll;
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll_new;
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
AscendC::MicroAPI::RegTensor<float> vreg_cur_max;
AscendC::MicroAPI::RegTensor<float> vreg_max_new;
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
AscendC::MicroAPI::RegTensor<float> vreg_in_max;
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
AscendC::MicroAPI::RegTensor<float> vreg_exp_even;
AscendC::MicroAPI::RegTensor<float> vreg_exp_odd;
// bfloat16_t
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_even_bf16;
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_odd_bf16;
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
// half
AscendC::MicroAPI::RegTensor<half> vreg_exp_even_fp16;
AscendC::MicroAPI::RegTensor<half> vreg_exp_odd_fp16;
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
AscendC::MicroAPI::UnalignRegForStore ureg_max;
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
AscendC::MicroAPI::MaskReg preg_all_b16 = AscendC::MicroAPI::CreateMask<uint16_t,
AscendC::MicroAPI::MaskPattern::ALL>();
AscendC::MicroAPI::MaskReg preg_n_b16 = AscendC::MicroAPI::UpdateMask<uint16_t>(pltN);
AscendC::MicroAPI::MaskReg preg_tail_n = AscendC::MicroAPI::UpdateMask<T>(pltTailN);
AscendC::MicroAPI::MaskReg preg_ori_tail_n = AscendC::MicroAPI::UpdateMask<T>(pltOriTailN);
AscendC::MicroAPI::Duplicate(vreg_min, minValue);
// x_max = max(src, axis=-1, keepdims=True); x_max = Max(x_max, inMax)
for (uint16_t i = 0; i < m; ++i) {
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
AscendC::MicroAPI::LoadAlign(vreg_input_x_unroll, srcUb + floatRepSize + i * s2BaseSize);
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_all); // Muls(scale)
AscendC::MicroAPI::Muls(vreg_input_x_unroll, vreg_input_x_unroll, scale, preg_ori_tail_n);
AscendC::MicroAPI::Select(vreg_input_x_unroll_new, vreg_input_x_unroll, vreg_min, preg_ori_tail_n);
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_all);
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)srcUb + floatRepSize + i * s2BaseSize, vreg_input_x_unroll_new, preg_tail_n);
AscendC::MicroAPI::Max(vreg_max_tmp, vreg_input_x, vreg_input_x_unroll_new, preg_all);
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
vreg_cur_max, vreg_max_tmp, preg_all);
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)tmpMaxUb), vreg_cur_max, ureg_max, 1);
}
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)tmpMaxUb), ureg_max, 0);
AscendC::MicroAPI::LoadAlign(vreg_in_max, inMaxUb);
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
AscendC::MicroAPI::LoadAlign(vreg_cur_max, tmpMaxUb2); // 获取新的max[s1, 1]
AscendC::MicroAPI::Max(vreg_max_new, vreg_cur_max, vreg_in_max, preg_all); // 计算新、旧max的最大值
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)tmpMaxUb2, vreg_max_new, preg_all);
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
for (uint16_t i = 0; i < m; ++i) {
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(
vreg_max_brc, tmpMaxUb2 + i);
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_DINTLV_B32>(
vreg_input_x, vreg_input_x_unroll, srcUb + i * s2BaseSize);
AscendC::MicroAPI::ExpSub(vreg_exp_even, vreg_input_x, vreg_max_brc, preg_all);
AscendC::MicroAPI::ExpSub(vreg_exp_odd, vreg_input_x_unroll, vreg_max_brc, preg_all);
// x_sum = sum(x_exp, axis=-1, keepdims=True)
AscendC::MicroAPI::Add(vreg_exp_sum, vreg_exp_even, vreg_exp_odd, preg_all);
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
vreg_exp_sum, vreg_exp_sum, preg_all);
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)tmpExpSumUb), vreg_exp_sum, ureg_exp_sum, 1);
if constexpr (IsSameType<T2, bfloat16_t>::value) {
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_bf16, vreg_exp_even, preg_all);
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_bf16, vreg_exp_odd, preg_all);
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_bf16, (RegTensor<uint16_t>&)vreg_exp_even_bf16,
(RegTensor<uint16_t>&)vreg_exp_odd_bf16, preg_all_b16);
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T2 *&)expUb), vreg_exp_bf16, blockStride, repeatStride, preg_n_b16);
} else if constexpr (IsSameType<T2, half>::value) {
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_fp16, vreg_exp_even, preg_all);
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_fp16, vreg_exp_odd, preg_all);
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_fp16, (RegTensor<uint16_t>&)vreg_exp_even_fp16,
(RegTensor<uint16_t>&)vreg_exp_odd_fp16, preg_all_b16);
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T2 *&)expUb), vreg_exp_fp16, blockStride, repeatStride, preg_n_b16);
}
}
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)tmpExpSumUb), ureg_exp_sum, 0);
}
// update, 64 < originN <= 128
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
__aicore__ inline void ProcessVec1UpdateGeneralImpl128(
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor, const LocalTensor<T>& inMaxTensor,
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
{
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
const uint32_t repeatStride = 1;
const uint32_t oriTailN = originN - floatRepSize;
const uint32_t tailN = s2BaseSize - floatRepSize;
uint32_t pltOriTailN = oriTailN;
uint32_t pltTailN = tailN;
uint32_t pltN = s2BaseSize;
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
__ubuf__ T * inMaxUb = (__ubuf__ T*)inMaxTensor.GetPhyAddr();
__ubuf__ T * tmpExpSumUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr();
__ubuf__ T * tmpMaxUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
__ubuf__ T * tmpMaxUb2 = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
ProcessVec1UpdateGeneralImpl128VF<T, T2, s1BaseSize, s2BaseSize>(
expUb, srcUb, inMaxUb, tmpExpSumUb, tmpMaxUb, tmpMaxUb2, blockStride, repeatStride,
m, scale, minValue, pltOriTailN, pltTailN, pltN);
}
} // namespace
#endif // VF_BASIC_BLOCK_UNALIGNED128_UPDATE_SFA_H

View File

@@ -0,0 +1,129 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file vf_basic_block_unaligned64_no_update_sfa.h
* \brief
*/
#ifndef VF_BASIC_BLOCK_UNALIGNED64_NO_UPDATE_SFA_H
#define VF_BASIC_BLOCK_UNALIGNED64_NO_UPDATE_SFA_H
#include "vf_basic_block_utils.h"
using namespace regbaseutil;
namespace FaVectorApi {
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
__simd_vf__ void ProcessVec1NoUpdateImpl64VF(
__ubuf__ T2 * expUb, __ubuf__ T * expSumUb, __ubuf__ T * maxUb, __ubuf__ T * maxUbStart,
__ubuf__ T * srcUb, const uint32_t blockStride, const uint32_t repeatStride,
const uint16_t m, const T scale, const T minValue, uint32_t pltOriginalN, uint32_t pltSrcN)
{
AscendC::MicroAPI::RegTensor<float> vreg_min;
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
AscendC::MicroAPI::RegTensor<float> vreg_input_max;
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
AscendC::MicroAPI::RegTensor<float> vreg_exp;
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
// bfloat16_t
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_dst_even_bf16;
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_dst_odd_bf16;
// half
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
AscendC::MicroAPI::RegTensor<half> vreg_dst_even_fp16;
AscendC::MicroAPI::RegTensor<half> vreg_dst_odd_fp16;
AscendC::MicroAPI::UnalignRegForStore ureg_max;
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
AscendC::MicroAPI::MaskReg preg_all_b16 =
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
AscendC::MicroAPI::MaskReg preg_src_n = AscendC::MicroAPI::UpdateMask<float>(pltSrcN);
AscendC::MicroAPI::MaskReg preg_src_n_b16 =
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::H>();
AscendC::MicroAPI::MaskReg preg_ori_src_n = AscendC::MicroAPI::UpdateMask<T>(pltOriginalN);
// x_max = max(src, axis=-1, keepdims=True)
for (uint16_t i = 0; i < m; ++i) {
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_ori_src_n); // Muls(scale)
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_src_n);
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
vreg_input_max, vreg_input_x, preg_ori_src_n);
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)maxUb), vreg_input_max, ureg_max, 1);
}
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)maxUb), ureg_max, 0);
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
for (uint16_t i = 0; i < m; ++i) {
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(
vreg_max_brc, maxUbStart + i);
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
AscendC::MicroAPI::ExpSub(vreg_exp, vreg_input_x, vreg_max_brc, preg_ori_src_n);
// x_sum = sum(x_exp, axis=-1, keepdims=True)
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
vreg_exp_sum, vreg_exp, preg_ori_src_n);
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)expSumUb), vreg_exp_sum, ureg_exp_sum, 1);
if constexpr (IsSameType<T2, bfloat16_t>::value) {
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_bf16, vreg_exp, preg_all_b16);
AscendC::MicroAPI::DeInterleave(vreg_dst_even_bf16, vreg_dst_odd_bf16,
vreg_exp_bf16, vreg_exp_bf16);
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T2 *&)expUb), vreg_dst_even_bf16, blockStride, repeatStride, preg_src_n_b16);
} else if constexpr (IsSameType<T2, half>::value) {
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_fp16, vreg_exp, preg_all_b16);
AscendC::MicroAPI::DeInterleave(vreg_dst_even_fp16, vreg_dst_odd_fp16,
vreg_exp_fp16, vreg_exp_fp16);
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T2 *&)expUb), vreg_dst_even_fp16, blockStride, repeatStride, preg_src_n_b16);
}
}
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)expSumUb), ureg_exp_sum, 0);
}
// no update, originN <= 64
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
__aicore__ inline void ProcessVec1NoUpdateImpl64(
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
{
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
__ubuf__ T * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
__ubuf__ T * maxUbStart = (__ubuf__ T*)maxTensor.GetPhyAddr();
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
const uint32_t repeatStride = 1;
uint32_t pltOriginalN = originN;
uint32_t pltSrcN = s2BaseSize;
ProcessVec1NoUpdateImpl64VF<T, T2, s1BaseSize, s2BaseSize>(
expUb, expSumUb, maxUb, maxUbStart, srcUb, blockStride, repeatStride, m, scale, minValue,
pltOriginalN, pltSrcN);
}
} // namespace
#endif // VF_BASIC_BLOCK_UNALIGNED64_NO_UPDATE_SFA_H

View File

@@ -0,0 +1,141 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file vf_basic_block_aligned64_update_sfa.h
* \brief
*/
#ifndef VF_BASIC_BLOCK_ALIGNED64_UPDATE_SFA_H
#define VF_BASIC_BLOCK_ALIGNED64_UPDATE_SFA_H
#include "vf_basic_block_utils.h"
using namespace regbaseutil;
namespace FaVectorApi {
// update, originN <= 64
template <typename T, typename T2, uint32_t s1BaseSize = 128, uint32_t s2BaseSize = 128>
__simd_vf__ void ProcessVec1UpdateImpl64VF(
__ubuf__ T2 * expUb, __ubuf__ T * srcUb, __ubuf__ T * inMaxUb,
__ubuf__ T * tmpExpSumUb, __ubuf__ T * tmpMaxUb, __ubuf__ T * tmpMaxUb2, const uint32_t blockStride,
const uint32_t repeatStride, const uint16_t m, const T scale, const T minValue, uint32_t pltOriginalN,
uint32_t pltSrcN)
{
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
AscendC::MicroAPI::RegTensor<float> vreg_in_max;
AscendC::MicroAPI::RegTensor<float> vreg_max_new;
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
AscendC::MicroAPI::RegTensor<float> vreg_cur_max;
AscendC::MicroAPI::RegTensor<float> vreg_exp;
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
// bfloat16_t
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_dst_even_bf16;
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_dst_odd_bf16;
// half
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
AscendC::MicroAPI::RegTensor<half> vreg_dst_even_fp16;
AscendC::MicroAPI::RegTensor<half> vreg_dst_odd_fp16;
AscendC::MicroAPI::UnalignRegForStore ureg_max;
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
AscendC::MicroAPI::MaskReg preg_all_b16 =
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
AscendC::MicroAPI::MaskReg preg_ori_src_n = AscendC::MicroAPI::UpdateMask<T>(pltOriginalN);
AscendC::MicroAPI::MaskReg preg_src_n = AscendC::MicroAPI::UpdateMask<T>(pltSrcN);
AscendC::MicroAPI::MaskReg preg_src_n_b16 =
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::H>();
// x_max = max(src, axis=-1, keepdims=True)
for (uint16_t i = 0; i < m; ++i) {
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_ori_src_n);
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_src_n);
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
vreg_cur_max, vreg_input_x, preg_ori_src_n);
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)tmpMaxUb), vreg_cur_max, ureg_max, 1);
}
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)tmpMaxUb), ureg_max, 0);
AscendC::MicroAPI::LoadAlign(vreg_in_max, inMaxUb);
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
AscendC::MicroAPI::LoadAlign(vreg_cur_max, tmpMaxUb2);
AscendC::MicroAPI::Max(vreg_max_new, vreg_cur_max, vreg_in_max, preg_all); // 计算新、旧的最大值
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)tmpMaxUb2, vreg_max_new, preg_all);
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
for (uint16_t i = 0; i < m; ++i) {
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(
vreg_max_brc, tmpMaxUb2 + i);
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
AscendC::MicroAPI::ExpSub(vreg_exp, vreg_input_x, vreg_max_brc, preg_ori_src_n);
// x_sum = sum(x_exp, axis=-1, keepdims=True)
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
vreg_exp_sum, vreg_exp, preg_ori_src_n);
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)tmpExpSumUb), vreg_exp_sum, ureg_exp_sum, 1);
if constexpr (IsSameType<T2, bfloat16_t>::value) {
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_bf16, vreg_exp, preg_all_b16);
AscendC::MicroAPI::DeInterleave(vreg_dst_even_bf16, vreg_dst_odd_bf16,
vreg_exp_bf16, vreg_exp_bf16);
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T2 *&)expUb), vreg_dst_even_bf16, blockStride, repeatStride, preg_src_n_b16);
} else if constexpr (IsSameType<T2, half>::value) {
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_fp16, vreg_exp, preg_all_b16);
AscendC::MicroAPI::DeInterleave(vreg_dst_even_fp16, vreg_dst_odd_fp16,
vreg_exp_fp16, vreg_exp_fp16);
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T2 *&)expUb), vreg_dst_even_fp16, blockStride, repeatStride, preg_src_n_b16);
}
}
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
((__ubuf__ T *&)tmpExpSumUb), ureg_exp_sum, 0);
}
// update, originN <= 64
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
__aicore__ inline void ProcessVec1UpdateImpl64(
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor, const LocalTensor<T>& inMaxTensor,
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
{
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
const uint32_t repeatStride = 1;
uint32_t pltOriginalN = originN;
uint32_t pltSrcN = s2BaseSize;
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
__ubuf__ T * inMaxUb = (__ubuf__ T*)inMaxTensor.GetPhyAddr();
__ubuf__ T * tmpExpSumUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr();
__ubuf__ T * tmpMaxUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
__ubuf__ T * tmpMaxUb2 = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
ProcessVec1UpdateImpl64VF <T, T2, s1BaseSize, s2BaseSize>(
expUb, srcUb, inMaxUb, tmpExpSumUb, tmpMaxUb, tmpMaxUb2, blockStride, repeatStride, m, scale, minValue,
pltOriginalN, pltSrcN);
}
} // namespace
#endif // VF_BASIC_BLOCK_ALIGNED64_UPDATE_SFA_H

View File

@@ -0,0 +1,112 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file vf_basic_block_utils.h
* \brief
*/
#ifndef VF_BASIC_BLOCK_UTILS_H
#define VF_BASIC_BLOCK_UTILS_H
#if ASC_DEVKIT_MAJOR >= 9
#include "kernel_basic_intf.h"
#else
#include "kernel_operator.h"
#endif
namespace FaVectorApi {
constexpr uint32_t floatRepSize = 64;
constexpr uint32_t halfRepSize = 128;
constexpr uint32_t blockBytesU8 = 32;
constexpr float fp8e4m3MaxValue = 448.0f;
constexpr float int8MaxValue = 127.0f;
constexpr float hifp8MaxValue = 32768.0f;
constexpr float floatEps = 2.220446049250313e-16;
/* **************************************************************************************************
* Muls + Select(optional) + SoftmaxFlashV2 + Cast(fp32->fp16/bf16) + ND2NZ
* ************************************************************************************************* */
using namespace MicroAPI;
constexpr static AscendC::MicroAPI::CastTrait castTraitZero = {
AscendC::MicroAPI::RegLayout::ZERO,
AscendC::MicroAPI::SatMode::SAT,
AscendC::MicroAPI::MaskMergeMode::ZEROING,
AscendC::RoundMode::CAST_ROUND,
};
constexpr static AscendC::MicroAPI::CastTrait castTraitOne = {
AscendC::MicroAPI::RegLayout::ONE,
AscendC::MicroAPI::SatMode::SAT,
AscendC::MicroAPI::MaskMergeMode::ZEROING,
AscendC::RoundMode::CAST_ROUND,
};
constexpr static AscendC::MicroAPI::CastTrait castTraitTwo = {
AscendC::MicroAPI::RegLayout::TWO,
AscendC::MicroAPI::SatMode::SAT,
AscendC::MicroAPI::MaskMergeMode::ZEROING,
AscendC::RoundMode::CAST_ROUND,
};
constexpr static AscendC::MicroAPI::CastTrait castTraitThree = {
AscendC::MicroAPI::RegLayout::THREE,
AscendC::MicroAPI::SatMode::SAT,
AscendC::MicroAPI::MaskMergeMode::ZEROING,
AscendC::RoundMode::CAST_ROUND,
};
constexpr static AscendC::MicroAPI::CastTrait castTraitRintZero = {
AscendC::MicroAPI::RegLayout::ZERO,
AscendC::MicroAPI::SatMode::SAT,
AscendC::MicroAPI::MaskMergeMode::ZEROING,
AscendC::RoundMode::CAST_RINT,
};
constexpr static AscendC::MicroAPI::CastTrait castTraitRintOne = {
AscendC::MicroAPI::RegLayout::ONE,
AscendC::MicroAPI::SatMode::SAT,
AscendC::MicroAPI::MaskMergeMode::ZEROING,
AscendC::RoundMode::CAST_RINT,
};
constexpr static AscendC::MicroAPI::CastTrait castTraitRintTwo = {
AscendC::MicroAPI::RegLayout::TWO,
AscendC::MicroAPI::SatMode::SAT,
AscendC::MicroAPI::MaskMergeMode::ZEROING,
AscendC::RoundMode::CAST_RINT,
};
constexpr static AscendC::MicroAPI::CastTrait castTraitRintThree = {
AscendC::MicroAPI::RegLayout::THREE,
AscendC::MicroAPI::SatMode::SAT,
AscendC::MicroAPI::MaskMergeMode::ZEROING,
AscendC::RoundMode::CAST_RINT,
};
#define USE_MLA_FULLQUANT_V1_P(vreg_exp, vreg_rowmax_p, MaskReg) \
do { \
Muls(vreg_exp, vreg_exp, fp8e4m3MaxValue, MaskReg); \
Div(vreg_exp, vreg_exp, vreg_rowmax_p, MaskReg); \
} while (0)
#define USE_MLA_FULLQUANT_V1_P_INT8(vreg_exp, vreg_rowmax_p, MaskReg) \
do { \
Muls(vreg_exp, vreg_exp, int8MaxValue, MaskReg); \
Div(vreg_exp, vreg_exp, vreg_rowmax_p, MaskReg); \
} while (0)
#define USE_MLA_FULLQUANT_V1_P_HIFP8(vreg_exp, vreg_rowmax_p, MaskReg) \
do { \
Muls(vreg_exp, vreg_exp, hifp8MaxValue, MaskReg); \
Div(vreg_exp, vreg_exp, vreg_rowmax_p, MaskReg); \
} while (0)
} // namespace
#endif // VF_BASIC_BLOCK_UTILS_H

View File

@@ -0,0 +1,727 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file vf_flashupdate_new.h
* \brief
*/
#ifndef MY_FLASH_UPDATE_NEW_INTERFACE_H
#define MY_FLASH_UPDATE_NEW_INTERFACE_H
#include "kernel_tensor.h"
namespace FaVectorApi {
// bf16->fp32
static constexpr MicroAPI::CastTrait castTraitFp16_32_update = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN,
MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
constexpr uint16_t REDUCE_SIZE = 1;
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize, bool isUpdatePre, bool isMlaFullQuant>
__simd_vf__ inline void FlashUpdateBasicVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * preUb,
__ubuf__ float * expMaxUb, __ubuf__ float * rowMaxUb, const uint16_t m, const uint16_t d,
const float deScaleV, const float deScaleVPre)
{
constexpr uint16_t floatRepSize = 64;
constexpr uint16_t dLoops = srcD / floatRepSize;
RegTensor<float> vreg_exp_max;
RegTensor<float> vreg_row_max;
RegTensor<float> vreg_input_pre;
RegTensor<float> vreg_input_cur;
RegTensor<float> vreg_mul;
RegTensor<float> vreg_add;
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
// dstTensor = preTensor * expMaxTensor + curTensor
for (uint16_t i = 0; i < m; ++i) {
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize); // [m,8]
if constexpr (isMlaFullQuant) {
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_row_max, rowMaxUb + i * reduceSize);
}
for (uint16_t j = 0; j < dLoops; ++j) {
LoadAlign(vreg_input_pre, preUb + i * d + j * floatRepSize);
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
if constexpr (isMlaFullQuant) {
Mul(vreg_input_cur, vreg_input_cur, vreg_row_max, preg_all);
}
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_all);
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
IsSameType<INPUT_T, hifloat8_t>::value ||
IsSameType<INPUT_T, int8_t>::value) {
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
if constexpr (isUpdatePre) {
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
}
}
Add(vreg_add, vreg_mul, vreg_input_cur, preg_all);
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_add, preg_all);
}
}
}
/* **************************************************************************************************
* FlashUpdate, fp32
* ************************************************************************************************* */
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize, bool isUpdatePre, bool isMlaFullQuant>
__aicore__ inline void FlashUpdateBasic(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
const LocalTensor<T>& preTensor, const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& rowMaxTensor,
const uint16_t m, const uint16_t d, const float deScaleV, const float deScaleVPre)
{
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
__ubuf__ float * preUb = (__ubuf__ T*)preTensor.GetPhyAddr();
__ubuf__ float * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
__ubuf__ float * rowMaxUb = (__ubuf__ T*)rowMaxTensor.GetPhyAddr();
FlashUpdateBasicVF<T, INPUT_T, OUTPUT_T, srcD, reduceSize, isUpdatePre, isMlaFullQuant>(
dstUb, curUb, preUb, expMaxUb, rowMaxUb, m, d, deScaleV, deScaleVPre);
}
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t reduceSize, bool isUpdatePre>
__simd_vf__ inline void FlashUpdateGeneralVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * preUb,
__ubuf__ float * expMaxUb, const uint16_t m, const uint16_t d,
const float deScaleV, const float deScaleVPre, const uint32_t pltTailD, const uint16_t hasTail)
{
RegTensor<float> vreg_exp_max;
RegTensor<float> vreg_input_pre;
RegTensor<float> vreg_input_cur;
RegTensor<float> vreg_mul;
RegTensor<float> vreg_add;
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
uint32_t tmpTailD = pltTailD;
MaskReg preg_tail_d = UpdateMask<float>(tmpTailD);
constexpr uint16_t floatRepSize = 64;
const uint16_t dLoops = d / floatRepSize;
// dstTensor = preTensor * expMaxTensor + curTensor
for (uint16_t i = 0; i < m; ++i) {
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize); // [m,8]
for (uint16_t j = 0; j < dLoops; ++j) {
LoadAlign(vreg_input_pre, preUb + i * d + j * floatRepSize);
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_all);
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
IsSameType<INPUT_T, hifloat8_t>::value ||
IsSameType<INPUT_T, int8_t>::value) {
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
if constexpr (isUpdatePre) {
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
}
}
Add(vreg_add, vreg_mul, vreg_input_cur, preg_all);
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_add, preg_all);
}
for (uint16_t t = 0; t < hasTail; ++t) {
LoadAlign(vreg_input_pre, preUb + i * d + dLoops * floatRepSize);
LoadAlign(vreg_input_cur, curUb + i * d + dLoops * floatRepSize);
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_tail_d);
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
IsSameType<INPUT_T, hifloat8_t>::value ||
IsSameType<INPUT_T, int8_t>::value) {
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
if constexpr (isUpdatePre) {
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
}
}
Add(vreg_add, vreg_mul, vreg_input_cur, preg_tail_d);
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)dstUb + i * d + dLoops * floatRepSize, vreg_add, preg_tail_d);
}
}
}
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t reduceSize, bool isUpdatePre>
__aicore__ inline void FlashUpdateGeneral(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
const LocalTensor<T>& preTensor, const LocalTensor<T>& expMaxTensor, const uint16_t m, const uint16_t d,
const float deScaleV, const float deScaleVPre)
{
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
__ubuf__ float * preUb = (__ubuf__ T*)preTensor.GetPhyAddr();
__ubuf__ float * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
constexpr uint16_t floatRepSize = 64;
const uint16_t tailD = d % floatRepSize;
uint32_t pltTailD = static_cast<uint32_t>(tailD);
uint16_t hasTail = 0;
if (tailD > 0) {
hasTail = 1;
}
FlashUpdateGeneralVF<T, INPUT_T, OUTPUT_T, reduceSize, isUpdatePre>(
dstUb, curUb, preUb, expMaxUb, m, d, deScaleV, deScaleVPre, pltTailD, hasTail);
}
/*
* @ingroup FlashUpdate
* @brief compute, dstTensor = preTensor * expMaxTensor + curTensor
* @param [out] dstTensor, output LocalTensor
* @param [in] curTensor, input LocalTensor
* @param [in] preTensor, input LocalTensor
* @param [in] expMaxTensor, input LocalTensor
* @param [in] m, input rows
* @param [in] d, input columns, should be 32 bytes aligned
*/
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, bool isUpdatePre, bool isMlaFullQuant>
__aicore__ inline void FlashUpdateNew(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
const LocalTensor<T>& preTensor, const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& rowMaxTensor, const uint16_t m, const uint16_t d,
const float deScaleV, const float deScaleVPre)
{
static_assert(IsSameType<T, float>::value, "VF FlashUpdate, T must be float");
constexpr uint16_t floatRepSize = 64;
if constexpr(srcD % floatRepSize == 0) {
FlashUpdateBasic<T, INPUT_T, OUTPUT_T, srcD, REDUCE_SIZE, isUpdatePre, isMlaFullQuant>(dstTensor, curTensor, preTensor, expMaxTensor, rowMaxTensor,
m, d, deScaleV, deScaleVPre);
} else {
FlashUpdateGeneral<T, INPUT_T, OUTPUT_T, REDUCE_SIZE, isUpdatePre>(dstTensor, curTensor, preTensor, expMaxTensor, m, d,
deScaleV, deScaleVPre);
}
}
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize, bool isUpdatePre, bool isMlaFullQuant>
__simd_vf__ inline void FlashUpdateLastBasicVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * preUb,
__ubuf__ float * expMaxUb, __ubuf__ float * expSumUb, __ubuf__ float * rowMaxUb, const uint16_t m, const uint16_t d,
const float deScaleV, const float deScaleVPre)
{
RegTensor<float> vreg_exp_max;
RegTensor<float> vreg_row_max;
RegTensor<float> vreg_input_pre;
RegTensor<float> vreg_input_cur;
RegTensor<float> vreg_mul;
RegTensor<float> vreg_add;
RegTensor<float> vreg_div;
RegTensor<half> vreg_cast;
RegTensor<float> vreg_exp_sum;
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
constexpr uint16_t floatRepSize = 64;
constexpr uint16_t dLoops = srcD / floatRepSize;
constexpr float fp8e4m3MaxValueRec = 1 / 448.0f;
constexpr float int8MaxValueRec = 1 / 127.0f;
constexpr float hifp8MaxValueRec = 1 / 32768.0f;
for (uint16_t i = 0; i < m; ++i) {
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize);
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_sum, expSumUb + i * reduceSize);
if constexpr (isMlaFullQuant) {
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_row_max, rowMaxUb + i * reduceSize);
}
for (uint16_t j = 0; j < dLoops; ++j) {
LoadAlign(vreg_input_pre, preUb + i * d + j * floatRepSize);
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
if constexpr (isMlaFullQuant) {
Mul(vreg_input_cur, vreg_input_cur, vreg_row_max, preg_all);
}
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_all);
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
IsSameType<INPUT_T, hifloat8_t>::value ||
IsSameType<INPUT_T, int8_t>::value) {
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
if constexpr (isUpdatePre) {
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
}
}
Add(vreg_add, vreg_mul, vreg_input_cur, preg_all);
Div(vreg_div, vreg_add, vreg_exp_sum, preg_all);
if constexpr (isMlaFullQuant) {
if constexpr (IsSameType<INPUT_T, fp8_e4m3fn_t>::value) {
Muls(vreg_div, vreg_div, fp8e4m3MaxValueRec, preg_all);
} else if constexpr (IsSameType<INPUT_T, int8_t>::value) {
Muls(vreg_div, vreg_div, int8MaxValueRec, preg_all);
} else {
Muls(vreg_div, vreg_div, hifp8MaxValueRec, preg_all);
}
}
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_div, preg_all);
}
}
}
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize, bool isUpdatePre, bool isMlaFullQuant>
__aicore__ inline void FlashUpdateLastBasic(const LocalTensor<T>& dstTensor,
const LocalTensor<T>& curTensor, const LocalTensor<T>& preTensor,
const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& rowMaxTensor, const LocalTensor<T>& expSumTensor,
const uint16_t m, const uint16_t d, const float deScaleV, const float deScaleVPre)
{
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
__ubuf__ float * preUb = (__ubuf__ T*)preTensor.GetPhyAddr();
__ubuf__ float * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
__ubuf__ float * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
__ubuf__ float * rowMaxUb = (__ubuf__ T*)rowMaxTensor.GetPhyAddr();
FlashUpdateLastBasicVF<T, INPUT_T, OUTPUT_T, srcD, reduceSize, isUpdatePre, isMlaFullQuant>(
dstUb, curUb, preUb, expMaxUb, expSumUb, rowMaxUb, m, d, deScaleV, deScaleVPre);
}
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t reduceSize, bool isUpdatePre>
__simd_vf__ inline void FlashUpdateLastGeneralVF(__ubuf__ float * dstUb, __ubuf__ float * curUb,
__ubuf__ float * preUb, __ubuf__ float * expMaxUb, __ubuf__ float * expSumUb, const uint16_t m, const uint16_t d,
const float deScaleV, const float deScaleVPre, const uint32_t pltTailD, const uint16_t hasTail)
{
RegTensor<float> vreg_exp_max;
RegTensor<float> vreg_input_pre;
RegTensor<float> vreg_input_cur;
RegTensor<float> vreg_mul;
RegTensor<float> vreg_add;
RegTensor<float> vreg_div;
RegTensor<half> vreg_cast;
RegTensor<float> vreg_exp_sum;
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
uint32_t tmpTailD = pltTailD;
MaskReg preg_tail_d = UpdateMask<float>(tmpTailD);
constexpr uint16_t floatRepSize = 64;
uint16_t dLoops = d / floatRepSize;
for (uint16_t i = 0; i < m; ++i) {
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize);
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_sum, expSumUb + i * reduceSize);
for (uint16_t j = 0; j < dLoops; ++j) {
LoadAlign(vreg_input_pre, preUb + i * d + j * floatRepSize);
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_all);
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
IsSameType<INPUT_T, hifloat8_t>::value ||
IsSameType<INPUT_T, int8_t>::value) {
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
if constexpr (isUpdatePre) {
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
}
}
Add(vreg_add, vreg_mul, vreg_input_cur, preg_all);
Div(vreg_div, vreg_add, vreg_exp_sum, preg_all);
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_div, preg_all);
}
for (uint16_t t = 0; t < hasTail; ++t) {
LoadAlign(vreg_input_pre, preUb + i * d + dLoops * floatRepSize);
LoadAlign(vreg_input_cur, curUb + i * d + dLoops * floatRepSize);
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_tail_d);
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
IsSameType<INPUT_T, hifloat8_t>::value ||
IsSameType<INPUT_T, int8_t>::value) {
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
if constexpr (isUpdatePre) {
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
}
}
Add(vreg_add, vreg_mul, vreg_input_cur, preg_tail_d);
Div(vreg_div, vreg_add, vreg_exp_sum, preg_tail_d);
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)dstUb + i * d + dLoops * floatRepSize, vreg_div, preg_tail_d);
}
}
}
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t reduceSize, bool isUpdatePre>
__aicore__ inline void FlashUpdateLastGeneral(const LocalTensor<T>& dstTensor,
const LocalTensor<T>& curTensor, const LocalTensor<T>& preTensor,
const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& expSumTensor,
const uint16_t m, const uint16_t d, const float deScaleV, const float deScaleVPre)
{
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
__ubuf__ float * preUb = (__ubuf__ T*)preTensor.GetPhyAddr();
__ubuf__ float * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
__ubuf__ float * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
constexpr uint16_t floatRepSize = 64;
uint16_t tailD = d % floatRepSize;
uint32_t pltTailD = tailD;
uint16_t hasTail = 0;
if (tailD > 0) {
hasTail = 1;
}
FlashUpdateLastGeneralVF<T, INPUT_T, OUTPUT_T, reduceSize, isUpdatePre>(
dstUb, curUb, preUb, expMaxUb, expSumUb, m, d, deScaleV, deScaleVPre, pltTailD, hasTail);
}
/*
* @ingroup FlashUpdateLast
* @brief compute, dstTensor = (preTensor * expMaxTensor + curTensor) / expSumTensor
* @param [out] dstTensor, output LocalTensor
* @param [in] curTensor, input LocalTensor
* @param [in] preTensor, input LocalTensor
* @param [in] expMaxTensor, input LocalTensor
* @param [in] expSumTensor, input LocalTensor
* @param [in] m, input rows
* @param [in] d, input columns, 32 bytes align
*/
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, bool isUpdatePre, bool isMlaFullQuant>
__aicore__ inline void FlashUpdateLastNew(const LocalTensor<T>& dstTensor,
const LocalTensor<T>& curTensor, const LocalTensor<T>& preTensor,
const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& rowMaxTensor, const LocalTensor<T>& expSumTensor,
uint16_t m, uint16_t d, const float deScaleV, const float deScaleVPre)
{
static_assert(IsSameType<T, float>::value, "VF FlashUpdateLast, T must be float");
constexpr uint16_t floatRepSize = 64;
if constexpr(srcD % floatRepSize == 0) {
FlashUpdateLastBasic<T, INPUT_T, OUTPUT_T, srcD, REDUCE_SIZE, isUpdatePre, isMlaFullQuant>(
dstTensor, curTensor, preTensor, expMaxTensor, rowMaxTensor, expSumTensor, m, d, deScaleV, deScaleVPre);
} else {
FlashUpdateLastGeneral<T, INPUT_T, OUTPUT_T, REDUCE_SIZE, isUpdatePre>(
dstTensor, curTensor, preTensor, expMaxTensor, expSumTensor, m, d, deScaleV, deScaleVPre);
}
}
template <typename T, typename INPUT_T, typename OUTPUT_T, uint32_t srcD, bool isMlaFullQuant>
__simd_vf__ inline void LastDivNewVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * expSumUb,
const uint16_t m, const uint16_t d, const float deScaleV)
{
RegTensor<float> vreg_input_cur;
RegTensor<float> vreg_div;
RegTensor<float> vreg_exp_sum;
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
constexpr uint16_t floatRepSize = 64;
const uint16_t dLoops = d >> 6;
constexpr float fp8e4m3MaxValueRec = 1 / 448.0f;
constexpr float int8MaxValueRec = 1 / 127.0f;
constexpr float hifp8MaxValueRec = 1 / 32768.0f;
for (uint16_t i = 0; i < m; ++i) {
uint32_t sreg_init = d;
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_sum, expSumUb + i * REDUCE_SIZE);
for (uint16_t j = 0; j < dLoops; ++j) {
MaskReg preg_update = UpdateMask<float>(sreg_init);
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
IsSameType<INPUT_T, hifloat8_t>::value ||
IsSameType<INPUT_T, int8_t>::value) {
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
}
Div(vreg_div, vreg_input_cur, vreg_exp_sum, preg_update);
if constexpr (isMlaFullQuant) {
if constexpr (IsSameType<INPUT_T, fp8_e4m3fn_t>::value) {
Muls(vreg_div, vreg_div, fp8e4m3MaxValueRec, preg_all);
} else if constexpr (IsSameType<INPUT_T, int8_t>::value) {
Muls(vreg_div, vreg_div, int8MaxValueRec, preg_all);
} else {
Muls(vreg_div, vreg_div, hifp8MaxValueRec, preg_all);
}
}
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_div, preg_update);
}
}
}
// dstTensor = curTensor / expSumTensor, curTensor: [64,128], expSumTensor: [64,8]
template <typename T, typename INPUT_T, typename OUTPUT_T, uint32_t srcD, bool isMlaFullQuant>
__aicore__ inline void LastDivNew(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
const LocalTensor<T>& expSumTensor, const uint16_t m, const uint16_t d, const float deScaleV)
{
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
__ubuf__ float * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
LastDivNewVF<T, INPUT_T, OUTPUT_T, srcD, isMlaFullQuant>(dstUb, curUb, expSumUb, m, d, deScaleV);
}
template <typename T, uint32_t srcD>
__simd_vf__ inline void InvalidLineUpdateVF(__ubuf__ T * dstUb, __ubuf__ T * srcUb, __ubuf__ T * maxUb,
const uint16_t m, const uint16_t d, const T minValue, const T invalidValue)
{
RegTensor<float> vreg_invalid_value;
RegTensor<float> vreg_max;
RegTensor<float> vreg_input;
RegTensor<float> vreg_input_brc;
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
MaskReg preg_compare;
const uint16_t dLoops = d >> 6;
Duplicate(vreg_invalid_value, invalidValue);
for (uint16_t i = 0; i < m; ++i) {
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_max, maxUb + i);
Compares<T, CMPMODE::EQ>(preg_compare, vreg_max, minValue, preg_all);
for (uint16_t j = 0; j < dLoops; ++j) {
LoadAlign(vreg_input, srcUb + i * d + j * floatRepSize);
Select(vreg_input_brc, vreg_invalid_value, vreg_input, preg_compare);
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_input_brc, preg_all);
}
}
}
template <typename T, uint32_t srcD>
__aicore__ inline void InvalidLineUpdate(const LocalTensor<T>& dstTensor, const LocalTensor<T>& srcTensor,
const LocalTensor<T>& maxTensor, const uint16_t m, const uint16_t d, const T minValue, const T invalidValue)
{
__ubuf__ T * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
constexpr uint16_t floatRepSize = 64;
uint16_t dLoops = d >> 6;
InvalidLineUpdateVF<T, srcD>(dstUb, srcUb, maxUb, m, d, minValue, invalidValue);
}
template <typename T>
__simd_vf__ inline void ComputeLseOutputVF(__ubuf__ T *srcSumUb, __ubuf__ T *srcMaxUb, __ubuf__ T *dstUb, const uint32_t dealCount)
{
MicroAPI::RegTensor<T> vregSum;
MicroAPI::RegTensor<T> vregMax;
MicroAPI::RegTensor<T> vregRes;
MicroAPI::RegTensor<T> vregResFinal;
MicroAPI::RegTensor<float> vregMinValue;
MicroAPI::RegTensor<float> vregInfValue;
MicroAPI::MaskReg pregCompare;
constexpr uint32_t dealRows = 8;
constexpr uint32_t floatRepSize = 64; // 64: 一个寄存器存64个float
constexpr float infValue = 3e+99; // 3e+99 for float inf
constexpr uint32_t tmpMin = 0xFF7FFFFF;
float minValue = *((float*)&tmpMin);
uint16_t updateLoops = dealCount / dealRows;
uint16_t tailLSize = dealCount % dealRows * 8;
uint32_t pltTail = static_cast<uint32_t>(tailLSize);
MicroAPI::MaskReg pregAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
MicroAPI::MaskReg pregTail = MicroAPI::UpdateMask<T>(pltTail);
MicroAPI::Duplicate<float, float>(vregMinValue, minValue);
MicroAPI::Duplicate<float, float>(vregInfValue, infValue);
for (uint16_t i = 0; i < updateLoops; ++i) {
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(vregSum, srcSumUb + (i * dealRows));
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(vregMax, srcMaxUb + (i * dealRows));
MicroAPI::Log<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSum, pregAll);
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, vregMax, pregAll);
MicroAPI::Compare<float, CMPMODE::EQ>(pregCompare, vregMax, vregMinValue, pregAll);
MicroAPI::Select<T>(vregResFinal, vregInfValue, vregRes, pregCompare);
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(dstUb + (i * floatRepSize), vregResFinal, pregAll);
}
if (tailLSize != 0) {
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(vregSum, srcSumUb + dealRows * updateLoops);
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(vregMax, srcMaxUb + dealRows * updateLoops);
MicroAPI::Log<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSum, pregTail);
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, vregMax, pregTail);
MicroAPI::Compare<float, CMPMODE::EQ>(pregCompare, vregMax, vregMinValue, pregTail);
MicroAPI::Select<T>(vregResFinal, vregInfValue, vregRes, pregCompare);
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(dstUb + floatRepSize * updateLoops, vregResFinal, pregTail);
}
}
template <typename T>
__aicore__ inline void ComputeLseOutputVF(const LocalTensor<T>& dstTensor, const LocalTensor<T>& softmaxSumTensor,
const LocalTensor<T>& softmaxMaxTensor, uint32_t dealCount)
{
__ubuf__ T * srcSumUb = (__ubuf__ T *)softmaxSumTensor.GetPhyAddr();
__ubuf__ T * srcMaxUb = (__ubuf__ T *)softmaxMaxTensor.GetPhyAddr();
__ubuf__ T * dstUb = (__ubuf__ T *)dstTensor.GetPhyAddr();
ComputeLseOutputVF<T>(srcSumUb, srcMaxUb, dstUb, dealCount);
}
template <typename T>
__simd_vf__ inline void SinkSubExpAddVF(__ubuf__ T *srcSumUb, __ubuf__ T *srcMaxUb, const T sinkValue, const uint32_t dealCount)
{
MicroAPI::RegTensor<T> vregSum;
MicroAPI::RegTensor<T> vregMax;
MicroAPI::RegTensor<T> vregRes;
MicroAPI::RegTensor<T> vregSink;
constexpr uint32_t floatRepSize = 64;
uint16_t updateLoops = dealCount / floatRepSize;
uint16_t tailSize = dealCount % floatRepSize;
uint32_t pltTail = static_cast<uint32_t>(tailSize);
//mask
MicroAPI::MaskReg pregAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
MicroAPI::MaskReg pregTail = MicroAPI::UpdateMask<T>(pltTail);
Duplicate(vregSink, sinkValue);
for (uint16_t i = 0; i < updateLoops; ++i) {
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregSum, srcSumUb + (i * floatRepSize));
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregMax, srcMaxUb + (i * floatRepSize));
MicroAPI::Sub<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSink, vregMax, pregAll);
MicroAPI::Exp<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, pregAll);
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregSum, vregSum, vregRes, pregAll);
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(srcSumUb + (i * floatRepSize), vregSum, pregAll);
}
for (uint16_t i = 0; i < tailSize; i = i + tailSize) {
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregSum, srcSumUb + (updateLoops * floatRepSize));
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregMax, srcMaxUb + (updateLoops * floatRepSize));
MicroAPI::Sub<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSink, vregMax, pregTail);
MicroAPI::Exp<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, pregTail);
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregSum, vregSum, vregRes, pregTail);
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(srcSumUb + (updateLoops * floatRepSize), vregSum, pregTail);
}
}
template <typename T>
__aicore__ inline void SinkSubExpAddVF(const LocalTensor<T>& softmaxSumTensor, const LocalTensor<T>& softmaxMaxTensor,
const T sinkValue, uint32_t dealCount)
{
__ubuf__ T * srcSumUb = (__ubuf__ T *)softmaxSumTensor.GetPhyAddr();
__ubuf__ T * srcMaxUb = (__ubuf__ T *)softmaxMaxTensor.GetPhyAddr();
SinkSubExpAddVF<T>(srcSumUb, srcMaxUb, sinkValue, dealCount);
}
template <typename T, typename SINK_T>
__simd_vf__ inline void SinkSubExpAddGSFusedVF(__ubuf__ T *srcSumUb, __ubuf__ T *srcMaxUb, __ubuf__ uint16_t *sinkUb, const uint32_t dealCount)
{
MicroAPI::RegTensor<T> vregSum;
MicroAPI::RegTensor<T> vregMax;
MicroAPI::RegTensor<T> vregRes;
MicroAPI::RegTensor<SINK_T> vregSink;
MicroAPI::RegTensor<T> vregSinkCast;
constexpr uint32_t floatRepSize = 64;
uint16_t updateLoops = dealCount / floatRepSize;
uint16_t tailSize = dealCount % floatRepSize;
uint32_t pltTail = static_cast<uint32_t>(tailSize);
//mask
MicroAPI::MaskReg pregAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
MicroAPI::MaskReg pregTail = MicroAPI::UpdateMask<T>(pltTail);
MicroAPI::MaskReg pregSinkAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
MicroAPI::LoadAlign<uint16_t, MicroAPI::LoadDist::DIST_UNPACK_B16>((MicroAPI::RegTensor<uint16_t>&)vregSink, sinkUb);
MicroAPI::Cast<T, SINK_T, castTraitFp16_32_update>(vregSinkCast, vregSink, pregSinkAll);
for (uint16_t i = 0; i < updateLoops; ++i) {
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregSum, srcSumUb + (i * floatRepSize));
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregMax, srcMaxUb + (i * floatRepSize));
MicroAPI::Sub<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSinkCast, vregMax, pregAll);
MicroAPI::Exp<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, pregAll);
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregSum, vregSum, vregRes, pregAll);
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(srcSumUb + (i * floatRepSize), vregSum, pregAll);
}
if (tailSize != 0) {
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregSum, srcSumUb + (updateLoops * floatRepSize));
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregMax, srcMaxUb + (updateLoops * floatRepSize));
MicroAPI::Sub<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSinkCast, vregMax, pregTail);
MicroAPI::Exp<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, pregTail);
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregSum, vregSum, vregRes, pregTail);
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(srcSumUb + (updateLoops * floatRepSize), vregSum, pregTail);
}
}
template <typename T, typename SINK_T>
__aicore__ inline void SinkSubExpAddGSFusedVF(const LocalTensor<SINK_T>& dstTensor, const LocalTensor<T>& softmaxSumTensor,
const LocalTensor<T>& softmaxMaxTensor, uint32_t dealCount)
{
__ubuf__ T * srcSumUb = (__ubuf__ T *)softmaxSumTensor.GetPhyAddr();
__ubuf__ T * srcMaxUb = (__ubuf__ T *)softmaxMaxTensor.GetPhyAddr();
__ubuf__ uint16_t * dstUb = (__ubuf__ uint16_t *)dstTensor.GetPhyAddr();
SinkSubExpAddGSFusedVF<T, SINK_T>(srcSumUb, srcMaxUb, dstUb, dealCount);
}
template <typename T>
__simd_vf__ inline void RowInvalidUpdateVF(__ubuf__ T *finalUb, __ubuf__ float *maxUb, const uint16_t m,
const uint16_t d, int64_t dSize, const uint32_t pltTailD, const uint16_t hasTail)
{
constexpr uint16_t floatRepSize = 64; // 64: 一个寄存器可以存储64个float类型数据
const uint16_t dLoops = d / floatRepSize;
constexpr uint32_t tmpZero = 0x00000000; // zero value of fp16 and fp32
const T zeroValue = *((T*)&tmpZero);
constexpr uint32_t tmpMin = 0xFF7FFFFF; // min value of float
const float minValue = *((float*)&tmpMin);
MicroAPI::RegTensor<float> vregMinValue;
MicroAPI::RegTensor<T> vregZeroValue;
MicroAPI::RegTensor<float> vregMax;
MicroAPI::RegTensor<T> vregFinal;
MicroAPI::RegTensor<T> vregFinalNew;
MicroAPI::MaskReg pregAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
uint32_t tmpTailD = pltTailD;
MicroAPI::MaskReg pregTailD = MicroAPI::UpdateMask<T>(tmpTailD);
MicroAPI::MaskReg pregCompare;
MicroAPI::Duplicate<float, float>(vregMinValue, minValue);
MicroAPI::Duplicate<T, T>(vregZeroValue, zeroValue);
for (uint16_t i = 0; i < m; ++i) {
MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(vregMax, maxUb + i);
MicroAPI::Compare<float, CMPMODE::EQ>(pregCompare, vregMax, vregMinValue, pregAll);
for (uint16_t j = 0; j < dLoops; ++j) {
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregFinal, finalUb + i * dSize + j * floatRepSize);
MicroAPI::Select<T>(vregFinalNew, vregZeroValue, vregFinal, pregCompare);
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(finalUb + i * dSize + j * floatRepSize,
vregFinalNew, pregAll);
}
for (uint16_t t = 0; t < hasTail; ++t) {
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregFinal, finalUb + i * dSize + dLoops * floatRepSize);
MicroAPI::Select<T>(vregFinalNew, vregZeroValue, vregFinal, pregCompare);
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(finalUb + i * dSize + dLoops * floatRepSize,
vregFinalNew, pregTailD);
}
}
}
template <typename T>
__aicore__ inline void RowInvalidUpdateVF(const LocalTensor<T>& finalTensor, const LocalTensor<float>& maxTensor,
const uint16_t m, const uint16_t d, int64_t dSize)
{
__ubuf__ T * finalUb = (__ubuf__ T*)finalTensor.GetPhyAddr();
__ubuf__ float * maxUb = (__ubuf__ float*)maxTensor.GetPhyAddr();
constexpr uint16_t floatRepSize = 64;
const uint16_t tailD = d % floatRepSize;
uint32_t pltTailD = static_cast<uint32_t>(tailD);
uint16_t hasTail = 0;
if (tailD > 0) {
hasTail = 1;
}
RowInvalidUpdateVF<T>(finalUb, maxUb, m, d, dSize, pltTailD, hasTail);
}
} // namespace
#endif // MY_FLASH_UPDATE_INTERFACE_H

View File

@@ -0,0 +1,164 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file vf_mul_sel_softmaxflashv2_cast_nz_sfa.h
* \brief
*/
#ifndef MUL_SEL_SOFTMAX_FLASH_V2_CAST_NZ_SFA_INTERFACE_H
#define MUL_SEL_SOFTMAX_FLASH_V2_CAST_NZ_SFA_INTERFACE_H
#include "vf_basic_block_aligned128_no_update_sfa.h"
#include "vf_basic_block_aligned128_update_sfa.h"
#include "vf_basic_block_unaligned64_update_sfa.h"
#include "vf_basic_block_unaligned64_no_update_sfa.h"
#include "vf_basic_block_unaligned128_no_update_sfa.h"
#include "vf_basic_block_unaligned128_update_sfa.h"
using namespace regbaseutil;
namespace FaVectorApi {
/* **************************************************************************************************
* Muls + Select(optional) + SoftmaxFlashV2 + Cast(fp32->fp16/bf16) + ND2NZ
* ************************************************************************************************* */
using AscendC::LocalTensor;
enum class OriginNRange {
EQ_128_SFA = 0, // originN == 128, better performance than GT_64_AND_LTE_128 (s2BaseSize=128)
GT_0_AND_LTE_64_SFA, // 0 < originN <= 64 (s2BaseSize <= 64 or tail s2)
GT_64_AND_LTE_128_SFA, // 64 < originN <= 128, support for non-alignment (s2BaseSize=128)
N_INVALID_SFA
};
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128,
OriginNRange oriNRange = OriginNRange::EQ_128_SFA>
__aicore__ inline void ProcessVec1NoUpdate(
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
{
if constexpr (oriNRange == OriginNRange::EQ_128_SFA) {
ProcessVec1NoUpdateImpl128<T, T2, s1BaseSize, s2BaseSize>(
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
} else if constexpr (oriNRange == OriginNRange::GT_0_AND_LTE_64_SFA) {
ProcessVec1NoUpdateImpl64<T, T2, s1BaseSize, s2BaseSize>(
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
} else if constexpr (oriNRange == OriginNRange::GT_64_AND_LTE_128_SFA) {
ProcessVec1NoUpdateGeneralImpl128<T, T2, s1BaseSize, s2BaseSize>(
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
}
}
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128,
OriginNRange oriNRange = OriginNRange::EQ_128_SFA>
__aicore__ inline void ProcessVec1Update(
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
{
if constexpr (oriNRange == OriginNRange::EQ_128_SFA) {
ProcessVec1UpdateImpl128<T, T2, s1BaseSize, s2BaseSize>(
dstTensor, srcTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
} else if constexpr (oriNRange == OriginNRange::GT_0_AND_LTE_64_SFA) {
ProcessVec1UpdateImpl64<T, T2, s1BaseSize, s2BaseSize>(
dstTensor, srcTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
} else if constexpr (oriNRange == OriginNRange::GT_64_AND_LTE_128_SFA) {
ProcessVec1UpdateGeneralImpl128<T, T2, s1BaseSize, s2BaseSize>(
dstTensor, srcTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
}
}
template <typename T, typename T2, bool isUpdate = false, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128,
OriginNRange oriNRange = OriginNRange::EQ_128_SFA>
__aicore__ inline void ProcessVec1Vf(
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
{
static_assert(IsSameType<T, float>::value, "VF mul_sel_softmaxflashv2_cast_nz, T must be float");
static_assert((IsSameType<T2, half>::value || IsSameType<T2, bfloat16_t>::value),
"VF mul_sel_softmaxflashv2_cast_nz, T2 must be half or bfloat16");
if constexpr (!isUpdate) {
ProcessVec1NoUpdate<T, T2, s1BaseSize, s2BaseSize, oriNRange>(
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
} else {
ProcessVec1Update<T, T2, s1BaseSize, s2BaseSize, oriNRange>(
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
}
}
template <typename T>
__simd_vf__ inline void UpdateExpSumAndExpMaxVF(__ubuf__ T * maxUb, __ubuf__ T * inMaxUb, __ubuf__ T * expMaxUb,
__ubuf__ T * expSumUb, __ubuf__ T * inExpSumUb, __ubuf__ T * tmpExpSumUb, __ubuf__ T * tmpMaxUb, const uint32_t m)
{
RegTensor<float> vreg_input_x;
RegTensor<float> vreg_input_x_unroll;
RegTensor<float> vreg_max;
RegTensor<float> vreg_in_max;
RegTensor<float> vreg_exp_sum;
RegTensor<float> vreg_in_exp_sum;
RegTensor<float> vreg_exp_max;
RegTensor<float> vreg_exp_sum_brc;
RegTensor<float> vreg_exp_sum_update;
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
// 注意:当m大于64的时候需要开启循环
LoadAlign(vreg_max, tmpMaxUb);
LoadAlign(vreg_in_max, inMaxUb);
FusedExpSub(vreg_exp_max, vreg_in_max, vreg_max, preg_all);
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)expMaxUb, vreg_exp_max, preg_all);
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)maxUb, vreg_max, preg_all);
LoadAlign(vreg_in_exp_sum, inExpSumUb);
// x_sum = exp_max * insum + x_sum
LoadAlign(vreg_exp_sum_brc, tmpExpSumUb);
Mul(vreg_exp_sum_update, vreg_exp_max, vreg_in_exp_sum, preg_all);
Add(vreg_exp_sum_update, vreg_exp_sum_update, vreg_exp_sum_brc, preg_all);
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
(__ubuf__ T *&)expSumUb, vreg_exp_sum_update, preg_all);
}
template <typename T>
__aicore__ inline void SFAUpdateExpSumAndExpMax(
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor,
const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& inExpSumTensor,
const LocalTensor<T>& inMaxTensor, const LocalTensor<T>& sharedTmpBuffer, const uint32_t m)
{
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
__ubuf__ T * inMaxUb = (__ubuf__ T*)inMaxTensor.GetPhyAddr();
__ubuf__ T * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
__ubuf__ T * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
__ubuf__ T * inExpSumUb = (__ubuf__ T*)inExpSumTensor.GetPhyAddr();
__ubuf__ T * tmpExpSumUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr();
__ubuf__ T * tmpMaxUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
UpdateExpSumAndExpMaxVF<T>(maxUb, inMaxUb, expMaxUb, expSumUb, inExpSumUb, tmpExpSumUb, tmpMaxUb, m);
}
template <typename T>
__simd_vf__ inline void DuplicateSumWithR0VF(__ubuf__ T * sumUb, const T R0, uint32_t m) {
AscendC::MicroAPI::RegTensor<T> vreg_sum;
AscendC::MicroAPI::MaskReg preg_m = AscendC::MicroAPI::UpdateMask<T>(m);
AscendC::MicroAPI::UnalignRegForStore ureg;
AscendC::MicroAPI::Duplicate<T, MicroAPI::MaskMergeMode::ZEROING, T>(vreg_sum, R0, preg_m);
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(sumUb, vreg_sum, preg_m);
}
template <typename T>
__aicore__ inline void DuplicateSumWithR0(const LocalTensor<T>& sumTensor, const T R0, uint32_t m)
{
__ubuf__ T * sumUb = (__ubuf__ T*)sumTensor.GetPhyAddr();
DuplicateSumWithR0VF<T>(sumUb, R0, m);
}
} // namespace
#endif // MUL_SEL_SOFTMAX_FLASH_V2_CAST_NZ_SFA_INTERFACE_H

View File

@@ -0,0 +1,292 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file buffer.h
* \brief同步管理
*/
#ifndef BUFFER_H
#define BUFFER_H
#include<type_traits>
#include"lib/matmul_intf.h"
#if ASC_DEVKIT_MAJOR >= 9
#include "kernel_basic_intf.h"
#else
#include "kernel_operator.h"
#endif
using namespace AscendC;
namespace fa_base_matmul {
__BLOCK_LOCAL__ __inline__ uint32_t idCounterNum;
#define MAKE_ID ((++idCounterNum) % 11)
// 核间同步中,AIC(flagId 0-10)对应AIV0(flagId 0-10),对应AIV1(flagId 16-26)
#define AIV0_AIV1_OFFSET 16
enum class BufferType {
L1 = 0,
L0A = 1,
L0B = 2,
L0C = 3,
UB = 4,
GM = 5,
C2 = 6,
};
enum class SyncType {
NO_SYNC,
INNER_CORE_SYNC,
CROSS_CORE_SYNC_FORWARD,
CROSS_CORE_SYNC_BOTH,
CROSS_CORE_SYNC_BACKWARD,
};
constexpr uint32_t INVALID_CROSS_CORE_EVENT_ID = 16;
static constexpr uint64_t CROSS_CORE_SYNC_MODE = 4;
template<BufferType Type>
struct BufferInfo{
// Cons 消费者,Prod 生产者
__aicore__ const static constexpr HardEvent ConsWaitProdStatus() {
if constexpr (Type == BufferType::L1) {
return HardEvent::MTE2_MTE1;
} else if constexpr (Type == BufferType::L0A) {
return HardEvent::MTE1_M;
} else if constexpr (Type == BufferType::L0B) {
return HardEvent::MTE1_M;
} else if constexpr (Type == BufferType::L0C) {
return HardEvent::M_FIX;
} else if constexpr (Type == BufferType::C2) {
return HardEvent::MTE1_M;
}
}
__aicore__ const static constexpr HardEvent ProdWaitConsStatus() {
if constexpr (Type == BufferType::L1) {
return HardEvent::MTE1_MTE2;
} else if constexpr (Type == BufferType::L0A) {
return HardEvent::M_MTE1;
} else if constexpr (Type == BufferType::L0B) {
return HardEvent::M_MTE1;
} else if constexpr (Type == BufferType::L0C) {
return HardEvent::FIX_M;
} else if constexpr (Type == BufferType::C2) {
return HardEvent::M_MTE1;
}
}
__aicore__ const static constexpr TPosition GetTPosition() {
if constexpr (Type == BufferType::L1) {
return TPosition::A1;
} else if constexpr (Type == BufferType::L0A) {
return TPosition::A2;
} else if constexpr (Type == BufferType::L0B) {
return TPosition::B2;
} else if constexpr (Type == BufferType::L0C) {
return TPosition::CO1;
} else if constexpr (Type == BufferType::UB) {
return TPosition::VECIN;
} else if constexpr (Type == BufferType::GM) {
return TPosition::GM;
} else if constexpr (Type == BufferType::C2) {
return TPosition::C2;
}
}
static constexpr HardEvent EventP2C = ConsWaitProdStatus(); // 生产者到消费者方向的HardEvent:消费者等生产者提供/生产者通知消费者已生成
static constexpr HardEvent EventC2P = ProdWaitConsStatus(); // 消费者到生产者方向的HardEvent:生产者等消费者消耗/消费者通知生产者已消耗’
static constexpr TPosition Position = GetTPosition();
};
// buffer绑定生产者、消费者关系
// L1 buffer的生产者为MTE2或者MTE3,消费者为MTE1
// L0A buffer的生产者为MTE1,消费者为M
// L0B buffer的生产者为MTE1,消费者为M
// L0C buffer的生产者为M,消费者为FIX
template<BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
class Buffer {
using TensorType = std::conditional_t<bufferType == BufferType::GM, GlobalTensor<uint8_t>, LocalTensor<uint8_t>>;
template <typename T>
using TargetTensorType = std::conditional_t<bufferType == BufferType::GM, GlobalTensor<T>, LocalTensor<T>>;
public:
__aicore__ inline Buffer() {}
__aicore__ inline Buffer(TensorType tensor, uint32_t size) {
tensor_ = tensor;
size_ = size;
if constexpr (syncType == SyncType::CROSS_CORE_SYNC_FORWARD) {
id0_ = MAKE_ID;
id1_ = INVALID_CROSS_CORE_EVENT_ID;
} else if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BACKWARD) {
id0_ = INVALID_CROSS_CORE_EVENT_ID;
id1_ = MAKE_ID;
} else if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BOTH) {
id0_ = MAKE_ID;
id1_ = MAKE_ID;
} else {
id0_ = INVALID_CROSS_CORE_EVENT_ID;
id1_ = INVALID_CROSS_CORE_EVENT_ID;
}
}
__aicore__ inline void Init() {
if ASCEND_IS_AIC {
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
p2cEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventP2C>(); // 确保只能被调用一次
c2pEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventC2P>();
SetFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_);
}
}
}
__aicore__ inline void UnInit() {
if ASCEND_IS_AIC {
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
WaitFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_);
GetTPipePtr()->ReleaseEventID<BufferInfo<bufferType>::EventP2C>(p2cEventId_); // 确保只能被调用一次
GetTPipePtr()->ReleaseEventID<BufferInfo<bufferType>::EventC2P>(c2pEventId_);
}
}
}
template<HardEvent EventType>
__aicore__ inline void Wait() {
if ASCEND_IS_AIC {
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
if constexpr (EventType == BufferInfo<bufferType>::EventP2C) {
WaitFlag<BufferInfo<bufferType>::EventP2C>(p2cEventId_); // 消费者等待生产者完成生产
} else {
WaitFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_); // 生产者等待消费者完成消费
}
}
}
}
template<HardEvent EventType>
__aicore__ inline void Set() {
if ASCEND_IS_AIC {
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
if constexpr (EventType == BufferInfo<bufferType>::EventP2C) {
SetFlag<BufferInfo<bufferType>::EventP2C>(p2cEventId_); // 生产者通知消费者已完成生产
} else {
SetFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_); // 消费者通知生产者已完成消费
}
}
}
}
__aicore__ inline void SetEventID() {
if ASCEND_IS_AIC {
p2cEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventP2C>(); // 确保只能被调用一次
c2pEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventC2P>();
}
}
template<HardEvent EventType>
__aicore__ inline TEventID GetEventID() {
if ASCEND_IS_AIC {
if constexpr (EventType == BufferInfo<bufferType>::EventP2C) {
return p2cEventId_; // 生产者通知消费者已完成生产
} else {
return c2pEventId_; // 消费者通知生产者已完成消费
}
}
}
template<bool isReuse = false>
__aicore__ inline void WaitCrossCore() {
if constexpr (bufferType == BufferType::GM && syncType == SyncType::CROSS_CORE_SYNC_BACKWARD) {
// AIC属于消费者,AIV属于生产者,且一个AIC对应两个AIV
if ASCEND_IS_AIC {
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE2>(id1_);
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE2>(id1_ + AIV0_AIV1_OFFSET);
} else {
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE2>(id0_);
}
} else if constexpr (bufferType == BufferType::UB || bufferType == BufferType::GM) {
// AIC属于生产者,AIV属于消费者,且一个AIC对应两个AIV
if ASCEND_IS_AIC {
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_FIX>(id1_);
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_FIX>(id1_ + AIV0_AIV1_OFFSET);
} else {
if constexpr (isReuse) {
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE3>(id0_);
} else {
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_V>(id0_);
}
}
} else if constexpr (bufferType == BufferType::L1) {
// AIC属于消费者,AIV属于生产者,且一个AIC对应两个AIV
if ASCEND_IS_AIC {
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE1>(id0_);
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE1>(id0_ + AIV0_AIV1_OFFSET);
} else {
if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BOTH) {
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE3>(id1_);
}
}
}
}
template<bool isReuse = false>
__aicore__ inline void SetCrossCore() {
if constexpr (bufferType == BufferType::GM && syncType == SyncType::CROSS_CORE_SYNC_BACKWARD) {
// AIC属于消费者,AIV属于生产者,且一个AIC对应两个AIV
if ASCEND_IS_AIC {
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_FIX>(id0_);
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_FIX>(id0_ + AIV0_AIV1_OFFSET);
} else {
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE3>(id1_);
}
} else if constexpr (bufferType == BufferType::UB || bufferType == BufferType::GM) {
// AIC属于生产者,AIV属于消费者,且一个AIC对应两个AIV
if ASCEND_IS_AIC {
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_FIX>(id0_);
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_FIX>(id0_ + AIV0_AIV1_OFFSET);
} else {
if constexpr (isReuse) {
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE3>(id1_);
} else {
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_V>(id1_);
}
}
} else if constexpr (bufferType == BufferType::L1) {
// AIC属于消费者,AIV属于生产者,且一个AIC对应两个AIV
if ASCEND_IS_AIC {
if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BOTH) {
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE1>(id1_);
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE1>(id1_ + AIV0_AIV1_OFFSET);
}
} else {
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE3>(id0_);
}
}
}
template<typename T>
__aicore__ inline TargetTensorType<T> GetTensor() {
return tensor_.template ReinterpretCast<T>();
}
template<typename T>
__aicore__ inline TargetTensorType<T> GetTensor(uint64_t startindex) {
TargetTensorType<T> tmpTensor = tensor_.template ReinterpretCast<T>();
return tmpTensor[startindex];
}
private:
TensorType tensor_;
uint32_t size_;
TEventID p2cEventId_;
TEventID c2pEventId_;
uint32_t id0_; // 用作正向同步:生产者通知消费者,或者消费者等待生产者;
uint32_t id1_; // 用作反向同步:消费者通知生产者,或者生产者等待消费者;
};
}
#endif

View File

@@ -0,0 +1,61 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file buffer_manager.h
* \brief buffer内存管理
*/
#ifndef BUFFER_MANAGER_H
#define BUFFER_MANAGER_H
#if (__NPU_ARCH__ == 5102)
#include "buffer_mix_core.h"
#else
#include "buffer.h"
#endif
// L1 TPosition::A1
// L0A TPosition::A2
// L0B TPosition::B2
// L0C TPosition::CO1
// UB TPosition::VECIN
namespace fa_base_matmul {
template<BufferType bufferType>
class BufferManager {
using TensorType = std::conditional_t<bufferType == BufferType::GM, GlobalTensor<uint8_t>, LocalTensor<uint8_t>>;
public:
__aicore__ inline void Init(TPipe *pipe, uint32_t size) {
static_assert(bufferType != BufferType::GM, "GM should use workspace.");
TBuf<BufferInfo<bufferType>::Position> tbuf;
pipe->InitBuffer(tbuf, size);
mem_ = tbuf.template Get<uint8_t>();
}
__aicore__ inline void Init(__gm__ uint8_t* workspace) {
static_assert(bufferType == BufferType::GM, "BufferType should be GM.");
mem_.SetGlobalBuffer((__gm__ uint8_t*)workspace);
}
template<SyncType syncType = SyncType::INNER_CORE_SYNC>
__aicore__ inline Buffer<bufferType, syncType> AllocBuffer(uint32_t size) {
TensorType temp = mem_[offset_];
offset_ += size;
return Buffer<bufferType, syncType>(temp, size);
}
template<SyncType syncType = SyncType::INNER_CORE_SYNC>
__aicore__ inline void FreeBuffer(Buffer<bufferType, syncType> &buffer){
}
private:
uint32_t offset_ = 0;
TensorType mem_;
};
}
#endif

View File

@@ -0,0 +1,216 @@
/**
 * 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 buffer_mix_core.h
* \brief同步管理
*/
#ifndef BUFFER_MIX_CORE_H
#define BUFFER_MIX_CORE_H
#include <type_traits>
#include "lib/matmul_intf.h"
#if ASC_DEVKIT_MAJOR >= 9
#include "kernel_basic_intf.h"
#else
#include "kernel_operator.h"
#endif
using namespace AscendC;
namespace fa_base_matmul {
__BLOCK_LOCAL__ __inline__ uint32_t idCounterNum;
#define MAKE_ID ((++idCounterNum) % 11)
// 核间同步中,AIC(flagId 0-10)对应AIV0(flagId 0-10),对应AIV1(flagId 16-26)
#define AIV0_AIV1_OFFSET 16
enum class BufferType {
L1 = 0,
L0A = 1,
L0B = 2,
L0C = 3,
UB = 4,
GM = 5,
};
enum class SyncType {
NO_SYNC,
INNER_CORE_SYNC,
CROSS_CORE_SYNC_FORWARD,
CROSS_CORE_SYNC_BOTH,
};
constexpr uint32_t INVALID_CROSS_CORE_EVENT_ID = 16;
static constexpr uint64_t CROSS_CORE_SYNC_MODE = 4;
template <BufferType Type>
struct BufferInfo {
// Cons 消费者,Prod 生产者
__aicore__ const static constexpr HardEvent ConsWaitProdStatus()
{
if constexpr (Type == BufferType::L1) {
return HardEvent::MTE2_MTE1;
} else if constexpr (Type == BufferType::L0A) {
return HardEvent::MTE1_M;
} else if constexpr (Type == BufferType::L0B) {
return HardEvent::MTE1_M;
} else if constexpr (Type == BufferType::L0C) {
return HardEvent::M_FIX;
}
}
__aicore__ const static constexpr HardEvent ProdWaitConsStatus()
{
if constexpr (Type == BufferType::L1) {
return HardEvent::MTE1_MTE2;
} else if constexpr (Type == BufferType::L0A) {
return HardEvent::M_MTE1;
} else if constexpr (Type == BufferType::L0B) {
return HardEvent::M_MTE1;
} else if constexpr (Type == BufferType::L0C) {
return HardEvent::FIX_M;
}
}
__aicore__ const static constexpr TPosition GetTPosition()
{
if constexpr (Type == BufferType::L1) {
return TPosition::A1;
} else if constexpr (Type == BufferType::L0A) {
return TPosition::A2;
} else if constexpr (Type == BufferType::L0B) {
return TPosition::B2;
} else if constexpr (Type == BufferType::L0C) {
return TPosition::CO1;
} else if constexpr (Type == BufferType::UB) {
return TPosition::VECIN;
} else if constexpr (Type == BufferType::GM) {
return TPosition::GM;
}
}
static constexpr HardEvent EventP2C =
ConsWaitProdStatus(); // 生产者到消费者方向的HardEvent:消费者等生产者提供/生产者通知消费者已生成
static constexpr HardEvent EventC2P =
ProdWaitConsStatus(); // 消费者到生产者方向的HardEvent:生产者等消费者消耗/消费者通知生产者已消耗’
static constexpr TPosition Position = GetTPosition();
};
// buffer绑定生产者、消费者关系
// L1 buffer的生产者为MTE2或者MTE3,消费者为MTE1
// L0A buffer的生产者为MTE1,消费者为M
// L0B buffer的生产者为MTE1,消费者为M
// L0C buffer的生产者为M,消费者为FIX
template <BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
class Buffer {
using TensorType = std::conditional_t<bufferType == BufferType::GM, GlobalTensor<uint8_t>, LocalTensor<uint8_t>>;
template <typename T>
using TargetTensorType = std::conditional_t<bufferType == BufferType::GM, GlobalTensor<T>, LocalTensor<T>>;
public:
__aicore__ inline Buffer()
{
}
__aicore__ inline Buffer(TensorType tensor, uint32_t size)
{
tensor_ = tensor;
size_ = size;
if constexpr (syncType == SyncType::CROSS_CORE_SYNC_FORWARD) {
id0_ = MAKE_ID;
id1_ = INVALID_CROSS_CORE_EVENT_ID;
} else if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BOTH) {
id0_ = MAKE_ID;
id1_ = MAKE_ID;
} else {
id0_ = INVALID_CROSS_CORE_EVENT_ID;
id1_ = INVALID_CROSS_CORE_EVENT_ID;
}
}
__aicore__ inline void Init()
{
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
p2cEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventP2C>(); // 确保只能被调用一次
c2pEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventC2P>();
SetFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_);
}
}
__aicore__ inline void UnInit()
{
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
WaitFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_);
GetTPipePtr()->ReleaseEventID<BufferInfo<bufferType>::EventP2C>(p2cEventId_); // 确保只能被调用一次
GetTPipePtr()->ReleaseEventID<BufferInfo<bufferType>::EventC2P>(c2pEventId_);
}
}
template <HardEvent EventType>
__aicore__ inline void Wait()
{
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
if constexpr (EventType == BufferInfo<bufferType>::EventP2C) {
WaitFlag<BufferInfo<bufferType>::EventP2C>(p2cEventId_); // 消费者等待生产者完成生产
} else {
WaitFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_); // 生产者等待消费者完成消费
}
}
}
template <HardEvent EventType>
__aicore__ inline void Set()
{
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
if constexpr (EventType == BufferInfo<bufferType>::EventP2C) {
SetFlag<BufferInfo<bufferType>::EventP2C>(p2cEventId_); // 生产者通知消费者已完成生产
} else {
SetFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_); // 消费者通知生产者已完成消费
}
}
}
__aicore__ inline void SetEventID()
{
p2cEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventP2C>(); // 确保只能被调用一次
c2pEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventC2P>();
}
template <HardEvent EventType>
__aicore__ inline TEventID GetEventID()
{
if constexpr (EventType == BufferInfo<bufferType>::EventP2C) {
return p2cEventId_; // 生产者通知消费者已完成生产
} else {
return c2pEventId_; // 消费者通知生产者已完成消费
}
}
template <typename T>
__aicore__ inline TargetTensorType<T> GetTensor()
{
return tensor_.template ReinterpretCast<T>();
}
template <typename T>
__aicore__ inline TargetTensorType<T> GetTensor(uint64_t startindex)
{
TargetTensorType<T> tmpTensor = tensor_.template ReinterpretCast<T>();
return tmpTensor[startindex];
}
private:
TensorType tensor_;
uint32_t size_;
TEventID p2cEventId_;
TEventID c2pEventId_;
uint32_t id0_; // 用作正向同步:生产者通知消费者,或者消费者等待生产者;
uint32_t id1_; // 用作反向同步:消费者通知生产者,或者生产者等待消费者;
};
} // namespace fa_base_matmul
#endif

View File

@@ -0,0 +1,407 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file buffers_policy.h
* \brief 综合管理buffer的内存和同步
*/
#ifndef BUFFERS_POLICY_H
#define BUFFERS_POLICY_H
#include "buffer_manager.h"
#define NUM_2 2
#define NUM_3 3
#define NUM_4 4
// Q复用 KV复用
// 申请单块buffer
namespace fa_base_matmul {
template<BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
class BuffersPolicySingleBuffer {
public:
__aicore__ inline void Init(BufferManager<bufferType> &bufferManager, uint32_t size){
buffer_ = bufferManager.template AllocBuffer<syncType>(size);
buffer_.Init();
}
__aicore__ inline void Uninit(BufferManager<bufferType> &bufferManager){
buffer_.UnInit();
bufferManager.FreeBuffer(buffer_);
}
__aicore__ inline Buffer<bufferType, syncType> &Get(){
return buffer_;
}
__aicore__ inline Buffer<bufferType, syncType> &GetPre(){
return Get();
}
__aicore__ inline Buffer<bufferType, syncType> &GetReused(){
return Get();
}
private:
Buffer<bufferType, syncType> buffer_;
};
// 申请2个buffer,乒乓轮转
template<BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
class BuffersPolicyDB {
public:
__aicore__ inline void Init(BufferManager<bufferType> &bufferManager, uint32_t size){
ping_ = bufferManager.template AllocBuffer<syncType>(size);
pong_ = bufferManager.template AllocBuffer<syncType>(size);
ping_.Init();
pong_.Init();
}
__aicore__ inline void Uninit(BufferManager<bufferType> &bufferManager){
ping_.UnInit();
pong_.UnInit();
bufferManager.FreeBuffer(ping_);
bufferManager.FreeBuffer(pong_);
}
__aicore__ inline Buffer<bufferType, syncType> &Get() {
if (flag1_) { // 1
flag1_ = 0;
return ping_;
} else { // 0
flag1_ = 1;
return pong_;
}
}
// 需要与Get联用, 首次调用Get,第二次调用GetPre(Q复用)
__aicore__ inline Buffer<bufferType, syncType> &GetPre() {
if (flag1_) { // 0->1
return pong_;
} else { // 1->0
return ping_;
}
}
// 需要与Get,GetPre联用, 首次调用Get,第二次调用GetPre,第三次复用时GetReused(KV复用)
__aicore__ inline Buffer<bufferType, syncType> &GetReused() {
if (flag2_ == 0) {
flag2_ = 1;
return pong_;
} else {
flag2_ = 0;
return ping_;
}
}
__aicore__ inline Buffer<bufferType, syncType> &GetReused(bool isNextS2IdxNoChange) {
if (isNextS2IdxNoChange) {
if (flag2_ == 0) {
return pong_;
} else {
return ping_;
}
} else {
return GetReused();
}
}
private:
Buffer<bufferType, syncType> ping_;
Buffer<bufferType, syncType> pong_;
uint32_t flag1_ = 0;
uint32_t flag2_ = 0;
};
// 申请3个buffer, 轮转
template<BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
class BuffersPolicy3buff {
public:
__aicore__ inline void Init(BufferManager<bufferType> &bufferManager, uint32_t size) {
a_ = bufferManager.template AllocBuffer<syncType>(size);
b_ = bufferManager.template AllocBuffer<syncType>(size);
c_ = bufferManager.template AllocBuffer<syncType>(size);
a_.Init();
b_.Init();
c_.Init();
}
__aicore__ inline void Uninit(BufferManager<bufferType> &bufferManager) {
a_.UnInit();
b_.UnInit();
c_.UnInit();
bufferManager.FreeBuffer(a_);
bufferManager.FreeBuffer(b_);
bufferManager.FreeBuffer(c_);
}
__aicore__ inline Buffer<bufferType, syncType> &Get() {
if (flag1_ == 0) {
flag1_ = 1;
return a_;
} else if (flag1_ == 1) {
flag1_ = NUM_2;
return b_;
} else {
flag1_ = 0;
return c_;
}
}
__aicore__ inline Buffer<bufferType, syncType> &GetVec() { // mixcore architecture
if (flag1_vec1_ == 0) {
flag1_vec1_ = 1;
return a_;
} else if (flag1_vec1_ == 1) {
flag1_vec1_ = NUM_2;
return b_;
} else {
flag1_vec1_ = 0;
return c_;
}
}
__aicore__ inline Buffer<bufferType, syncType> &GetCube() { // mixcore architecture
if (flag1_bmm2_ == 0) {
flag1_bmm2_ = 1;
return a_;
} else if (flag1_bmm2_ == 1) {
flag1_bmm2_ = NUM_2;
return b_;
} else {
flag1_bmm2_ = 0;
return c_;
}
}
// Q复用
__aicore__ inline Buffer<bufferType, syncType> &GetPre() {
if (flag1_ == 0) {
return c_;
} else if (flag1_ == 1) {
return a_;
} else {
return b_;
}
}
// KV复用
__aicore__ inline Buffer<bufferType, syncType> &GetReused() {
if (flag2_ == 0) {
flag2_ = 1;
return a_;
} else if (flag2_ == 1){
flag2_ = NUM_2;
return b_;
} else {
flag2_ = 0;
return c_;
}
}
private:
Buffer<bufferType, syncType> a_;
Buffer<bufferType, syncType> b_;
Buffer<bufferType, syncType> c_;
uint32_t flag1_ = 0;
uint32_t flag1_vec1_ = 0;
uint32_t flag1_bmm2_ = 0;
uint32_t flag2_ = 0;
};
// 申请4个buffer + kv复用
template<BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
class BuffersPolicy4buff {
public:
__aicore__ inline void Init(BufferManager<bufferType> &bufferManager, uint32_t size) {
a_ = bufferManager.template AllocBuffer<syncType>(size);
b_ = bufferManager.template AllocBuffer<syncType>(size);
c_ = bufferManager.template AllocBuffer<syncType>(size);
d_ = bufferManager.template AllocBuffer<syncType>(size);
a_.Init();
b_.Init();
c_.Init();
d_.Init();
}
__aicore__ inline void Uninit(BufferManager<bufferType> &bufferManager) {
a_.UnInit();
b_.UnInit();
c_.UnInit();
d_.UnInit();
bufferManager.FreeBuffer(a_);
bufferManager.FreeBuffer(b_);
bufferManager.FreeBuffer(c_);
bufferManager.FreeBuffer(d_);
}
__aicore__ inline Buffer<bufferType, syncType> &Get(uint32_t id) {
uint32_t flag = id % 4;
if (flag == 0) {
return a_;
} else if (flag == 1) {
return b_;
} else if (flag == 2) { // 2:c_
return c_;
} else {
return d_;
}
}
__aicore__ inline Buffer<bufferType, syncType> &Get() {
auto& buffer = Get(head_);
head_++;
return buffer;
}
__aicore__ inline Buffer<bufferType, syncType> &GetReused() {
auto& buffer = Get(used_);
used_ = (used_ - tail_ + 1) % (head_ - tail_) + tail_;
return buffer;
}
__aicore__ inline Buffer<bufferType, syncType> &GetFree() {
if (tail_ == used_) {
used_++;
}
auto& buffer = Get(tail_);
tail_++;
return buffer;
}
private:
Buffer<bufferType, syncType> a_;
Buffer<bufferType, syncType> b_;
Buffer<bufferType, syncType> c_;
Buffer<bufferType, syncType> d_;
uint32_t tail_ = 0; // 表示当前正在使用的buffer队列队尾
uint32_t head_ = 0; // 表示当前正在使用的buffer队列队首+1
uint32_t used_ = 0; // 表示当前正在使用的buffer,于首尾间,左闭右开
};
template<BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
class Matrix2x2BufferPolicy { // 4buffer
// 二维buffer管理,地址行优先,使用列优先
// MracBuffer:memory address with row first, alloc/use/free with column first
public:
__aicore__ inline void Init(BufferManager<bufferType> &bufferManager, uint32_t size) {
bufferM0k0_ = bufferManager.template AllocBuffer<syncType>(size);
bufferM0k1_ = bufferManager.template AllocBuffer<syncType>(size);
bufferM1k0_ = bufferManager.template AllocBuffer<syncType>(size);
bufferM1k1_ = bufferManager.template AllocBuffer<syncType>(size);
bufferM0k0_.Init();
bufferM0k1_.Init();
bufferM1k0_.Init();
bufferM1k1_.Init();
}
__aicore__ inline void Uninit(BufferManager<bufferType> &bufferManager) {
bufferM0k0_.UnInit();
bufferM0k1_.UnInit();
bufferM1k0_.UnInit();
bufferM1k1_.UnInit();
bufferManager.FreeBuffer(bufferM0k0_);
bufferManager.FreeBuffer(bufferM0k1_);
bufferManager.FreeBuffer(bufferM1k0_);
bufferManager.FreeBuffer(bufferM1k1_);
}
__aicore__ inline void SetMExtent(int32_t mExtent) {
aIdx_ = -1;
amIdx_ = (amIdx_ + mSize_ - 1) % mSize_; // 翻转 0->1, 1->0
akIdx_ = 0;
uIdx_ = -1;
umIdx_ = (umIdx_ + mSize_ - 1) % mSize_;
ukIdx_ = 0;
fIdx_ = -1;
fmIdx_ = (fmIdx_ + mSize_ - 1) % mSize_;
fkIdx_ = 0;
mExtent_ = mExtent;
}
__aicore__ inline Buffer<bufferType, syncType> &AllocNext() {
aIdx_++;
return GetBuffer(aIdx_, amIdx_, akIdx_);
}
__aicore__ inline Buffer<bufferType, syncType> &ReuseNext() {
uIdx_++;
return GetBuffer(uIdx_, umIdx_, ukIdx_);
}
__aicore__ inline Buffer<bufferType, syncType> &FreeNext() {
fIdx_++;
return GetBuffer(fIdx_, fmIdx_, fkIdx_);
}
__aicore__ inline Buffer<bufferType, syncType> &PeekNextK() { // 在Alloc阶段使用,k方向取下一个
return PeekBuffer(amIdx_, (1 - akIdx_)); // k翻转
}
private:
__aicore__ inline Buffer<bufferType, syncType> &GetBuffer(int32_t xIdx, int32_t &mIdx, int32_t &kIdx) {
// xIdx为入参,表示当前alloc/use/free的idx,mIdx和kIdx为下标出参,移动到下一个buffer并获取
mIdx = (mIdx + mExtent_ - 1) % mExtent_;
kIdx = (xIdx / mExtent_) % kSize_;
if (mIdx == 0 && kIdx == 0) {
return bufferM0k0_;
} else if (mIdx == 0 && kIdx == 1) {
return bufferM0k1_;
} else if (mIdx == 1 && kIdx == 0) {
return bufferM1k0_;
} else { // 该分支条件为:mIdx == 1 && kIdx == 1
return bufferM1k1_;
}
}
__aicore__ inline Buffer<bufferType, syncType> &PeekBuffer(int32_t mIdx, int32_t kIdx) {
// 只访问buffer,不进行下标移动
if (mIdx == 0 && kIdx == 0) {
return bufferM0k0_;
} else if (mIdx == 0 && kIdx == 1) {
return bufferM0k1_;
} else if ((mIdx == 1) && (kIdx == 0)) {
return bufferM1k0_;
} else { // mIdx == 1 && kIdx == 1
return bufferM1k1_;
}
}
Buffer<bufferType, syncType> bufferM0k0_;
Buffer<bufferType, syncType> bufferM0k1_;
Buffer<bufferType, syncType> bufferM1k0_;
Buffer<bufferType, syncType> bufferM1k1_;
int32_t mSize_ = 2; // m的总buffer数
int32_t kSize_ = 2; // k的总buffer数
// Alloc
int32_t aIdx_ = -1; // 当前第几次Alloc Buffer
int32_t amIdx_ = 0; // 当前Alloc Buffer的m下标
int32_t akIdx_ = 0; // 当前Alloc Buffer的k下标
// Reuse
int32_t uIdx_ = -1;
int32_t umIdx_ = 0;
int32_t ukIdx_ = 0;
// Free
int32_t fIdx_ = -1;
int32_t fmIdx_ = 0;
int32_t fkIdx_ = 0;
int32_t mExtent_ = 0; // m实际使用的大小,可以为1或者2
};
}
#endif

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,35 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file fa_gm_tensor.h
* \brief
*/
#ifndef FA_GM_TENSOR_H
#define FA_GM_TENSOR_H
#if ASC_DEVKIT_MAJOR >= 9
#include "kernel_vec_intf.h"
#include "kernel_cube_intf.h"
#else
#include "kernel_operator.h"
#endif
#include "gm_layout.h"
#include "offset_calculator_v2.h"
using AscendC::GlobalTensor;
template <typename Q_T, GmFormat FORMAT, typename ACTLEN_T = uint64_t>
struct FaGmTensor {
GlobalTensor<Q_T> gmTensor;
OffsetCalculator<FORMAT, ACTLEN_T> offsetCalculator;
};
#endif

View File

@@ -0,0 +1,43 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file fa_l1_tensor.h
* \brief
*/
#ifndef FA_L1_TENSOR_H
#define FA_L1_TENSOR_H
#if ASC_DEVKIT_MAJOR >= 9
#include "kernel_vec_intf.h"
#include "kernel_cube_intf.h"
#else
#include "kernel_operator.h"
#endif
using AscendC::LocalTensor;
enum class L1Format {
NZ = 0
};
enum class ScaleTrans {
NO_TRANS = 0,
ND2NZ = 1,
DN2NZ = 2
};
template <typename Q_T, L1Format FORMAT>
struct FaL1Tensor {
LocalTensor<Q_T> tensor;
uint32_t rowCount;
};
#endif

View File

@@ -0,0 +1,26 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file gm_coord.h
* \brief
*/
#ifndef GM_COORD_H
#define GM_COORD_H
struct GmCoord {
uint32_t bIdx;
uint32_t n2Idx;
uint32_t gS1Idx;
uint32_t dIdx;
uint32_t gS1DealSize;
uint32_t dDealSize;
};
#endif

View File

@@ -0,0 +1,427 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file gm_layout.h
* \brief
*/
#ifndef GM_LAYOUT_H
#define GM_LAYOUT_H
#if ASC_DEVKIT_MAJOR >= 9
#include "kernel_vec_intf.h"
#include "kernel_cube_intf.h"
#else
#include "kernel_operator.h"
#endif
// ----------------------------------------------GmLayout--------------------------------
enum class GmFormat {
BSNGD = 0,
BNGSD = 1,
NGBSD = 2,
TNGD = 3,
NGTD = 4,
BSND = 5,
BNSD = 6,
TND = 7,
NTD = 8,
PA_BnBsND = 9,
PA_BnNBsD = 10,
PA_NZ = 11,
NGD = 12, // post_quant
ND = 13, //antiquant no PA
BS2 = 14,
BNS2 = 15,
PA_BnBs = 16, //antiquant PA
PA_BnNBs = 17,
BN2GS1S2 = 18, //PSE_GmFormat
SBNGD = 19,
SBND = 20,
NTGD = 21,
PA_NZ_K_SCALE = 22,
};
template <GmFormat FORMAT>
struct GmLayout {
};
template <>
struct GmLayout<GmFormat::BSNGD> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t g, uint32_t s, uint32_t d) {
shape = AscendC::MakeShape(b, n, g, s, d);
uint64_t dStride = 1;
uint64_t gStride = dStride * d;
uint64_t nStride = gStride * g;
uint64_t sStride = nStride * n;
uint64_t bStride = sStride * s;
stride = AscendC::MakeStride(bStride, nStride, gStride, sStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::BNGSD> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t g, uint32_t s, uint32_t d) {
shape = AscendC::MakeShape(b, n, g, s, d);
uint64_t dStride = 1;
uint64_t sStride = dStride * d;
uint64_t gStride = sStride * s;
uint64_t nStride = gStride * g;
uint64_t bStride = nStride * n;
stride = AscendC::MakeStride(bStride, nStride, gStride, sStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::NGBSD> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t g, uint32_t s, uint32_t d) {
shape = AscendC::MakeShape(b, n, g, s, d);
uint64_t dStride = 1;
uint64_t sStride = dStride * d;
uint64_t bStride = sStride * s;
uint64_t gStride = bStride * b;
uint64_t nStride = gStride * g;
stride = AscendC::MakeStride(bStride, nStride, gStride, sStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::TNGD> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t t, uint32_t n, uint32_t g, uint32_t d) {
shape = AscendC::MakeShape(t, n, g, d);
uint64_t dStride = 1;
uint64_t gStride = dStride * d;
uint64_t nStride = gStride * g;
uint64_t tStride = nStride * n;
stride = AscendC::MakeStride(tStride, nStride, gStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::NGTD> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t t, uint32_t n, uint32_t g, uint32_t d) {
shape = AscendC::MakeShape(t, n, g, d);
uint64_t dStride = 1;
uint64_t tStride = dStride * d;
uint64_t gStride = tStride * t;
uint64_t nStride = gStride * g;
stride = AscendC::MakeStride(tStride, nStride, gStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::NTGD> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t t, uint32_t n, uint32_t g, uint32_t d) {
shape = AscendC::MakeShape(t, n, g, d);
uint64_t dStride = 1;
uint64_t gStride = dStride * d;
uint64_t tStride = gStride * g;
uint64_t nStride = tStride * t;
stride = AscendC::MakeStride(tStride, nStride, gStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::BSND> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t s, uint32_t d) {
shape = AscendC::MakeShape(b, n, s, d);
uint64_t dStride = 1;
uint64_t nStride = dStride * d;
uint64_t sStride = nStride * n;
uint64_t bStride = sStride * s;
stride = AscendC::MakeStride(bStride, nStride, sStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::BNSD> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t s, uint32_t d) {
shape = AscendC::MakeShape(b, n, s, d);
uint64_t dStride = 1;
uint64_t sStride = dStride * d;
uint64_t nStride = sStride * s;
uint64_t bStride = nStride * n;
stride = AscendC::MakeStride(bStride, nStride, sStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::TND> {
AscendC::Shape<uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t t, uint32_t n, uint32_t d) {
shape = AscendC::MakeShape(t, n, d);
uint64_t dStride = 1;
uint64_t nStride = dStride * d;
uint64_t tStride = nStride * n;
stride = AscendC::MakeStride(tStride, nStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::NTD> {
AscendC::Shape<uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t t, uint32_t n, uint32_t d) {
shape = AscendC::MakeShape(t, n, d);
uint64_t dStride = 1;
uint64_t tStride = dStride * d;
uint64_t nStride = tStride * t;
stride = AscendC::MakeStride(tStride, nStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::PA_BnBsND> {
AscendC::Shape<uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t n, uint32_t blockSize, uint32_t d) {
shape = AscendC::MakeShape(n, blockSize, d);
uint64_t dStride = 1;
uint64_t nStride = dStride * d;
uint64_t bsStride = nStride * n;
uint64_t bnStride = bsStride * blockSize;
stride = AscendC::MakeStride(bnStride, nStride, bsStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::PA_BnNBsD> {
AscendC::Shape<uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t n, uint32_t blockSize, uint32_t d) {
shape = AscendC::MakeShape(n, blockSize, d);
uint64_t dStride = 1;
uint64_t bsStride = dStride * d;
uint64_t nStride = bsStride * blockSize;
uint64_t bnStride = nStride * n;
stride = AscendC::MakeStride(bnStride, nStride, bsStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::PA_NZ> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t n, uint32_t blockSize, uint32_t d1, uint32_t d0) {
shape = AscendC::MakeShape(n, d1, blockSize, d0);
uint64_t d0Stride = 1;
uint64_t bsStride = d0Stride * d0;
uint64_t d1Stride = bsStride * blockSize;
uint64_t nStride = d1Stride * d1;
uint64_t bnStride = nStride * n;
stride = AscendC::MakeStride(bnStride, nStride, d1Stride, bsStride, d0Stride);
}
};
template <>
struct GmLayout<GmFormat::PA_NZ_K_SCALE> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t n, uint32_t blockSize1, uint32_t d, uint32_t blockSize0) {
shape = AscendC::MakeShape(n, blockSize1, d, blockSize0);
uint64_t bs0Stride = 1;
uint64_t dStride = bs0Stride * blockSize0;
uint64_t bs1Stride = dStride * d;
uint64_t nStride = bs1Stride * blockSize1;
uint64_t bnStride = nStride * n;
stride = AscendC::MakeStride(bnStride, nStride, bs1Stride, dStride, bs0Stride);
}
};
// post_quant
template <>
struct GmLayout<GmFormat::NGD> {
AscendC::Shape<uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t n, uint32_t g, uint32_t d) {
shape = AscendC::MakeShape(n, g, d);
uint64_t dStride = 1;
uint64_t gStride = dStride * d;
uint64_t nStride = gStride * g;
stride = AscendC::MakeStride(nStride, gStride, dStride);
}
};
//antiquant
template <>
struct GmLayout<GmFormat::ND> {
AscendC::Shape<uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t n, uint32_t d) {
shape = AscendC::MakeShape(n, d);
uint64_t dStride = 1;
uint64_t nStride = dStride * d; //headDim
stride = AscendC::MakeStride(nStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::BS2> {
AscendC::Shape<uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t b, uint32_t s) {
shape = AscendC::MakeShape(b, s);
uint64_t sStride = 1;
uint64_t bStride = sStride * s;
stride = AscendC::MakeStride(bStride, sStride);
}
};
template <>
struct GmLayout<GmFormat::BNS2> {
AscendC::Shape<uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t s) {
shape = AscendC::MakeShape(b, n, s);
uint64_t sStride = 1;
uint64_t nStride = sStride * s;
uint64_t bStride = nStride * n;
stride = AscendC::MakeStride(bStride, nStride, sStride);
}
};
template <>
struct GmLayout<GmFormat::PA_BnBs> {
AscendC::Shape<uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t blockSize) {
shape = AscendC::MakeShape(blockSize);
uint64_t bsStride = 1;
uint64_t bnStride = bsStride * blockSize;
stride = AscendC::MakeStride(bnStride, bsStride);
}
};
template <>
struct GmLayout<GmFormat::PA_BnNBs> {
AscendC::Shape<uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t n, uint32_t blockSize) {
shape = AscendC::MakeShape(n, blockSize);
uint64_t bsStride = 1;
uint64_t nStride = bsStride * blockSize;
uint64_t bnStride = nStride * n; //blockSize * kvHeadNum
stride = AscendC::MakeStride(bnStride, nStride, bsStride);
}
};
//PSE_GmLayout
template <>
struct GmLayout<GmFormat::BN2GS1S2> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t g, uint32_t s1, uint32_t s2)
{
shape = AscendC::MakeShape(b, n, g, s1, s2);
uint64_t s2Stride = 1;
uint64_t s1Stride = s2Stride * s2;
uint64_t gStride = s1Stride * s1;
uint64_t nStride = gStride * g;
uint64_t bStride = nStride * n;
stride = AscendC::MakeStride(bStride, nStride, gStride, s1Stride, s2Stride);
}
};
template <>
struct GmLayout<GmFormat::SBNGD> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t g, uint32_t s, uint32_t d) {
shape = AscendC::MakeShape(b, n, g, s, d);
uint64_t dStride = 1;
uint64_t gStride = dStride * d;
uint64_t nStride = gStride * g;
uint64_t bStride = nStride * n;
uint64_t sStride = bStride * b;
stride = AscendC::MakeStride(bStride, nStride, gStride, sStride, dStride);
}
};
template <>
struct GmLayout<GmFormat::SBND> {
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
__aicore__ inline GmLayout() = default;
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t s, uint32_t d) {
shape = AscendC::MakeShape(b, n, s, d);
uint64_t dStride = 1;
uint64_t nStride = dStride * d;
uint64_t bStride = nStride * n;
uint64_t sStride = bStride * b;
stride = AscendC::MakeStride(bStride, nStride, sStride, dStride);
}
};
#endif

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,140 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file parser.h
* \brief
*/
#ifndef PARSER_H
#define PARSER_H
#if ASC_DEVKIT_MAJOR >= 9
#include "kernel_vec_intf.h"
#include "kernel_cube_intf.h"
#else
#include "kernel_operator.h"
#endif
using AscendC::GlobalTensor;
// ----------------------------------------------ActualSeqLensParser--------------------------------
enum class ActualSeqLensMode
{
BY_BATCH = 0,
ACCUM = 1,
};
template <ActualSeqLensMode MODE, typename ACTLEN_T = uint64_t>
class ActualSeqLensParser {
};
template <typename ACTLEN_T>
class ActualSeqLensParser<ActualSeqLensMode::ACCUM, ACTLEN_T> {
public:
__aicore__ inline ActualSeqLensParser() = default;
__aicore__ inline void Init(GlobalTensor<ACTLEN_T> actualSeqLengthsGm, uint32_t actualLenDims,
uint64_t defaultVal = 0)
{
this->actualSeqLengthsGm = actualSeqLengthsGm;
this->actualLenDims = actualLenDims;
}
__aicore__ inline uint64_t GetTBase(uint32_t bIdx) const
{
if (bIdx == 0) {
return 0;
}
return actualSeqLengthsGm.GetValue(bIdx - 1);
}
__aicore__ inline uint64_t GetMxVscaleTBase(uint32_t bIdx) const
{
if (bIdx == 0) {
return 0;
}
uint64_t vScaleTBaseOffset = 0;
for (uint32_t idx = 0; idx < bIdx; idx++) {
vScaleTBaseOffset += ((GetActualSeqLength(idx) + 63) >> 6);
}
return vScaleTBaseOffset;
}
__aicore__ inline uint64_t GetActualSeqLength(uint32_t bIdx) const
{
if (bIdx == 0) {
return actualSeqLengthsGm.GetValue(0);
}
return (actualSeqLengthsGm.GetValue(bIdx) - actualSeqLengthsGm.GetValue(bIdx - 1));
}
__aicore__ inline uint64_t GetTSize() const
{
return actualSeqLengthsGm.GetValue(actualLenDims - 1);
}
private:
GlobalTensor<ACTLEN_T> actualSeqLengthsGm;
uint32_t actualLenDims;
};
template <typename ACTLEN_T>
class ActualSeqLensParser<ActualSeqLensMode::BY_BATCH, ACTLEN_T> {
public:
__aicore__ inline ActualSeqLensParser() = default;
__aicore__ inline void Init(GlobalTensor<ACTLEN_T> actualSeqLengthsGm, uint32_t actualLenDims, uint64_t defaultVal)
{
this->actualSeqLengthsGm = actualSeqLengthsGm;
this->actualLenDims = actualLenDims;
this->defaultVal = defaultVal;
}
__aicore__ inline uint64_t GetActualSeqLength(uint32_t bIdx) const
{
if (actualLenDims == 0) {
return defaultVal;
}
if (actualLenDims == 1) {
return actualSeqLengthsGm.GetValue(0);
}
return actualSeqLengthsGm.GetValue(bIdx);
}
__aicore__ inline uint32_t GetActualLenDims() const
{
return actualLenDims;
}
private:
GlobalTensor<ACTLEN_T> actualSeqLengthsGm;
uint32_t actualLenDims = 0;
uint64_t defaultVal = 0;
};
// ----------------------------------------------BlockTableParser--------------------------------
class BlockTableParser {
public:
__aicore__ inline BlockTableParser() = default;
__aicore__ inline void Init(GlobalTensor<int32_t> blockTableGm, uint32_t maxblockNumPerBatch)
{
this->blockTableGm = blockTableGm;
this->maxblockNumPerBatch = maxblockNumPerBatch;
}
__aicore__ inline int32_t GetBlockIdx(uint32_t bIdx, uint32_t blockIdxInBatch) const
{
return blockTableGm.GetValue(bIdx * maxblockNumPerBatch + blockIdxInBatch);
}
private:
GlobalTensor<int32_t> blockTableGm;
uint32_t maxblockNumPerBatch;
};
#endif

View File

@@ -0,0 +1,31 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file offset_calculator.h
* \brief
*/
#ifndef OFFSET_CALCULATOR_H
#define OFFSET_CALCULATOR_H
#if ASC_DEVKIT_MAJOR >= 9
#include "kernel_basic_intf.h"
#else
#include "kernel_operator.h"
#endif
#include "memcopy/gm_layout.h"
#include "memcopy/parser.h"
#include "memcopy/offset_calculator_v2.h"
#include "memcopy/fa_gm_tensor.h"
#include "memcopy/fa_l1_tensor.h"
#include "memcopy/gm_coord.h"
#endif