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,89 @@
/**
* 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 hc_pre.cpp
* \brief
*/
#include "kernel_operator.h"
#include "kernel_operator_intf.h"
#if defined(__DAV_C310__)
#include "hc_pre_m_k_split_core_arch35.h"
#include "hc_pre_m_split_core_arch35.h"
#include "hc_pre_base_arch35.h"
using namespace HcPreNs;
#else
#include "lib/matmul_intf.h"
#include "hc_pre_m_k_split_core.h"
#include "hc_pre_base.h"
using namespace HcPre;
#endif
using namespace AscendC;
extern "C" __global__ __aicore__ void hc_pre(GM_ADDR x, GM_ADDR hc_fn, GM_ADDR hc_scale, GM_ADDR hc_base,
GM_ADDR y, GM_ADDR post, GM_ADDR comb_frag, GM_ADDR workspace,
GM_ADDR tiling)
{
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2);
if (workspace == nullptr) {
return;
}
GM_ADDR userWs = GetUserWorkspace(workspace);
if (userWs == nullptr) {
return;
}
TPipe pipe;
// 950PR 950DT
#if defined(__DAV_C310__)
if (TILING_KEY_IS(1000)) {
GET_TILING_DATA_WITH_STRUCT(HcPreTilingData, tiling_data_in, tiling);
const HcPreTilingData *__restrict tilingData = &tiling_data_in;
HcPreNs::HcPreMKSplitCorePart1<DTYPE_X> op;
op.Init(x, hc_fn, userWs, tilingData, &pipe);
op.Process();
pipe.Destroy();
TPipe pipeStage2;
HcPreNs::HcPreMKSplitCorePart2<DTYPE_X> op2;
op2.Init(x, hc_scale, hc_base, y, post, comb_frag, userWs, tilingData, &pipeStage2);
op2.Process();
pipeStage2.Destroy();
} else if (TILING_KEY_IS(1001)) {
GET_TILING_DATA_WITH_STRUCT(HcPreTilingData, tiling_data_in, tiling);
const HcPreTilingData *__restrict tilingData = &tiling_data_in;
HcPreNs::HcPreMSplitCorePart1<DTYPE_X> op;
op.Init(x, hc_fn, hc_scale, hc_base, y, post, comb_frag, tilingData, &pipe);
op.Process();
}
#else
// A3
if (TILING_KEY_IS(0)) {
GET_TILING_DATA_WITH_STRUCT(HcPreTilingData, tiling_data_in, tiling);
const HcPreTilingData *__restrict tilingData = &tiling_data_in;
HcPre::HcPreMembaseKSplitCorePart1<DTYPE_X> op;
op.Init(x, hc_fn, userWs, tilingData, &pipe);
op.Process();
pipe.Destroy();
TPipe pipeStage2;
HcPre::HcPreMembaseKSplitCorePart2<DTYPE_X> op2;
op2.Init(x, hc_scale, hc_base, y, post, comb_frag, userWs, tilingData, &pipeStage2);
op2.Process();
pipeStage2.Destroy();
}
#endif
}

View File

@@ -0,0 +1,688 @@
/**
* 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 hc_pre_base.h
* \brief
*/
#ifndef HC_PRE_VECTOR_BASE_H
#define HC_PRE_VECTOR_BASE_H
#include "kernel_operator.h"
namespace HcPre {
using namespace AscendC;
constexpr int32_t BLOCK_SIZE = 32;
constexpr int32_t DEFAULT_BLOCK_STRIDE = 1;
constexpr int32_t DEFAULT_REPEAT_STRIDE = 8;
constexpr int32_t ONE_REPEAT_BLOCK_NUMS = 8;
constexpr int32_t REPEAT_SIZE = 256;
constexpr int32_t MAX_REPEAT_STRIDE = 255;
constexpr int32_t WORKSPACE_ALIGN_SIZE = 512;
constexpr uint32_t ONE = 1;
constexpr uint32_t SHIFT_COEFF = 17;
constexpr int32_t DOUBLE_BUFFER = 2;
constexpr int32_t CV_RATIO = 2;
constexpr uint64_t N_SIZE = 24;
constexpr uint64_t NUM_TWO = 2;
constexpr uint64_t SQUARE_SUM_SIZE = 16;
constexpr uint64_t MASK_PATTERN_BASE_SIZE = 16;
constexpr uint64_t MASK_PATTERN_REPEAT_SIZE = 8;
constexpr uint64_t MASK_PATTERN_DIM_SIZE = 16;
constexpr uint64_t MM_CACHE_LINE_BYTES = 512;
__aicore__ inline int32_t CeilDiv(int32_t a, int32_t b)
{
if (b == 0) {
return a;
}
return (a + b - 1) / b;
}
__aicore__ inline int32_t CeilAlign(int32_t a, int32_t b)
{
return CeilDiv(a, b) * b;
}
template <typename T>
__aicore__ inline int32_t RoundUp(int32_t num)
{
int32_t elemNum = BLOCK_SIZE / sizeof(T);
return CeilAlign(num, elemNum);
}
__aicore__ inline void SetGatherMaskPattern(const LocalTensor<uint32_t>& maskPattern)
{
uint32_t base = ONE;
for (uint32_t i = 0; i < 16; i++) {
Duplicate(maskPattern[i * 8], base, 8);
base = (base << 1);
}
PipeBarrier<PIPE_V>();
}
__aicore__ inline void GatherMaskByDiagonal(const LocalTensor<float>& output, const LocalTensor<float>& input,
const LocalTensor<uint32_t> maskPattern, uint16_t dim0)
{
uint32_t totalCount = MASK_PATTERN_DIM_SIZE * MASK_PATTERN_DIM_SIZE;
uint32_t remainCount = (dim0 % MASK_PATTERN_DIM_SIZE) * MASK_PATTERN_DIM_SIZE;
uint32_t loopCount = dim0 / MASK_PATTERN_DIM_SIZE;
uint64_t rsvdCnt = 0;
for (uint32_t loopIdx = 0; loopIdx < loopCount; loopIdx++) {
GatherMask(output[loopIdx * MASK_PATTERN_DIM_SIZE],
input[loopIdx * totalCount], maskPattern, true, totalCount,
{1, 1, 8, 8}, rsvdCnt);
}
GatherMask(output[loopCount * MASK_PATTERN_DIM_SIZE],
input[loopCount * totalCount], maskPattern, true, remainCount,
{1, 1, 8, 8}, rsvdCnt);
PipeBarrier<PIPE_V>();
}
template <typename T, bool needBrc = true>
__aicore__ inline void MulABLastDimBrcInline(const LocalTensor<T> &output, const LocalTensor<T> &input0,
const LocalTensor<T> &input1, const LocalTensor<T> &tmpBuffer,
const int32_t curRowNum, const int32_t curColNum)
{
if constexpr (needBrc) {
uint32_t repeatTimes = CeilDiv(curRowNum, ONE_REPEAT_BLOCK_NUMS);
Brcb(tmpBuffer, input1, repeatTimes, {DEFAULT_BLOCK_STRIDE, DEFAULT_REPEAT_STRIDE});
}
PipeBarrier<PIPE_V>();
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
uint32_t curColNumAlign = RoundUp<T>(curColNum);
if (curColNum <= elemInOneBlock) {
Mul(output, input0, tmpBuffer, curRowNum * curColNumAlign);
} else {
int32_t numRepeatPerLine = curColNum / elemInOneRepeat;
int32_t numRemainPerLine = curColNum % elemInOneRepeat;
int32_t dstRepStridePerLine = CeilDiv(curColNum, elemInOneBlock);
BinaryRepeatParams instrParams;
if (numRepeatPerLine > 0) {
if (dstRepStridePerLine > MAX_REPEAT_STRIDE || curRowNum < numRepeatPerLine) {
// 在Col方向开Repeat, 并且Repeat小于255
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src1RepStride = 0;
for (uint32_t i = 0; i < curRowNum; i++) {
Mul(output[i * curColNumAlign], input0[i * curColNumAlign], tmpBuffer[i * elemInOneBlock],
elemInOneRepeat, numRepeatPerLine, instrParams);
}
} else {
// 在Row方向开Repeat
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = dstRepStridePerLine;
instrParams.src0RepStride = dstRepStridePerLine;
instrParams.src1RepStride = 1;
for (uint32_t i = 0; i < numRepeatPerLine; i++) {
Mul(output[i * elemInOneRepeat], input0[i * elemInOneRepeat], tmpBuffer, elemInOneRepeat, curRowNum,
instrParams);
}
}
}
if (numRemainPerLine > 0) {
if (dstRepStridePerLine > MAX_REPEAT_STRIDE) {
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = 0;
instrParams.src0RepStride = 0;
instrParams.src1RepStride = 0;
for (uint32_t i = 0; i < curRowNum; i++) {
Mul(output[numRepeatPerLine * elemInOneRepeat + i * curColNumAlign],
input0[numRepeatPerLine * elemInOneRepeat + i * curColNumAlign], tmpBuffer[i * elemInOneBlock],
numRemainPerLine, 1, instrParams);
}
} else {
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = dstRepStridePerLine;
instrParams.src0RepStride = dstRepStridePerLine;
instrParams.src1RepStride = 0;
Mul(output[numRepeatPerLine * elemInOneRepeat], input0[numRepeatPerLine * elemInOneRepeat], tmpBuffer,
numRemainPerLine, curRowNum, instrParams);
}
}
}
PipeBarrier<PIPE_V>();
}
template <typename T, bool needBrc = true>
__aicore__ inline void SubABLastDimBrcInline(const LocalTensor<T> &output, const LocalTensor<T> &input0,
const LocalTensor<T> &input1, const LocalTensor<T> &tmpBuffer,
const int32_t curRowNum, const int32_t curColNum)
{
if constexpr (needBrc) {
uint32_t repeatTimes = CeilDiv(curRowNum, ONE_REPEAT_BLOCK_NUMS);
Brcb(tmpBuffer, input1, repeatTimes, {DEFAULT_BLOCK_STRIDE, DEFAULT_REPEAT_STRIDE});
}
PipeBarrier<PIPE_V>();
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
uint32_t curColNumAlign = RoundUp<T>(curColNum);
if (curColNum <= elemInOneBlock) {
Sub(output, input0, tmpBuffer, curRowNum * curColNumAlign);
} else {
int32_t numRepeatPerLine = curColNum / elemInOneRepeat;
int32_t numRemainPerLine = curColNum % elemInOneRepeat;
int32_t dstRepStridePerLine = CeilDiv(curColNum, elemInOneBlock);
BinaryRepeatParams instrParams;
if (numRepeatPerLine > 0) {
if (dstRepStridePerLine > MAX_REPEAT_STRIDE || curRowNum < numRepeatPerLine) {
// 在Col方向开Repeat, 并且Repeat小于255
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src1RepStride = 0;
for (uint32_t i = 0; i < curRowNum; i++) {
Sub(output[i * curColNumAlign], input0[i * curColNumAlign], tmpBuffer[i * elemInOneBlock],
elemInOneRepeat, numRepeatPerLine, instrParams);
}
} else {
// 在Row方向开Repeat
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = dstRepStridePerLine;
instrParams.src0RepStride = dstRepStridePerLine;
instrParams.src1RepStride = 1;
for (uint32_t i = 0; i < numRepeatPerLine; i++) {
Sub(output[i * elemInOneRepeat], input0[i * elemInOneRepeat], tmpBuffer, elemInOneRepeat, curRowNum,
instrParams);
}
}
}
if (numRemainPerLine > 0) {
if (dstRepStridePerLine > MAX_REPEAT_STRIDE) {
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = 0;
instrParams.src0RepStride = 0;
instrParams.src1RepStride = 0;
for (uint32_t i = 0; i < curRowNum; i++) {
Sub(output[numRepeatPerLine * elemInOneRepeat + i * curColNumAlign],
input0[numRepeatPerLine * elemInOneRepeat + i * curColNumAlign], tmpBuffer[i * elemInOneBlock],
numRemainPerLine, 1, instrParams);
}
} else {
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = dstRepStridePerLine;
instrParams.src0RepStride = dstRepStridePerLine;
instrParams.src1RepStride = 0;
Sub(output[numRepeatPerLine * elemInOneRepeat], input0[numRepeatPerLine * elemInOneRepeat], tmpBuffer,
numRemainPerLine, curRowNum, instrParams);
}
}
}
PipeBarrier<PIPE_V>();
}
template <typename T, bool needBrc = true>
__aicore__ inline void DivABLastDimBrcInline(const LocalTensor<T> &output, const LocalTensor<T> &input0,
const LocalTensor<T> &input1, const LocalTensor<T> &tmpBuffer,
const int32_t curRowNum, const int32_t curColNum)
{
if constexpr (needBrc) {
uint32_t repeatTimes = CeilDiv(curRowNum, ONE_REPEAT_BLOCK_NUMS);
Brcb(tmpBuffer, input1, repeatTimes, {DEFAULT_BLOCK_STRIDE, DEFAULT_REPEAT_STRIDE});
}
PipeBarrier<PIPE_V>();
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
uint32_t curColNumAlign = RoundUp<T>(curColNum);
if (curColNum <= elemInOneBlock) {
Div(output, input0, tmpBuffer, curRowNum * curColNumAlign);
} else {
int32_t numRepeatPerLine = curColNum / elemInOneRepeat;
int32_t numRemainPerLine = curColNum % elemInOneRepeat;
int32_t dstRepStridePerLine = CeilDiv(curColNum, elemInOneBlock);
BinaryRepeatParams instrParams;
if (numRepeatPerLine > 0) {
if (dstRepStridePerLine > MAX_REPEAT_STRIDE || curRowNum < numRepeatPerLine) {
// 在Col方向开Repeat, 并且Repeat小于255
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src1RepStride = 0;
for (uint32_t i = 0; i < curRowNum; i++) {
Div(output[i * curColNumAlign], input0[i * curColNumAlign], tmpBuffer[i * elemInOneBlock],
elemInOneRepeat, numRepeatPerLine, instrParams);
}
} else {
// 在Row方向开Repeat
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = dstRepStridePerLine;
instrParams.src0RepStride = dstRepStridePerLine;
instrParams.src1RepStride = 1;
for (uint32_t i = 0; i < numRepeatPerLine; i++) {
Div(output[i * elemInOneRepeat], input0[i * elemInOneRepeat], tmpBuffer, elemInOneRepeat, curRowNum,
instrParams);
}
}
}
if (numRemainPerLine > 0) {
if (dstRepStridePerLine > MAX_REPEAT_STRIDE) {
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = 0;
instrParams.src0RepStride = 0;
instrParams.src1RepStride = 0;
for (uint32_t i = 0; i < curRowNum; i++) {
Div(output[numRepeatPerLine * elemInOneRepeat + i * curColNumAlign],
input0[numRepeatPerLine * elemInOneRepeat + i * curColNumAlign], tmpBuffer[i * elemInOneBlock],
numRemainPerLine, 1, instrParams);
}
} else {
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = dstRepStridePerLine;
instrParams.src0RepStride = dstRepStridePerLine;
instrParams.src1RepStride = 0;
Div(output[numRepeatPerLine * elemInOneRepeat], input0[numRepeatPerLine * elemInOneRepeat], tmpBuffer,
numRemainPerLine, curRowNum, instrParams);
}
}
}
PipeBarrier<PIPE_V>();
}
template <typename T>
__aicore__ inline void AddBAFirstDimBrcInline(const LocalTensor<T> &output, const LocalTensor<T> &input0,
const LocalTensor<T> &input1, const int32_t curRowNum,
const int32_t curColNum)
{
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
uint32_t curColNumAlign = RoundUp<T>(curColNum);
int32_t numRepeatPerLine = curColNum / elemInOneRepeat;
int32_t numRemainPerLine = curColNum % elemInOneRepeat;
int32_t dstRepStridePerLine = CeilDiv(curColNum, elemInOneBlock);
BinaryRepeatParams instrParams;
if (numRepeatPerLine > 0) {
if (dstRepStridePerLine > MAX_REPEAT_STRIDE || curRowNum < numRepeatPerLine) {
// 在Col方向开Repeat, 并且Repeat小于255
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 1;
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src1RepStride = DEFAULT_REPEAT_STRIDE;
for (uint32_t i = 0; i < curRowNum; i++) {
Add(output[i * curColNumAlign], input0[i * curColNumAlign], input1, elemInOneRepeat, numRepeatPerLine,
instrParams);
}
} else {
// 在Row方向开Repeat
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 1;
instrParams.dstRepStride = dstRepStridePerLine;
instrParams.src0RepStride = dstRepStridePerLine;
instrParams.src1RepStride = 0;
for (uint32_t i = 0; i < numRepeatPerLine; i++) {
Add(output[i * elemInOneRepeat], input0[i * elemInOneRepeat], input1[i * elemInOneRepeat],
elemInOneRepeat, curRowNum, instrParams);
}
}
}
if (numRemainPerLine > 0) {
if (dstRepStridePerLine > MAX_REPEAT_STRIDE) {
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 1;
instrParams.dstRepStride = 0;
instrParams.src0RepStride = 0;
instrParams.src1RepStride = 0;
for (uint32_t i = 0; i < curRowNum; i++) {
Add(output[numRepeatPerLine * elemInOneRepeat], input0[numRepeatPerLine * elemInOneRepeat], input1,
numRemainPerLine, 1, instrParams);
}
} else {
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 1;
instrParams.dstRepStride = dstRepStridePerLine;
instrParams.src0RepStride = dstRepStridePerLine;
instrParams.src1RepStride = 0;
Add(output[numRepeatPerLine * elemInOneRepeat], input0[numRepeatPerLine * elemInOneRepeat], input1,
numRemainPerLine, curRowNum, instrParams);
}
}
PipeBarrier<PIPE_V>();
}
template <typename T>
__aicore__ inline void CalcDenominator(const LocalTensor<T> &output, const LocalTensor<T> &input,
const uint32_t calCount)
{
Muls(output, input, static_cast<T>(-1.0), calCount);
PipeBarrier<PIPE_V>();
Exp(output, output, calCount);
PipeBarrier<PIPE_V>();
Adds(output, output, static_cast<T>(1.0), calCount);
PipeBarrier<PIPE_V>();
}
// 暂时不处理repeat超限场景
template <typename T>
__aicore__ inline void SigmoidPerf(const LocalTensor<T> &output, const LocalTensor<T> &input,
const LocalTensor<T> &tmpBuffer, const int64_t calCount)
{
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
Duplicate(tmpBuffer, static_cast<T>(1.0), elemInOneBlock);
CalcDenominator(output, input, calCount);
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
int32_t numRepeatPerLine = calCount / elemInOneRepeat;
int32_t numRemainPerLine = calCount % elemInOneRepeat;
BinaryRepeatParams instrParams;
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 0;
instrParams.src1BlkStride = 1;
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src0RepStride = 0;
instrParams.src1RepStride = DEFAULT_REPEAT_STRIDE;
Div(output, tmpBuffer, output, elemInOneRepeat, numRepeatPerLine, instrParams);
if (numRemainPerLine != 0) {
Div(output[numRepeatPerLine * elemInOneRepeat], tmpBuffer, output[numRepeatPerLine * elemInOneRepeat],
numRemainPerLine, 1, instrParams);
}
PipeBarrier<PIPE_V>();
}
__aicore__ inline void ProcessPre(const LocalTensor<float> &preLocal, const LocalTensor<float> &mixLocal,
const LocalTensor<float> &hcBaseLocal, const LocalTensor<float> &rsqrtLocal,
const LocalTensor<float> &tmpBuffer0, const LocalTensor<float> &tmpBuffer1,
float scale, float eps, const int32_t curRowNum, const int32_t curColNum)
{
int32_t curColNumAlign = RoundUp<float>(curColNum);
MulABLastDimBrcInline<float, true>(mixLocal, mixLocal, rsqrtLocal, tmpBuffer0, curRowNum, curColNum);
Muls(mixLocal, mixLocal, scale, curRowNum * curColNumAlign);
PipeBarrier<PIPE_V>();
AddBAFirstDimBrcInline<float>(mixLocal, mixLocal, hcBaseLocal, curRowNum, curColNum);
SigmoidPerf(preLocal, mixLocal, tmpBuffer1, curRowNum * curColNumAlign);
Adds(preLocal, preLocal, eps, curRowNum * curColNumAlign);
PipeBarrier<PIPE_V>();
}
__aicore__ inline void ReduceSumARAPerf(const LocalTensor<float> &output, const LocalTensor<float> &input,
const uint32_t dim0, const uint32_t dim1, const uint32_t dim2)
{
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(float);
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(float);
uint32_t dim2Align = RoundUp<float>(dim2);
// 拷贝第一个R到output上
DataCopyParams copyParams;
copyParams.blockCount = dim0;
copyParams.blockLen = dim2Align / elemInOneBlock;
copyParams.srcStride = (dim1 - 1) * (dim2Align / elemInOneBlock);
copyParams.dstStride = 0;
DataCopy(output, input, copyParams);
PipeBarrier<PIPE_V>();
uint32_t dim2RepeatTimes = dim2 / elemInOneRepeat;
uint32_t dim2Reminder = dim2 % elemInOneRepeat;
// 沿着dim2方向开repeat
BinaryRepeatParams instrParams;
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 1;
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src1RepStride = DEFAULT_REPEAT_STRIDE;
for (uint32_t i = 0; i < dim0; i++) {
for (uint32_t j = 1; j < dim1; j++) {
Add(output[i * dim2Align], output[i * dim2Align], input[i * dim1 * dim2Align + j * dim2Align],
elemInOneRepeat, dim2RepeatTimes, instrParams);
if (dim2Reminder != 0) {
Add(output[i * dim2Align + dim2RepeatTimes * elemInOneRepeat],
output[i * dim2Align + dim2RepeatTimes * elemInOneRepeat],
input[i * dim1 * dim2Align + j * dim2Align + +dim2RepeatTimes * elemInOneRepeat], dim2Reminder, 1,
instrParams);
}
PipeBarrier<PIPE_V>();
}
}
PipeBarrier<PIPE_V>();
}
template <typename T0, typename T1>
__aicore__ inline void CastTwoDim(const LocalTensor<T0> &output, const LocalTensor<T1> &input, const uint32_t dim0,
const uint32_t dim1)
{
uint32_t dim1AlignT0 = RoundUp<T0>(dim1);
uint32_t dim1AlignT1 = RoundUp<T1>(dim1);
if constexpr (IsSameType<T1, bfloat16_t>::value && IsSameType<T0, float>::value) {
for (uint32_t i = 0; i < dim0; i++) {
Cast(output[i * dim1AlignT0], input[i * dim1AlignT1], AscendC::RoundMode::CAST_NONE, dim1);
}
} else {
for (uint32_t i = 0; i < dim0; i++) {
Cast(output[i * dim1AlignT0], input[i * dim1AlignT1], AscendC::RoundMode::CAST_RINT, dim1);
}
}
PipeBarrier<PIPE_V>();
}
template <typename T>
__aicore__ void inline ProcessY(const LocalTensor<T> &yLocal, const LocalTensor<T> &xLocal,
const LocalTensor<float> &mix01Local, const LocalTensor<float> &hcBrcbLocal1,
const LocalTensor<float> &xCastLocal, const LocalTensor<float> &yCastLocal,
const uint32_t dim0, const uint32_t dim1, const uint32_t dim2)
{
CastTwoDim(xCastLocal, xLocal, dim0 * dim1, dim2);
MulABLastDimBrcInline<float, true>(xCastLocal, xCastLocal, mix01Local, hcBrcbLocal1, dim0 * dim1, dim2);
ReduceSumARAPerf(yCastLocal, xCastLocal, dim0, dim1, dim2);
CastTwoDim(yLocal, yCastLocal, dim0, dim2);
}
__aicore__ inline void ProcessPost(const LocalTensor<float> &postLocal, const LocalTensor<float> &mixLocal,
const LocalTensor<float> &hcBaseLocal, const LocalTensor<float> &rsqrtLocal,
const LocalTensor<float> &tmpBuffer0, const LocalTensor<float> &tmpBuffer1,
float scale, const int32_t curRowNum, const int32_t curColNum)
{
int32_t curColNumAlign = RoundUp<float>(curColNum);
MulABLastDimBrcInline<float, false>(mixLocal, mixLocal, rsqrtLocal, tmpBuffer0, curRowNum, curColNum);
Muls(mixLocal, mixLocal, scale, curRowNum * curColNumAlign);
PipeBarrier<PIPE_V>();
AddBAFirstDimBrcInline<float>(mixLocal, mixLocal, hcBaseLocal, curRowNum, curColNum);
SigmoidPerf(postLocal, mixLocal, tmpBuffer1, curRowNum * curColNumAlign);
Muls(postLocal, postLocal, static_cast<float>(2.0f), curRowNum * curColNumAlign);
PipeBarrier<PIPE_V>();
}
__aicore__ inline void LastDimReduceMaxPerf(const LocalTensor<float> &output, const LocalTensor<float> &input,
const uint32_t curRowNum, const uint32_t curColNum)
{
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(float);
WholeReduceMax(output, input, curColNum, curRowNum, 1, 1, CeilDiv(curColNum, elemInOneBlock),
ReduceOrder::ORDER_ONLY_VALUE);
PipeBarrier<PIPE_V>();
}
__aicore__ inline void LastDimReduceSumPerf(const LocalTensor<float> &output, const LocalTensor<float> &input,
const uint32_t curRowNum, const uint32_t curColNum)
{
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(float);
WholeReduceSum(output, input, curColNum, curRowNum, 1, 1, CeilDiv(curColNum, elemInOneBlock));
PipeBarrier<PIPE_V>();
}
// 暂时只支持R轴小于64,既curColNum不能超过64
__aicore__ inline void SoftmaxFP32Perf(const LocalTensor<float> &output, const LocalTensor<float> &input,
const LocalTensor<float> &tmpReduceBuffer,
const LocalTensor<float> tmpBrcbBuffer, const int32_t curRowNum,
const int32_t curColNum, float eps)
{
LastDimReduceMaxPerf(tmpReduceBuffer, input, curRowNum, curColNum);
SubABLastDimBrcInline<float, true>(output, input, tmpReduceBuffer, tmpBrcbBuffer, curRowNum, curColNum);
uint32_t curColNumAlign = RoundUp<float>(curColNum);
Exp(output, output, curRowNum * curColNumAlign);
PipeBarrier<PIPE_V>();
LastDimReduceSumPerf(tmpReduceBuffer, output, curRowNum, curColNum);
DivABLastDimBrcInline<float, true>(output, output, tmpReduceBuffer, tmpBrcbBuffer, curRowNum, curColNum);
Adds(output, output, eps, curRowNum * curColNumAlign);
PipeBarrier<PIPE_V>();
}
// (bs, hc_mult, hc_mult) = (bs, hc_mult, hc_mult) + (bs, 1, hc_mult)
template <typename T>
__aicore__ inline void DivABABrcInline(const LocalTensor<T> &output, const LocalTensor<T> &input0,
const LocalTensor<T> &input1, const uint32_t dim0, const uint32_t dim1,
const uint32_t dim2)
{
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
uint32_t dim2Align = RoundUp<T>(dim2);
uint32_t dim2RepeatTimes = dim2 / elemInOneRepeat;
uint32_t dim2Reminder = dim2 % elemInOneRepeat;
uint32_t dim2RepeatStride = CeilDiv(dim2, elemInOneBlock);
// 在dim1方向开repeat
BinaryRepeatParams instrParams;
if (dim1 >= dim2RepeatTimes) {
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 1;
instrParams.dstRepStride = dim2RepeatStride;
instrParams.src0RepStride = dim2RepeatStride;
instrParams.src1RepStride = 0;
for (uint32_t i = 0; i < dim0; i++) {
for (uint32_t j = 0; j < dim2RepeatTimes; j++) {
Div(output[i * dim1 * dim2Align + j * elemInOneRepeat],
input0[i * dim1 * dim2Align + j * elemInOneRepeat], input1[i * dim2Align + j * elemInOneRepeat],
elemInOneRepeat, dim1, instrParams);
}
if (dim2Reminder != 0) {
Div(output[i * dim1 * dim2Align + dim2RepeatTimes * elemInOneRepeat],
input0[i * dim1 * dim2Align + dim2RepeatTimes * elemInOneRepeat],
input1[i * dim2Align + dim2RepeatTimes * elemInOneRepeat], dim2Reminder, dim1, instrParams);
}
}
} else {
// 在dim2方向开repeat
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 1;
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src1RepStride = DEFAULT_REPEAT_STRIDE;
for (uint32_t i = 0; i < dim0; i++) {
for (uint32_t j = 0; j < dim1; j++) {
Div(output[i * dim1 * dim2Align + j * dim2Align], input0[i * dim1 * dim2Align + j * dim2Align],
input1[i * dim2Align], dim2);
}
}
}
PipeBarrier<PIPE_V>();
}
template <typename T>
__aicore__ inline void CopyIn(const GlobalTensor<T> &inputGm, const LocalTensor<T> &inputTensor, const uint16_t nBurst,
const uint32_t copyLen, uint32_t srcStride = 0, uint32_t dstStride = 0)
{
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
DataCopyPadExtParams<T> dataCopyPadExtParams;
dataCopyPadExtParams.isPad = false;
dataCopyPadExtParams.leftPadding = 0;
dataCopyPadExtParams.rightPadding = 0;
dataCopyPadExtParams.paddingValue = 0;
DataCopyExtParams dataCoptExtParams;
dataCoptExtParams.blockCount = nBurst;
dataCoptExtParams.blockLen = copyLen * sizeof(T);
dataCoptExtParams.srcStride = srcStride * sizeof(T);
dataCoptExtParams.dstStride = dstStride / elemInOneBlock;
DataCopyPad(inputTensor, inputGm, dataCoptExtParams, dataCopyPadExtParams);
}
// (bs, hc_mult, hc_mult) --> (bs, hc_mult, hc_mult_align)
template <typename T>
__aicore__ inline void CopyInWithOuterFor(const GlobalTensor<T> &inputGm, const LocalTensor<T> &inputTensor,
const uint16_t outerLoop, const uint16_t nBurst, const uint32_t copyLen,
const uint32_t gmLastDim)
{
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
uint32_t ubLastDimAlign = RoundUp<T>(copyLen);
for (uint16_t i = 0; i < outerLoop; i++) {
CopyIn(inputGm[i * nBurst * gmLastDim], inputTensor[i * nBurst * ubLastDimAlign], nBurst, copyLen);
}
}
template <typename T>
__aicore__ inline void CopyInWithOuterFor(const GlobalTensor<T> &inputGm,
const LocalTensor<T> &inputTensor, const uint16_t outerLoop,
const uint16_t nBurst, const uint32_t copyLen, const uint32_t gmFirstDim,
const uint32_t gmLastDim)
{
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
uint32_t ubLastDimAlign = RoundUp<T>(copyLen);
if (outerLoop <= nBurst) {
for (uint32_t i = 0; i < outerLoop; i++) {
CopyIn(inputGm[i * gmFirstDim * gmLastDim],
inputTensor[i * nBurst * ubLastDimAlign], nBurst, copyLen,
gmLastDim - copyLen);
}
} else {
uint32_t srcStride = (gmLastDim - copyLen) + (gmFirstDim - 1) * gmLastDim;
uint32_t dstStride = (nBurst - 1) * ubLastDimAlign;
for (uint32_t i = 0; i < nBurst; i++) {
CopyIn(inputGm[i * gmLastDim], inputTensor[i * ubLastDimAlign], outerLoop, copyLen, srcStride, dstStride);
}
}
}
template <typename T>
__aicore__ inline void CopyOut(const LocalTensor<T> &outputTensor, const GlobalTensor<T> &outputGm,
const uint16_t nBurst, const uint32_t copyLen, uint32_t dstStride = 0)
{
DataCopyExtParams dataCopyParams;
dataCopyParams.blockCount = nBurst;
dataCopyParams.blockLen = copyLen * sizeof(T);
dataCopyParams.srcStride = 0;
dataCopyParams.dstStride = dstStride * sizeof(T);
DataCopyPad(outputGm, outputTensor, dataCopyParams);
}
} // namespace HcPre
#endif

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,481 @@
/**
* 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 hc_pre_cube_compute.h
* \brief
*/
#ifndef HC_PRE_CUBE_COMPUTE_H
#define HC_PRE_CUBE_COMPUTE_H
#include "kernel_operator.h"
#include "kernel_operator_intf.h"
#include "hc_pre_base.h"
using AscendC::BLOCK_CUBE;
using AscendC::GlobalTensor;
using AscendC::HardEvent;
using AscendC::LocalTensor;
using AscendC::Nd2NzParams;
using AscendC::SetFlag;
using AscendC::TBuf;
using AscendC::TPipe;
using AscendC::TPosition;
using AscendC::WaitFlag;
using namespace AscendC;
namespace HcPre {
struct MmParams {
uint64_t curML1;
uint64_t curKL1;
uint64_t curNL1;
uint64_t singleCoreK;
uint64_t kGmBaseOffset;
uint64_t nGmSize;
uint64_t kGmSize;
uint64_t nOutSize;
uint64_t xWsKSize;
bool isLastK;
bool isFirstK;
};
#define HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM template <bool enableSquareSum>
#define HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS HcCubeCompute<enableSquareSum>
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
class HcCubeCompute {
public:
__aicore__ inline HcCubeCompute(){};
__aicore__ inline void Init(const GlobalTensor<float>& xGm, const GlobalTensor<float>& fnGm, TPipe *tpipe);
__aicore__ inline void ComputeDecode(const AscendC::GlobalTensor<float> &xGm, const AscendC::GlobalTensor<float> &workspaceGlobalA2,
const AscendC::GlobalTensor<float> &workspaceGlobalAB, const MmParams &mmParams);
__aicore__ inline void CopyInB1(
uint64_t mGmOffset, uint64_t kGmOffset, uint64_t kL1Size, const MmParams &mmParams);
__aicore__ inline void SetBL1Mte1ToMte2Flag();
__aicore__ inline void WaitBL1Mte1ToMte2Flag();
__aicore__ inline void End();
private:
__aicore__ inline void CopyInA1(
uint64_t kL1Size,
const GlobalTensor<float> &aGlobal, const LocalTensor<float> &al1Local, const MmParams &mmParams);
__aicore__ inline void CopyOut(const AscendC::GlobalTensor<float> &workspaceGlobal,
const AscendC::LocalTensor<float> &c1Local, uint64_t baseM, uint64_t baseN, bool enableNz2Nd, uint64_t N);
__aicore__ inline void Fixp(const AscendC::GlobalTensor<float> &workspaceGlobalA2,
const AscendC::GlobalTensor<float> &workspaceGlobalAB, const MmParams &mmParams);
__aicore__ inline void LoadAToL0A(
uint64_t kL1Offset, uint64_t kL0Size, uint64_t l1LoopIdx, const MmParams &mmParams);
__aicore__ inline void LoadAToL0B(
uint64_t kL1Offset, uint64_t kL0Size, uint64_t l1LoopIdx, const MmParams &mmParams);
__aicore__ inline void MmadA2(uint64_t kGmOffset, uint64_t kL0Size, bool isLastK, const MmParams &mmParams);
__aicore__ inline void LoadBToL0B(
uint64_t kL1Offset, uint64_t kL0Size, uint64_t l1LoopIdx, const MmParams &mmParams);
__aicore__ inline void MmadAB(
uint64_t kGmOffset, uint64_t kL0Size, bool isLastK, const MmParams &mmParams);
TPipe *pipe_;
const HcPreTilingData *tiling_;
int32_t blkIdx_ = -1;
int64_t batch_ = 0;
int64_t hcParam_ = 0;
int64_t dParam_ = 0;
static constexpr int32_t ONE_BLOCK_SIZE = 32;
int32_t perBlock32 = ONE_BLOCK_SIZE / sizeof(float);
GlobalTensor<float> fnGm_;
GlobalTensor<float> yGm_;
static constexpr uint64_t MM1_MTE2_MTE1_EVENT = 2;
static constexpr uint64_t X_MTE1_MTE2_EVENT = 2;
static constexpr uint64_t B_MTE1_MTE2_EVENT = 4;
static constexpr uint64_t M_MTE1_EVENT_L0A = 3;
static constexpr uint64_t M_MTE1_EVENT_L0B = 5;
static constexpr uint64_t MTE1_M_EVENT = 2;
static constexpr uint64_t L1_BUF_NUM = 2;
static constexpr uint64_t L0A_BUF_NUM = 2;
static constexpr uint64_t L0B_BUF_NUM = 2;
static constexpr uint64_t L0AB_BUF_NUM = 2;
static constexpr uint64_t L0C_BUF_NUM = 2;
static constexpr uint64_t L1_BUF_OFFSET = 128 * 256;
static constexpr uint64_t L0AB_BUF_OFFSET = 32 * 256;
static constexpr uint64_t L0C_BUF_OFFSET = 64 * 256;
static constexpr uint64_t L0C_A2_BUF_OFFSET = 256 * 16;
static constexpr uint16_t UNIT_FLAG_ENABLE = 2;
static constexpr uint16_t UNIT_FLAG_ENABLE_AUTO_CLOSE = 3;
constexpr static uint32_t FINAL_ACCUMULATION = 3;
constexpr static uint32_t NON_FINAL_ACCUMULATION = 2;
static constexpr uint64_t K_L0_SIZE = 32UL;
static constexpr uint64_t FLOAT_C0_SIZE = 8UL;
uint64_t l1aLoopIdx_ = 0;
uint64_t l1bLoopIdx_ = 0;
uint64_t l0aLoopIdx_ = 0;
uint64_t l0bLoopIdx_ = 0;
uint64_t l0cLoopIdx_ = 0;
LocalTensor<float> l1a_;
LocalTensor<float> l1b_;
LocalTensor<float> l0a_;
LocalTensor<float> l0b_;
LocalTensor<float> l0c_;
uint64_t k_ = 0;
uint64_t n_ = 0;
};
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::Init(const GlobalTensor<float>& xGm, const GlobalTensor<float>& fnGm, TPipe *tpipe)
{
fnGm_ = fnGm;
TBuf<TPosition::A1> l1aBuffer;
tpipe->InitBuffer(l1aBuffer, 256 * 1024);
l1a_ = l1aBuffer.Get<float>();
TBuf<TPosition::B1> l1bBuffer;
tpipe->InitBuffer(l1bBuffer, 256 * 1024);
l1b_ = l1bBuffer.Get<float>();
TBuf<TPosition::A2> l0aBuffer;
tpipe->InitBuffer(l0aBuffer, 64 * 1024);
l0a_ = l0aBuffer.Get<float>();
TBuf<TPosition::B2> l0bBuffer;
tpipe->InitBuffer(l0bBuffer, 64 * 1024);
l0b_ = l0bBuffer.Get<float>();
// loc
TBuf<TPosition::CO1> l0cBuffer;
tpipe->InitBuffer(l0cBuffer, 64 * 1024);
l0c_ = l0cBuffer.Get<float>();
for (int i = 0; i < L0A_BUF_NUM; i++) {
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0A + i);
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0B + i);
}
for (int i = 0; i < L1_BUF_NUM; i++) {
SetFlag<HardEvent::MTE1_MTE2>(X_MTE1_MTE2_EVENT + i);
SetFlag<HardEvent::MTE1_MTE2>(B_MTE1_MTE2_EVENT + i);
}
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::CopyInA1(
uint64_t kL1Size,
const GlobalTensor<float> &aGlobal, const LocalTensor<float> &al1Local, const MmParams &mmParams)
{
Nd2NzParams nd2nzParams;
nd2nzParams.ndNum = 1;
nd2nzParams.nValue = mmParams.curML1;
nd2nzParams.dValue = kL1Size;
nd2nzParams.srcNdMatrixStride = 1;
nd2nzParams.srcDValue = mmParams.xWsKSize; // vec处理的singleK
nd2nzParams.dstNzC0Stride = (mmParams.curML1 + BLOCK_CUBE - 1) / BLOCK_CUBE * BLOCK_CUBE;
nd2nzParams.dstNzNStride = 1;
nd2nzParams.dstNzMatrixStride = 1;
DataCopy(al1Local, aGlobal, nd2nzParams);
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::CopyInB1(
uint64_t mGmOffset, uint64_t kGmOffset, uint64_t kL1Size, const MmParams &mmParams)
{
AscendC::Nd2NzParams nd2nzParams;
nd2nzParams.ndNum = 1;
nd2nzParams.nValue = mmParams.curNL1;
nd2nzParams.dValue = kL1Size;
nd2nzParams.srcNdMatrixStride = 1;
nd2nzParams.srcDValue = mmParams.kGmSize; // 原始k
nd2nzParams.dstNzC0Stride = (mmParams.curNL1 + BLOCK_CUBE - 1) / BLOCK_CUBE * BLOCK_CUBE;
nd2nzParams.dstNzNStride = 1;
nd2nzParams.dstNzMatrixStride = 1;
DataCopy(l1b_[(l1bLoopIdx_ % L1_BUF_NUM) * L1_BUF_OFFSET], fnGm_[kGmOffset], nd2nzParams);
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::CopyOut(const AscendC::GlobalTensor<float> &workspaceGlobal,
const AscendC::LocalTensor<float> &c1Local, uint64_t baseM, uint64_t baseN, bool enableNz2Nd, uint64_t N)
{
AscendC::DataCopyCO12DstParams intriParams;
intriParams.nSize = baseN;
intriParams.mSize = baseM;
// set mode to float32, then cast in ub
intriParams.quantPre = QuantMode_t::NoQuant;
intriParams.nz2ndEn = enableNz2Nd;
if (enableNz2Nd) {
intriParams.dstStride = N; // NZ -> ND
intriParams.srcStride = CeilAlign(baseM, AscendC::BLOCK_CUBE);
AscendC::SetFixpipeNz2ndFlag(1, 1, 1);
} else {
intriParams.dstStride = CeilAlign(intriParams.nSize, AscendC::BLOCK_CUBE); // NZ -> NZ
intriParams.srcStride = CeilAlign(baseM, AscendC::BLOCK_CUBE);
}
intriParams.unitFlag = UNIT_FLAG_ENABLE_AUTO_CLOSE; // 3
AscendC::DataCopy(workspaceGlobal, c1Local, intriParams);
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::Fixp(const AscendC::GlobalTensor<float> &workspaceGlobalA2,
const AscendC::GlobalTensor<float> &workspaceGlobalAB, const MmParams &mmParams)
{
// Copy MmadA2
CopyOut(workspaceGlobalA2,
l0c_[(l0cLoopIdx_ % L0C_BUF_NUM) * L0C_BUF_OFFSET],
mmParams.curML1,
BLOCK_CUBE,
false,
BLOCK_CUBE); // nz m,16
// Copy MmadAB
CopyOut(workspaceGlobalAB,
l0c_[(l0cLoopIdx_ % L0C_BUF_NUM) * L0C_BUF_OFFSET + L0C_A2_BUF_OFFSET],
mmParams.curML1,
mmParams.curNL1,
true,
mmParams.nOutSize); // ND M/N are 512B aligned.
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::SetBL1Mte1ToMte2Flag()
{
SetFlag<HardEvent::MTE1_MTE2>(B_MTE1_MTE2_EVENT + l1bLoopIdx_ % L1_BUF_NUM);
l1bLoopIdx_++;
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::WaitBL1Mte1ToMte2Flag()
{
WaitFlag<HardEvent::MTE1_MTE2>(B_MTE1_MTE2_EVENT + l1bLoopIdx_ % L1_BUF_NUM);
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::ComputeDecode(
const AscendC::GlobalTensor<float> &xGm,
const AscendC::GlobalTensor<float> &workspaceGlobalA2, const AscendC::GlobalTensor<float> &workspaceGlobalAB,
const MmParams &mmParams)
{
WaitFlag<HardEvent::MTE1_MTE2>(X_MTE1_MTE2_EVENT + l1aLoopIdx_ % L1_BUF_NUM);
uint64_t kGmOffset = 0;
uint64_t curKL1Size = (kGmOffset + mmParams.curKL1 >= mmParams.singleCoreK) ? (mmParams.singleCoreK - kGmOffset) : mmParams.curKL1;
CopyInA1(curKL1Size, xGm[kGmOffset], l1a_[(l1aLoopIdx_ % DOUBLE_BUFFER) * L1_BUF_OFFSET], mmParams);
SetFlag<HardEvent::MTE2_MTE1>(MM1_MTE2_MTE1_EVENT + l1aLoopIdx_ % L1_BUF_NUM);
kGmOffset += mmParams.curKL1;
while (kGmOffset < mmParams.singleCoreK) {
WaitFlag<HardEvent::MTE1_MTE2>(X_MTE1_MTE2_EVENT + (l1aLoopIdx_ + 1) % L1_BUF_NUM);
uint64_t nextKL1Size = (kGmOffset + mmParams.curKL1 >= mmParams.singleCoreK) ? (mmParams.singleCoreK - kGmOffset) : mmParams.curKL1;
CopyInA1(nextKL1Size, xGm[kGmOffset], l1a_[((l1aLoopIdx_ + 1) % DOUBLE_BUFFER) * L1_BUF_OFFSET], mmParams);
SetFlag<HardEvent::MTE2_MTE1>(MM1_MTE2_MTE1_EVENT + (l1aLoopIdx_ + 1) % L1_BUF_NUM);
WaitFlag<HardEvent::MTE2_MTE1>(MM1_MTE2_MTE1_EVENT + l1aLoopIdx_ % L1_BUF_NUM);
for (uint64_t kL1Offset = 0; kL1Offset < curKL1Size; kL1Offset += K_L0_SIZE) { // K_L0_SIZE 32
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0A + l0aLoopIdx_ % L0A_BUF_NUM);
LoadAToL0A(kL1Offset, K_L0_SIZE, l1aLoopIdx_, mmParams);
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0B + l0bLoopIdx_ % L0B_BUF_NUM);
LoadAToL0B(kL1Offset, K_L0_SIZE, l1aLoopIdx_, mmParams);
MmadA2(kGmOffset + kL1Offset - mmParams.curKL1, K_L0_SIZE,
false,
mmParams);
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0B + l0bLoopIdx_ % L0B_BUF_NUM);
l0bLoopIdx_++;
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0B + l0bLoopIdx_ % L0B_BUF_NUM);
LoadBToL0B((kGmOffset + kL1Offset - mmParams.curKL1) % mmParams.xWsKSize, K_L0_SIZE, l1bLoopIdx_, mmParams);
MmadAB(kGmOffset + kL1Offset - mmParams.curKL1, K_L0_SIZE,
false,
mmParams);
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0B + l0bLoopIdx_ % L0B_BUF_NUM);
l0bLoopIdx_++;
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0A + l0aLoopIdx_ % L0A_BUF_NUM);
l0aLoopIdx_++;
}
SetFlag<HardEvent::MTE1_MTE2>(X_MTE1_MTE2_EVENT + l1aLoopIdx_ % L1_BUF_NUM);
l1aLoopIdx_++;
kGmOffset += mmParams.curKL1;
curKL1Size = nextKL1Size;
}
WaitFlag<HardEvent::MTE2_MTE1>(MM1_MTE2_MTE1_EVENT + l1aLoopIdx_ % L1_BUF_NUM);
for (uint64_t kL1Offset = 0; kL1Offset < curKL1Size; kL1Offset += K_L0_SIZE) {
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0A + l0aLoopIdx_ % L0A_BUF_NUM);
LoadAToL0A(kL1Offset, K_L0_SIZE, l1aLoopIdx_, mmParams); // to l0a
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0B + l0bLoopIdx_ % L0B_BUF_NUM);
LoadAToL0B(kL1Offset, K_L0_SIZE, l1aLoopIdx_, mmParams); // to l0b
MmadA2(kGmOffset + kL1Offset - mmParams.curKL1, K_L0_SIZE,
mmParams.isLastK && kL1Offset + K_L0_SIZE >= curKL1Size, mmParams);
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0B + l0bLoopIdx_ % L0B_BUF_NUM);
l0bLoopIdx_++;
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0B + l0bLoopIdx_ % L0B_BUF_NUM);
LoadBToL0B((kGmOffset + kL1Offset - mmParams.curKL1) % mmParams.xWsKSize, K_L0_SIZE, l1bLoopIdx_, mmParams); // to l0b
MmadAB(kGmOffset + kL1Offset - mmParams.curKL1, K_L0_SIZE,
mmParams.isLastK && kL1Offset + K_L0_SIZE >= curKL1Size,
mmParams);
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0B + l0bLoopIdx_ % L0B_BUF_NUM);
l0bLoopIdx_++;
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0A + l0aLoopIdx_ % L0A_BUF_NUM);
l0aLoopIdx_++;
}
SetFlag<HardEvent::MTE1_MTE2>(X_MTE1_MTE2_EVENT + l1aLoopIdx_ % L1_BUF_NUM);
l1aLoopIdx_++;
if (mmParams.isLastK) {
Fixp(workspaceGlobalA2, workspaceGlobalAB, mmParams); // l0cLoopIdx_++;
}
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::End()
{
for (int i = 0; i < L0A_BUF_NUM; i++) {
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0A + i);
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT_L0B + i);
}
for (int i = 0; i < L1_BUF_NUM; i++) {
WaitFlag<HardEvent::MTE1_MTE2>(X_MTE1_MTE2_EVENT + i);
WaitFlag<HardEvent::MTE1_MTE2>(B_MTE1_MTE2_EVENT + i);
}
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::LoadAToL0A(
uint64_t kL1Offset, uint64_t kL0Size, uint64_t l1LoopIdx, const MmParams &mmParams)
{
static constexpr IsResetLoad3dConfig LOAD3DV2_CONFIG = {true, true};
LoadData3DParamsV2<float> loadData3DParams;
// SetFmatrixParams
loadData3DParams.l1H = CeilDiv(mmParams.curML1, BLOCK_CUBE); // Hin=M1=8
loadData3DParams.l1W = BLOCK_CUBE; // Win=M0
loadData3DParams.channelSize = kL0Size; // Cin=K
loadData3DParams.padList[0] = 0;
loadData3DParams.padList[1] = 0;
loadData3DParams.padList[2] = 0;
loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果
// SetLoadToA0Params
loadData3DParams.mExtension = CeilAlign(mmParams.curML1, BLOCK_CUBE); // M height维度目的
loadData3DParams.kExtension = kL0Size; // K width维度目的
loadData3DParams.mStartPt = 0;
loadData3DParams.kStartPt = 0;
loadData3DParams.strideW = 1;
loadData3DParams.strideH = 1;
loadData3DParams.filterW = 1;
loadData3DParams.filterSizeW = (1 >> 8) & 255;
loadData3DParams.filterH = 1;
loadData3DParams.filterSizeH = (1 >> 8) & 255;
loadData3DParams.dilationFilterW = 1;
loadData3DParams.dilationFilterH = 1;
loadData3DParams.enTranspose = 0;
loadData3DParams.fMatrixCtrl = 0;
LoadData<float, LOAD3DV2_CONFIG>(l0a_[(l0aLoopIdx_ % L0AB_BUF_NUM) * L0AB_BUF_OFFSET],
l1a_[(l1LoopIdx % L1_BUF_NUM) * L1_BUF_OFFSET +
CeilAlign(mmParams.curML1, static_cast<uint64_t>(BLOCK_CUBE)) * kL1Offset],
loadData3DParams);
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::LoadAToL0B(
uint64_t kL1Offset, uint64_t kL0Size, uint64_t l1LoopIdx, const MmParams &mmParams)
{
// mk nz -> m,k zz
for (uint64_t mL0Offset = 0; mL0Offset < mmParams.curML1; mL0Offset += BLOCK_CUBE) {
LoadData2DParams l1ToL0bParams;
l1ToL0bParams.startIndex = 0;
l1ToL0bParams.repeatTimes = CeilDiv(kL0Size, (uint64_t)BLOCK_CUBE >> 1);
l1ToL0bParams.srcStride = CeilDiv(mmParams.curML1, (uint64_t)BLOCK_CUBE);
l1ToL0bParams.dstGap = 0;
LoadData(l0b_[(l0bLoopIdx_ % L0AB_BUF_NUM) * L0AB_BUF_OFFSET +
mL0Offset * CeilAlign(kL0Size, (uint64_t)BLOCK_CUBE >> 1)],
l1a_[(l1LoopIdx % L1_BUF_NUM) * L1_BUF_OFFSET +
CeilAlign(mmParams.curML1, static_cast<uint64_t>(BLOCK_CUBE)) * kL1Offset +
mL0Offset * (uint64_t)(BLOCK_CUBE >> 1)],
l1ToL0bParams);
}
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::MmadA2(
uint64_t kGmOffset, uint64_t kL0Size, bool isLastK, const MmParams &mmParams)
{
SetFlag<HardEvent::MTE1_M>(MTE1_M_EVENT);
WaitFlag<HardEvent::MTE1_M>(MTE1_M_EVENT);
for (uint64_t mL0Offset = 0; mL0Offset < mmParams.curML1; mL0Offset += BLOCK_CUBE) {
MmadParams mmadParams;
mmadParams.m = BLOCK_CUBE;
mmadParams.n = BLOCK_CUBE;
mmadParams.k = kL0Size;
mmadParams.cmatrixInitVal = mmParams.isFirstK && kGmOffset == 0;
mmadParams.cmatrixSource = false;
mmadParams.unitFlag = isLastK ? UNIT_FLAG_ENABLE_AUTO_CLOSE : UNIT_FLAG_ENABLE;
// mk zz @ mk zz
Mmad(l0c_[(l0cLoopIdx_ % L0C_BUF_NUM) * L0C_BUF_OFFSET + mL0Offset * BLOCK_CUBE],
l0a_[(l0aLoopIdx_ % L0A_BUF_NUM) * L0AB_BUF_OFFSET + mL0Offset * CeilAlign(kL0Size, 8)],
l0b_[(l0bLoopIdx_ % L0AB_BUF_NUM) * L0AB_BUF_OFFSET + mL0Offset * CeilAlign(kL0Size, 8)],
mmadParams);
}
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::LoadBToL0B(
uint64_t kL1Offset, uint64_t kL0Size, uint64_t l1LoopIdx, const MmParams &mmParams)
{
LoadData2DParams l1ToL0bParams;
l1ToL0bParams.startIndex = 0;
l1ToL0bParams.repeatTimes =
CeilDiv(mmParams.curNL1, (uint64_t)BLOCK_CUBE) * CeilDiv(kL0Size, (uint64_t)BLOCK_CUBE >> 1);
l1ToL0bParams.srcStride = 1;
l1ToL0bParams.dstGap = 0;
// n,k nz -> n,k nz
LoadData(l0b_[(l0bLoopIdx_ % L0AB_BUF_NUM) * L0AB_BUF_OFFSET],
l1b_[(l1LoopIdx % L1_BUF_NUM) * L1_BUF_OFFSET +
CeilAlign(mmParams.curNL1, static_cast<uint64_t>(BLOCK_CUBE)) * kL1Offset],
l1ToL0bParams);
}
HC_PRE_CUBE_COMPUTE_TEMPLATE_PARAM
__aicore__ inline void HC_PRE_CUBE_COMPUTE_TEMPLATE_CLASS::MmadAB(
uint64_t kGmOffset, uint64_t kL0Size, bool isLastK, const MmParams &mmParams)
{
SetFlag<HardEvent::MTE1_M>(MTE1_M_EVENT);
WaitFlag<HardEvent::MTE1_M>(MTE1_M_EVENT);
AscendC::SetHF32Mode(1);
AscendC::SetHF32TransMode(1);
MmadParams mmadParams;
mmadParams.m = CeilAlign(mmParams.curML1, BLOCK_CUBE);
mmadParams.n = mmParams.curNL1;
mmadParams.k = kL0Size; // kl0Size
mmadParams.cmatrixInitVal = mmParams.isFirstK && kGmOffset == 0;
mmadParams.cmatrixSource = false;
mmadParams.unitFlag = isLastK ? UNIT_FLAG_ENABLE_AUTO_CLOSE : UNIT_FLAG_ENABLE;
// mk zz @ nk nz
Mmad(l0c_[(l0cLoopIdx_ % L0C_BUF_NUM) * L0C_BUF_OFFSET + L0C_A2_BUF_OFFSET],
l0a_[(l0aLoopIdx_ % L0A_BUF_NUM) * L0AB_BUF_OFFSET],
l0b_[(l0bLoopIdx_ % L0AB_BUF_NUM) * L0AB_BUF_OFFSET],
mmadParams);
AscendC::SetHF32Mode(0);
}
} // namespace HcPre
#endif // HC_PRE_CUBE_COMPUTE_H

View File

@@ -0,0 +1,241 @@
/**
* 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 hc_pre_cube_compute_arch35.h
* \brief
*/
#ifndef HC_PRE_CUBE_COMPUTE_ARCH35_H
#define HC_PRE_CUBE_COMPUTE_ARCH35_H
#include "kernel_operator.h"
#include "hc_pre_base_arch35.h"
namespace HcPreNs
{
using namespace AscendC;
constexpr static uint32_t FINAL_ACCUMULATION = 3;
constexpr static uint32_t NON_FINAL_ACCUMULATION = 2;
// constexpr static int32_t C0_SIZE = AscendC::AuxGetC0Size<float>();
constexpr static bool splitM_ = true;
constexpr static uint64_t SPLIT_M_ALIGN = 2;
static constexpr uint64_t L1_ALLOC_SIZE = 512 * 1024;
static constexpr uint64_t L1_BUF_NUM = 2;
static constexpr uint64_t L1_BUF_OFFSET = 128 * 256;
constexpr static int SYNC_MODE4 = 4;
class HcPreCubeCompute
{
public:
uint64_t m_{0};
uint64_t n_{0};
uint64_t k_{0};
uint64_t baseM_{0};
uint64_t baseN_{0};
uint64_t baseK_{0};
uint64_t kL1_{0};
public:
AscendC::LocalTensor<float> aL0Ping_;
AscendC::LocalTensor<float> aL0Pong_;
AscendC::LocalTensor<float> bL0Ping_;
AscendC::LocalTensor<float> bL0Pong_;
AscendC::LocalTensor<float> cL0Ping_;
AscendC::LocalTensor<float> cL0Pong_;
uint8_t bL1BufferID_{0};
uint8_t l0PingPongID_{0};
uint8_t crossPingPongID_{0};
uint8_t cl0PingPongID_{0};
__aicore__ inline HcPreCubeCompute()
{
}
__aicore__ inline uint8_t GetBL1BufferId()
{
return bL1BufferID_;
}
__aicore__ inline void Init()
{
// L0 空间的分配
uint32_t aL0OneBuffer = 256 * 32;
uint32_t bL0OneBuffer = 256 * 32;
uint32_t cL0OneBuffer = 256 * 128;
aL0Ping_ = AscendC::LocalTensor<float>(AscendC::TPosition::A2, 0, aL0OneBuffer);
aL0Pong_ = AscendC::LocalTensor<float>(AscendC::TPosition::A2, aL0OneBuffer * sizeof(float), aL0OneBuffer);
bL0Ping_ = AscendC::LocalTensor<float>(AscendC::TPosition::B2, 0, bL0OneBuffer);
bL0Pong_ = AscendC::LocalTensor<float>(AscendC::TPosition::B2, bL0OneBuffer * sizeof(float), bL0OneBuffer);
cL0Ping_ = AscendC::LocalTensor<float>(AscendC::TPosition::CO1, 0, cL0OneBuffer);
cL0Pong_ = AscendC::LocalTensor<float>(AscendC::TPosition::CO1, cL0OneBuffer * sizeof(float), cL0OneBuffer);
// 同步
// B 的 gm2L1 的 pingpong id
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(0);
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(1);
// l12l0a & l12l0b 的 pingpong id (共用)
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(3);
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(4);
// hf32 compute
AscendC::SetHF32Mode(1);
AscendC::SetHF32TransMode(1);
}
__aicore__ inline void CopyInB1Nd2Nz(uint64_t k, uint64_t currentK, uint64_t baseN, const AscendC::GlobalTensor<float> &bGlobal,
const AscendC::LocalTensor<float> &bl1Local)
{
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(bL1BufferID_);
k_ = k;
kL1_ = currentK;
baseN_ = baseN;
AscendC::Nd2NzParams nd2nzParam;
nd2nzParam.ndNum = 1;
nd2nzParam.srcNdMatrixStride = 1;
nd2nzParam.dstNzNStride = 1;
nd2nzParam.dstNzMatrixStride = 1;
nd2nzParam.nValue = baseN_;
nd2nzParam.dValue = kL1_;
nd2nzParam.srcDValue = k_;
nd2nzParam.dstNzC0Stride = (baseN_ + AscendC::BLOCK_CUBE - 1) / AscendC::BLOCK_CUBE * AscendC::BLOCK_CUBE;
AscendC::DataCopy(bl1Local, bGlobal, nd2nzParam);
AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE1>(bL1BufferID_);
AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE1>(bL1BufferID_);
}
// note: baseM * baseK must <= 256 * 32; baseN * baseK must <= 256 * 32; baseM * baseN must <= 256 * 128
__aicore__ inline void Process(uint64_t m, uint64_t n, uint64_t baseM, uint64_t baseK,
bool isFirstKL1, bool isLastKL1, const AscendC::LocalTensor<float> &al1Local,
const AscendC::LocalTensor<float> &bl1Local)
{
m_ = m;
n_ = n;
baseM_ = baseM;
baseK_ = baseK;
uint64_t kL1Offset = 0;
for (uint64_t kb = 0; kb < kL1_; kb += baseK_)
{
bool isLastKL0 = (kb + baseK_) >= kL1_;
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0PingPongID_ + 3);
CopyInA2(kb, kL1Offset, al1Local);
CopyInB2(kb, kL1Offset, bl1Local);
AscendC::SetFlag<AscendC::HardEvent::MTE1_M>(l0PingPongID_);
AscendC::WaitFlag<AscendC::HardEvent::MTE1_M>(l0PingPongID_);
MmadBase(kb, isFirstKL1, isLastKL1, isLastKL0);
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0PingPongID_ + 3);
l0PingPongID_ = l0PingPongID_ ^ 1;
kL1Offset += baseK_;
}
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(bL1BufferID_);
bL1BufferID_ ^= 1;
}
__aicore__ inline void CopyInA2(uint64_t kOffset, uint64_t kAL1Offset, const AscendC::LocalTensor<float> &al1Local_)
{
uint64_t mAL1 = Align(baseM_, AscendC::BLOCK_CUBE);
uint64_t offsetAL1 = Align(kAL1Offset, C0_SIZE) * mAL1;
AscendC::LoadData2DParamsV2 loadData2dParams;
uint64_t currM = baseM_;
uint64_t currK = AscendC::Std::min(baseK_, kL1_ - kOffset);
loadData2dParams.mStartPosition = 0;
loadData2dParams.kStartPosition = 0;
loadData2dParams.mStep = CeilDiv(currM, AscendC::BLOCK_CUBE);
loadData2dParams.kStep = CeilDiv(currK, C0_SIZE);
loadData2dParams.srcStride = CeilDiv(currM, AscendC::BLOCK_CUBE);
loadData2dParams.dstStride = loadData2dParams.mStep;
loadData2dParams.ifTranspose = false;
AscendC::LoadData(l0PingPongID_ == 0 ? aL0Ping_ : aL0Pong_, al1Local_[offsetAL1], loadData2dParams);
}
__aicore__ inline void CopyInB2(uint64_t kOffset, uint64_t kBL1Offset, const AscendC::LocalTensor<float> &bl1Local_)
{
uint64_t nBL1 = Align(baseN_, AscendC::BLOCK_CUBE);
uint64_t offsetBL1 = Align(kBL1Offset, C0_SIZE) * nBL1;
AscendC::LoadData2DParamsV2 loadData2dParams;
uint64_t currN = baseN_;
uint64_t currK = AscendC::Std::min(baseK_, kL1_ - kOffset);
loadData2dParams.mStartPosition = 0;
loadData2dParams.kStartPosition = 0;
loadData2dParams.mStep = CeilDiv(currN, AscendC::BLOCK_CUBE);
loadData2dParams.kStep = CeilDiv(currK, C0_SIZE);
loadData2dParams.srcStride = CeilDiv(currN, AscendC::BLOCK_CUBE);
loadData2dParams.dstStride = loadData2dParams.mStep;
loadData2dParams.ifTranspose = false;
AscendC::LoadData(l0PingPongID_ == 0 ? bL0Ping_ : bL0Pong_, bl1Local_[offsetBL1], loadData2dParams);
}
__aicore__ inline void MmadBase(uint64_t kOffset, bool isFirstKL1, bool isLastKL1, bool isLastKL0)
{
uint32_t mmadK = AscendC::Std::min(baseK_, kL1_ - kOffset);
AscendC::MmadParams mmadParams;
mmadParams.m = baseM_;
mmadParams.n = baseN_;
mmadParams.k = mmadK;
mmadParams.disableGemv = true;
mmadParams.cmatrixInitVal = (isFirstKL1 && kOffset == 0); // kOffset == 0: isFirstKL0
mmadParams.cmatrixSource = false;
mmadParams.unitFlag = (isLastKL1 && isLastKL0) ? FINAL_ACCUMULATION : NON_FINAL_ACCUMULATION;
AscendC::Mmad(cl0PingPongID_ == 0 ? cL0Ping_ : cL0Pong_, l0PingPongID_ == 0 ? aL0Ping_ : aL0Pong_,
l0PingPongID_ == 0 ? bL0Ping_ : bL0Pong_, mmadParams);
}
// fixpipe CopyOut实现c01拷贝到UB
__aicore__ inline void CopyOut(const AscendC::LocalTensor<float>& dstLocal)
{
AscendC::FixpipeParamsC310<AscendC::CO2Layout::ROW_MAJOR> fixpipeParams; // ROW_MAJOR默认使能NZ2ND
uint64_t c0 = AscendC::AuxGetC0Size<float>();
fixpipeParams.nSize = Align(baseN_, c0);
fixpipeParams.mSize = splitM_ ? Align(baseM_, SPLIT_M_ALIGN) : baseM_; // 切m需要m是2对齐
fixpipeParams.dstStride = fixpipeParams.nSize;
fixpipeParams.srcStride = Align(baseM_, AscendC::BLOCK_CUBE); // 单位CO_SIZE (16*sizeof(C_T))
fixpipeParams.quantPre = QuantMode_t::NoQuant;
// fixpipeParams.quantPre = 0;
// set cvRatio=1:2 默认splitM
fixpipeParams.dualDstCtl = splitM_ ? static_cast<uint8_t>(AscendC::McgShfMode::DUAL_DST_SPLIT_M) : 0;
fixpipeParams.unitFlag = FINAL_ACCUMULATION; // 3 unitflag
fixpipeParams.params.ndNum = 1; // ndNum
fixpipeParams.params.srcNdStride = 1; // srcNdStride
fixpipeParams.params.dstNdStride = 1; // dstNdStride
AscendC::Fixpipe<float, float, AscendC::Impl::CFG_ROW_MAJOR_UB>(dstLocal, cl0PingPongID_ == 0 ? cL0Ping_ : cL0Pong_, fixpipeParams);
cl0PingPongID_ ^= 1;
}
// fixpipe CopyOut实现c01拷贝到GM
__aicore__ inline void CopyOut(const AscendC::GlobalTensor<float> &cGlobal)
{
AscendC::DataCopyCO12DstParams intriParams;
intriParams.nSize = baseN_;
intriParams.mSize = baseM_;
intriParams.dstStride = n_;
intriParams.srcStride = Align(baseM_, AscendC::BLOCK_CUBE);
// set mode according to dtype
intriParams.quantPre = QuantMode_t::NoQuant;
intriParams.nz2ndEn = true;
intriParams.unitFlag = FINAL_ACCUMULATION; // 3 unitflag
AscendC::SetFixpipeNz2ndFlag(1, 1, 1);
AscendC::DataCopy(cGlobal, cl0PingPongID_ == 0 ? cL0Ping_ : cL0Pong_, intriParams);
cl0PingPongID_ ^= 1;
}
__aicore__ inline void End()
{
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(0);
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(1);
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(3);
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(4);
AscendC::SetHF32Mode(0);
}
};
}
#endif

View File

@@ -0,0 +1,564 @@
/**
* 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 hc_pre_m_split_core.h
* \brief
*/
#ifndef HC_PRE_M_K_SPLIT_A3_CORE_H
#define HC_PRE_M_K_SPLIT_A3_CORE_H
#include "kernel_operator.h"
#include "hc_pre_base.h"
#include "hc_pre_cube_compute.h"
namespace HcPre {
using namespace AscendC;
template <typename T>
class HcPreMembaseKSplitCorePart1 {
public:
__aicore__ inline HcPreMembaseKSplitCorePart1()
{
}
__aicore__ inline void Init(GM_ADDR x, GM_ADDR hcFn, GM_ADDR workspace,
const HcPreTilingData *tilingDataPtr, TPipe *pipePtr)
{
pipe = pipePtr;
tilingData = tilingDataPtr;
xGm.SetGlobalBuffer((__gm__ T *)x);
hcFnGm.SetGlobalBuffer((__gm__ float *)hcFn);
workspaceGm.SetGlobalBuffer((__gm__ float *)workspace);
uint64_t curVectorBlockIdx = GetBlockIdx();
uint64_t curCubeBlockIdx = curVectorBlockIdx;
if ASCEND_IS_AIV {
curCubeBlockIdx = curCubeBlockIdx / CV_RATIO;
}
xCastFp32BufSize_ = tilingData->mL1Size *
CeilAlign(tilingData->cvLoopKSize, MM_CACHE_LINE_BYTES / sizeof(float));
int64_t SingleCubeCoreXCastSize = xCastFp32BufSize_ * DOUBLE_BUFFER;
xCastFp32WsGm.SetGlobalBuffer((__gm__ float *)workspace + curCubeBlockIdx * SingleCubeCoreXCastSize);
uint64_t wsOffset = tilingData->cubeCoreNum * SingleCubeCoreXCastSize;
mmOutFp32WsGm.SetGlobalBuffer((__gm__ float *)workspace + wsOffset);
mmOuterInnerSize_ = CeilAlign(N_SIZE, MM_CACHE_LINE_BYTES / sizeof(float));
mmOutFp32BufSize_ = tilingData->bs * mmOuterInnerSize_;
wsOffset += CeilAlign(tilingData->cubeBlockDimK * mmOutFp32BufSize_ * sizeof(float),
static_cast<uint64_t>(WORKSPACE_ALIGN_SIZE)) / sizeof(float);
squareSumFp32WsGm.SetGlobalBuffer((__gm__ float *)workspace + wsOffset);
squareSumFp32BufSize_ = CeilAlign(tilingData->bs, BLOCK_CUBE) * BLOCK_CUBE;
if ASCEND_IS_AIC {
cubeCompute_.Init(xCastFp32WsGm, hcFnGm, pipePtr);
return;
}
// InQue
int64_t xQueNum = tilingData->stage1MFactor * RoundUp<T>(tilingData->cvLoopKSize);
pipe->InitBuffer(xQue, NUM_TWO, xQueNum * sizeof(T));
// OutQue
pipe->InitBuffer(mmInQue, NUM_TWO, xQueNum * sizeof(float));
}
__aicore__ inline void Process()
{
if ASCEND_IS_AIC{
// 初始设置
CrossCoreSetFlag<SYNC_MODE2, PIPE_FIX>(SYNC_AIC_TO_AIV_FLAG);
CrossCoreSetFlag<SYNC_MODE2, PIPE_FIX>(SYNC_AIC_TO_AIV_FLAG);
}
uint64_t curBlockIdx = GetBlockIdx();
uint64_t curVectorBlockIdx = curBlockIdx;
if ASCEND_IS_AIV {
curBlockIdx = curBlockIdx / NUM_TWO;
}
uint64_t mBlkDimIdx = curBlockIdx / tilingData->cubeBlockDimK;
uint64_t kBlkDimIdx = curBlockIdx % tilingData->cubeBlockDimK;
uint64_t kGmStartOffset = tilingData->multCoreSplitKSize * kBlkDimIdx;
uint64_t kGmEndOffset = kGmStartOffset + tilingData->multCoreSplitKSize;
if (kGmEndOffset > tilingData->k) {
kGmEndOffset = tilingData->k;
}
uint64_t mGmBaseOffset = mBlkDimIdx * tilingData->multCoreSplitMSize;
uint64_t mGmEndOffset = mGmBaseOffset + tilingData->multCoreSplitMSize;
if (mGmEndOffset > tilingData->bs) {
mGmEndOffset = tilingData->bs;
}
for (uint64_t mGmOffset = mGmBaseOffset; mGmOffset < mGmEndOffset;
mGmOffset += tilingData->mL1Size) {
uint64_t realMSize = mGmOffset + tilingData->mL1Size > mGmEndOffset ?
mGmEndOffset - mGmOffset : tilingData->mL1Size;
for (uint64_t kGmBaseOffset = kGmStartOffset; kGmBaseOffset < kGmEndOffset;
kGmBaseOffset += tilingData->cvLoopKSize) {
uint64_t realKGmSize = kGmBaseOffset + tilingData->cvLoopKSize > kGmEndOffset ?
kGmEndOffset - kGmBaseOffset : tilingData->cvLoopKSize;
if ASCEND_IS_AIC{
MmParams mmParams;
mmParams.curML1 = realMSize;
mmParams.curKL1 = tilingData->kL1Size;
mmParams.curNL1 = N_SIZE;
mmParams.singleCoreK = realKGmSize;
mmParams.xWsKSize = tilingData->cvLoopKSize;
mmParams.nOutSize = mmOuterInnerSize_;
mmParams.kGmBaseOffset = kGmBaseOffset;
mmParams.nGmSize = N_SIZE;
mmParams.kGmSize = tilingData->k;
mmParams.isLastK = kGmBaseOffset + tilingData->cvLoopKSize >= kGmEndOffset;
mmParams.isFirstK = kGmBaseOffset == kGmStartOffset;
cubeCompute_.WaitBL1Mte1ToMte2Flag();
cubeCompute_.CopyInB1(mGmOffset, kGmBaseOffset, realKGmSize, mmParams);
CrossCoreWaitFlag(SYNC_AIV_TO_AIC_FLAG);
cubeCompute_.ComputeDecode(xCastFp32WsGm[cvLoopIdx_ % DOUBLE_BUFFER * xCastFp32BufSize_],
squareSumFp32WsGm[kBlkDimIdx * squareSumFp32BufSize_ + mGmOffset * BLOCK_CUBE],
mmOutFp32WsGm[kBlkDimIdx * mmOutFp32BufSize_ + mGmOffset * mmOuterInnerSize_],
mmParams);
cubeCompute_.SetBL1Mte1ToMte2Flag();
CrossCoreSetFlag<SYNC_MODE2, PIPE_FIX>(SYNC_AIC_TO_AIV_FLAG);
} else {
CrossCoreWaitFlag(SYNC_AIC_TO_AIV_FLAG);
// vec.compute();
int64_t mVectorOffset = mGmOffset;
int64_t mVectorLength = (realMSize + 1) / NUM_TWO;
if ((curVectorBlockIdx % NUM_TWO) == 1) {
mVectorOffset = mGmOffset + (realMSize + 1) / NUM_TWO;
mVectorLength = realMSize - mVectorLength;
}
int64_t curUbLoops = (mVectorLength + tilingData->stage1MFactor - 1) / tilingData->stage1MFactor;
int64_t curUbMfactorTail = mVectorLength - ((curUbLoops - 1) * tilingData->stage1MFactor);
for (int64_t i = 0; i < curUbLoops; ++i) {
int64_t curUbMFactor = (i != (curUbLoops - 1)) ? tilingData->stage1MFactor : curUbMfactorTail;
xLocal = xQue.template AllocTensor<T>();
int64_t curGlobalxOffset = (mVectorOffset + i * tilingData->stage1MFactor) *
tilingData->k + kGmBaseOffset;
CopyIn(xGm[curGlobalxOffset], xLocal, curUbMFactor, realKGmSize, tilingData->k - realKGmSize);
xQue.template EnQue(xLocal);
xLocal = xQue.template DeQue<T>();
xCastLocal = mmInQue.AllocTensor<float>();
CastTwoDim(xCastLocal, xLocal, curUbMFactor, realKGmSize);
xQue.template FreeTensor(xLocal);
mmInQue.template EnQue(xCastLocal);
xCastLocal = mmInQue.template DeQue<float>();
int64_t cutMmInOffset = cvLoopIdx_ % DOUBLE_BUFFER * xCastFp32BufSize_ +
(i * tilingData->stage1MFactor + mVectorOffset - mGmOffset) *
tilingData->cvLoopKSize;
CopyOut(xCastLocal, xCastFp32WsGm[cutMmInOffset], curUbMFactor, realKGmSize,
tilingData->cvLoopKSize - realKGmSize);
mmInQue.FreeTensor(xCastLocal);
}
CrossCoreSetFlag<SYNC_MODE2, PIPE_MTE3>(SYNC_AIV_TO_AIC_FLAG);
}
cvLoopIdx_++;
}
}
if ASCEND_IS_AIC {
cubeCompute_.End();
} else {
CrossCoreWaitFlag(SYNC_AIC_TO_AIV_FLAG);
CrossCoreWaitFlag(SYNC_AIC_TO_AIV_FLAG);
}
SyncAll<false>(); // cv全部同步
}
private:
TPipe *pipe;
const HcPreTilingData *tilingData;
GlobalTensor<float> workspaceGm;
GlobalTensor<float> xCastFp32WsGm;
GlobalTensor<float> mmOutFp32WsGm;
GlobalTensor<float> squareSumFp32WsGm;
GlobalTensor<float> hcFnGm;
GlobalTensor<T> xGm;
TQue<QuePosition::VECIN, 1> xQue;
TQue<QuePosition::VECOUT, 1> mmInQue;
LocalTensor<T> xLocal;
LocalTensor<float> xCastLocal;
LocalTensor<float> yCastLocal;
LocalTensor<float> mmOutLocal;
HcCubeCompute<false> cubeCompute_;
static constexpr uint64_t SYNC_AIV_TO_AIC_FLAG = 8;
static constexpr uint64_t SYNC_AIC_TO_AIV_FLAG = 9;
static constexpr uint64_t SYNC_MODE2 = NUM_TWO;
uint64_t cvLoopIdx_ = 0;
uint64_t xCastFp32BufSize_;
uint64_t mmOuterInnerSize_;
uint64_t mmOutFp32BufSize_;
uint64_t squareSumFp32BufSize_;
};
template <typename T>
class HcPreMembaseKSplitCorePart2 {
public:
__aicore__ inline HcPreMembaseKSplitCorePart2()
{
}
__aicore__ inline void InitGlobalBuffers(GM_ADDR x, GM_ADDR hcScale,
GM_ADDR hcBase, GM_ADDR y, GM_ADDR post, GM_ADDR combFrag,
GM_ADDR workspace)
{
xGm.SetGlobalBuffer((__gm__ T *)x);
hcScaleGm.SetGlobalBuffer((__gm__ float *)hcScale);
hcBaseGm.SetGlobalBuffer((__gm__ float *)hcBase);
yGm.SetGlobalBuffer((__gm__ T *)y);
postGm.SetGlobalBuffer((__gm__ float *)post);
combFragGm.SetGlobalBuffer((__gm__ float *)combFrag);
workspaceGm.SetGlobalBuffer((__gm__ float *)workspace);
}
__aicore__ inline void InitQueBuffers(int64_t stage1UsedCoreNum,
int64_t xQueNum2)
{
int64_t mixesQue01Size = stage1UsedCoreNum * tilingData->stage2RowFactor *
tilingData->hcMultAlign * NUM_TWO * sizeof(float);
pipe->InitBuffer(mixesQue01, NUM_TWO, mixesQue01Size);
pipe->InitBuffer(mixesQue2, NUM_TWO,
stage1UsedCoreNum * tilingData->stage2RowFactor *
tilingData->hcMult * tilingData->hcMultAlign * sizeof(float));
pipe->InitBuffer(squareSumQue, NUM_TWO, stage1UsedCoreNum *
tilingData->stage2RowFactor * SQUARE_SUM_SIZE * sizeof(float));
pipe->InitBuffer(xQue, NUM_TWO, xQueNum2 * sizeof(T));
pipe->InitBuffer(squareSumQue, NUM_TWO, stage1UsedCoreNum *
tilingData->stage2RowFactor * SQUARE_SUM_SIZE * sizeof(float));
pipe->InitBuffer(yQue, NUM_TWO,
tilingData->stage2RowFactor * RoundUp<T>(tilingData->dFactor) * sizeof(T));
pipe->InitBuffer(postQue, NUM_TWO,
tilingData->stage2RowFactor * tilingData->hcMultAlign * sizeof(float));
pipe->InitBuffer(combFragQue, NUM_TWO,
tilingData->stage2RowFactor * tilingData->hcMult *
tilingData->hcMultAlign * sizeof(float));
}
__aicore__ inline void InitTBufBuffers(int64_t xQueNum2)
{
pipe->InitBuffer(hcBaseBuf0, tilingData->hcMultAlign * sizeof(float));
pipe->InitBuffer(hcBaseBuf1, tilingData->hcMultAlign * sizeof(float));
pipe->InitBuffer(hcBaseBuf2,
tilingData->hcMult * tilingData->hcMultAlign * sizeof(float));
pipe->InitBuffer(rowBrcbBuf0,
RoundUp<float>(tilingData->stage2RowFactor) * BLOCK_SIZE);
pipe->InitBuffer(hcBrcbBuf1,
RoundUp<float>(tilingData->stage2RowFactor *
tilingData->hcMultAlign) * BLOCK_SIZE);
pipe->InitBuffer(reduceBuf,
tilingData->stage2RowFactor * tilingData->hcMultAlign * sizeof(float));
pipe->InitBuffer(mixes01ReduceBuf, tilingData->stage2RowFactor *
tilingData->hcMultAlign * NUM_TWO * sizeof(float));
pipe->InitBuffer(mixes02ReduceBuf, tilingData->stage2RowFactor *
tilingData->hcMultAlign * tilingData->hcMult * sizeof(float));
pipe->InitBuffer(squareReduceBuf,
tilingData->stage2RowFactor * tilingData->hcMultAlign * sizeof(float));
pipe->InitBuffer(xCastBuf, xQueNum2 * sizeof(float));
pipe->InitBuffer(yCastBuf, tilingData->stage2RowFactor *
RoundUp<T>(tilingData->dFactor) * sizeof(float));
pipe->InitBuffer(rsqrtBuf,
RoundUp<float>(tilingData->stage2RowFactor) * sizeof(float));
pipe->InitBuffer(maskPatternBuf,
RoundUp<uint32_t>(MASK_PATTERN_BASE_SIZE * MASK_PATTERN_REPEAT_SIZE) *
sizeof(uint32_t));
}
__aicore__ inline void GetLocalTensors()
{
hcBase0Local = hcBaseBuf0.Get<float>();
hcBase1Local = hcBaseBuf1.Get<float>();
hcBase2Local = hcBaseBuf2.Get<float>();
rowBrcbLocal0 = rowBrcbBuf0.Get<float>();
hcBrcbLocal1 = hcBrcbBuf1.Get<float>();
reduceLocal = reduceBuf.Get<float>();
squareReduceLocal = squareReduceBuf.Get<float>();
mixes01ReduceLocal = mixes01ReduceBuf.Get<float>();
mixes02ReduceLocal = mixes02ReduceBuf.Get<float>();
xCastLocal = xCastBuf.Get<float>();
yCastLocal = yCastBuf.Get<float>();
rsqrtLocal = rsqrtBuf.Get<float>();
maskPatternLocal = maskPatternBuf.Get<uint32_t>();
SetGatherMaskPattern(maskPatternLocal);
}
__aicore__ inline void Init(GM_ADDR x, GM_ADDR hcScale, GM_ADDR hcBase,
GM_ADDR y, GM_ADDR post, GM_ADDR combFrag, GM_ADDR workspace,
const HcPreTilingData *tilingDataPtr, TPipe *pipePtr)
{
pipe = pipePtr;
tilingData = tilingDataPtr;
InitGlobalBuffers(x, hcScale, hcBase, y, post, combFrag, workspace);
int64_t stage1UsedCoreNum = tilingData->cubeBlockDimK;
int64_t xQueNum2 = tilingData->stage2RowFactor * tilingData->hcMult *
RoundUp<T>(tilingData->dFactor);
InitQueBuffers(stage1UsedCoreNum, xQueNum2);
InitTBufBuffers(xQueNum2);
GetLocalTensors();
}
__aicore__ inline void Process()
{
if ASCEND_IS_AIV {
int64_t stage1UsedCoreNum = tilingData->cubeBlockDimK;// todo check 此处不应该写死32
int64_t stage2BlockIdx = GetBlockIdx();
int64_t stage2UsedCoreNum = tilingData->secondUsedCoreNum;
if (stage2BlockIdx >= stage2UsedCoreNum) {
return;
}
int64_t mmLastAxisSize = CeilAlign(tilingData->hcMix, MM_CACHE_LINE_BYTES / sizeof(float));
int64_t xCastFp32BufSize = tilingData->mL1Size *
CeilAlign(tilingData->cvLoopKSize, MM_CACHE_LINE_BYTES / sizeof(float));
int64_t workspaceSize1 = tilingData->cubeCoreNum * DOUBLE_BUFFER * xCastFp32BufSize;
int64_t workspaceSize2 = CeilAlign(stage1UsedCoreNum * tilingData->bs *
mmLastAxisSize * sizeof(float), WORKSPACE_ALIGN_SIZE) / sizeof(float);
CopyIn(hcBaseGm, hcBase0Local, 1, tilingData->hcMult);
CopyIn(hcBaseGm[tilingData->hcMult], hcBase1Local, 1, tilingData->hcMult);
CopyIn(hcBaseGm[tilingData->hcMult * NUM_TWO], hcBase2Local, tilingData->hcMult, tilingData->hcMult);
event_t eventId = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V));
SetFlag<HardEvent::MTE2_V>(eventId);
WaitFlag<HardEvent::MTE2_V>(eventId);
int64_t rowOuterLoop =
(stage2BlockIdx == stage2UsedCoreNum - 1) ?
tilingData->rowLoopOfTailBlock : tilingData->rowLoopOfFormerBlock;
int64_t tailRowFactor = (stage2BlockIdx == stage2UsedCoreNum - 1) ? tilingData->tailRowFactorOfTailBlock :
tilingData->tailRowFactorOfFormerBlock;
int64_t xGmBlockBaseOffsetPart2 = stage2BlockIdx *
tilingData->rowOfFormerBlock * tilingData->hcMult * tilingData->d;
for (int64_t rowOuterIdx = 0; rowOuterIdx < rowOuterLoop; rowOuterIdx++) {
int64_t xGmBsBaseOffsetPart2 = rowOuterIdx * tilingData->stage2RowFactor *
tilingData->hcMult * tilingData->d;
int64_t curRowFactor = (rowOuterIdx == rowOuterLoop - 1) ? tailRowFactor : tilingData->stage2RowFactor;
squareSumOutLocal = squareSumQue.AllocTensor<float>();
//todo
CopyIn(workspaceGm[workspaceSize1 + workspaceSize2 +
stage2BlockIdx * tilingData->rowOfFormerBlock * SQUARE_SUM_SIZE +
rowOuterIdx * tilingData->stage2RowFactor * SQUARE_SUM_SIZE],
squareSumOutLocal, stage1UsedCoreNum, curRowFactor * SQUARE_SUM_SIZE,
CeilAlign(tilingData->bs, SQUARE_SUM_SIZE) * SQUARE_SUM_SIZE -
curRowFactor * SQUARE_SUM_SIZE);
squareSumQue.EnQue(squareSumOutLocal);
squareSumOutLocal = squareSumQue.DeQue<float>();
ReduceSumARAPerf(squareReduceLocal, squareSumOutLocal, 1, stage1UsedCoreNum,
curRowFactor * SQUARE_SUM_SIZE);
int64_t curBsIdxForAll = (stage2BlockIdx * tilingData->rowLoopOfFormerBlock +
rowOuterIdx) * tilingData->stage2RowFactor;
GatherMaskByDiagonal(rsqrtLocal, squareReduceLocal,
maskPatternLocal[(curBsIdxForAll % SQUARE_SUM_SIZE) * 8], curRowFactor);
float coeff = 1.0f / static_cast<float>(tilingData->k);
Muls(rsqrtLocal, rsqrtLocal, coeff, curRowFactor);
PipeBarrier<PIPE_V>();
Adds(rsqrtLocal, rsqrtLocal, tilingData->normEps, curRowFactor);
PipeBarrier<PIPE_V>();
Sqrt(rsqrtLocal, rsqrtLocal, curRowFactor);
Duplicate(rowBrcbLocal0, static_cast<float>(1.0f), curRowFactor);
PipeBarrier<PIPE_V>();
Div(rsqrtLocal, rowBrcbLocal0, rsqrtLocal, curRowFactor);
mixes01Local = mixesQue01.AllocTensor<float>();
uint64_t mixBaseOffset = workspaceSize1 +
stage2BlockIdx * tilingData->rowOfFormerBlock *
CeilAlign(tilingData->hcMix, WORKSPACE_ALIGN_SIZE / sizeof(float)) +
rowOuterIdx * tilingData->stage2RowFactor *
CeilAlign(tilingData->hcMix, WORKSPACE_ALIGN_SIZE / sizeof(float));
CopyInWithOuterFor(workspaceGm[mixBaseOffset], mixes01Local, stage1UsedCoreNum,
curRowFactor, tilingData->hcMult, tilingData->bs,
CeilAlign(tilingData->hcMix, WORKSPACE_ALIGN_SIZE / sizeof(float)));
CopyInWithOuterFor(workspaceGm[mixBaseOffset + tilingData->hcMult],
mixes01Local[stage1UsedCoreNum * tilingData->stage2RowFactor *
tilingData->hcMultAlign], stage1UsedCoreNum, curRowFactor, tilingData->hcMult,
tilingData->bs,
CeilAlign(tilingData->hcMix, WORKSPACE_ALIGN_SIZE / sizeof(float)));
mixesQue01.EnQue(mixes01Local);
mixes01Local = mixesQue01.DeQue<float>();
ReduceSumARAPerf(mixes01ReduceLocal, mixes01Local, NUM_TWO, stage1UsedCoreNum,
curRowFactor * tilingData->hcMultAlign);
ProcessPre(mixes01ReduceLocal, mixes01ReduceLocal, hcBase0Local, rsqrtLocal,
rowBrcbLocal0, hcBrcbLocal1, hcScaleGm.GetValue(0), tilingData->hcEps,
curRowFactor, tilingData->hcMult);
for (int64_t dLoopIdx = 0; dLoopIdx < tilingData->dLoop; dLoopIdx++) {
int64_t curDFactor =
(dLoopIdx == tilingData->dLoop - 1) ? tilingData->tailDFactor : tilingData->dFactor;
xLocal = xQue.template AllocTensor<T>();
CopyIn(xGm[xGmBlockBaseOffsetPart2 + xGmBsBaseOffsetPart2 +
dLoopIdx * tilingData->dFactor], xLocal,
tilingData->stage2RowFactor * tilingData->hcMult, curDFactor,
tilingData->d - curDFactor);
xQue.template EnQue(xLocal);
xLocal = xQue.template DeQue<T>();
yLocal = yQue.template AllocTensor<T>();
ProcessY(yLocal, xLocal, mixes01ReduceLocal, hcBrcbLocal1, xCastLocal, yCastLocal, curRowFactor,
tilingData->hcMult, curDFactor);
xQue.template FreeTensor(xLocal);
yQue.template EnQue(yLocal);
yLocal = yQue.template DeQue<T>();
CopyOut(yLocal,
yGm[stage2BlockIdx * tilingData->rowOfFormerBlock * tilingData->d +
rowOuterIdx * tilingData->stage2RowFactor * tilingData->d +
dLoopIdx * tilingData->dFactor],
curRowFactor, curDFactor, tilingData->d - curDFactor);
yQue.template FreeTensor(yLocal);
}
// post
postLocal = postQue.AllocTensor<float>();
ProcessPost(postLocal,
mixes01ReduceLocal[tilingData->stage2RowFactor * tilingData->hcMultAlign],
hcBase1Local, rsqrtLocal, rowBrcbLocal0, hcBrcbLocal1, hcScaleGm.GetValue(1),
curRowFactor, tilingData->hcMult);
mixesQue01.template FreeTensor(mixes01Local);
postQue.EnQue(postLocal);
postLocal = postQue.DeQue<float>();
CopyOut(postLocal,
postGm[stage2BlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMult +
rowOuterIdx * tilingData->stage2RowFactor * tilingData->hcMult],
curRowFactor, tilingData->hcMult);
postQue.FreeTensor(postLocal);
// combFrag
mixes2Local = mixesQue2.AllocTensor<float>();
for (int64_t i = 0; i < stage1UsedCoreNum; ++i) {
for (int64_t j = 0; j < curRowFactor; ++j) {
CopyIn(workspaceGm[workspaceSize1 + i * tilingData->bs * mmLastAxisSize +
j * mmLastAxisSize +
stage2BlockIdx * tilingData->rowOfFormerBlock * mmLastAxisSize +
rowOuterIdx * tilingData->stage2RowFactor * mmLastAxisSize +
tilingData->hcMult * NUM_TWO],
mixes2Local[(i * curRowFactor + j) * tilingData->hcMult *
tilingData->hcMultAlign], tilingData->hcMult, tilingData->hcMult);
}
}
mixesQue2.EnQue(mixes2Local);
mixes2Local = mixesQue2.DeQue<float>();
ReduceSumARAPerf(mixes02ReduceLocal, mixes2Local, 1, stage1UsedCoreNum,
curRowFactor * tilingData->hcMult * tilingData->hcMultAlign);
combFragLocal = combFragQue.AllocTensor<float>();
MulABLastDimBrcInline<float, false>(mixes02ReduceLocal, mixes02ReduceLocal,
rsqrtLocal, rowBrcbLocal0, curRowFactor,
tilingData->hcMult * tilingData->hcMultAlign);
Muls(mixes02ReduceLocal, mixes02ReduceLocal, hcScaleGm.GetValue(NUM_TWO),
curRowFactor * tilingData->hcMult * tilingData->hcMultAlign);
PipeBarrier<PIPE_V>();
AddBAFirstDimBrcInline<float>(mixes02ReduceLocal, mixes02ReduceLocal, hcBase2Local, curRowFactor,
tilingData->hcMult * tilingData->hcMultAlign);
SoftmaxFP32Perf(mixes02ReduceLocal, mixes02ReduceLocal, reduceLocal, hcBrcbLocal1,
curRowFactor * tilingData->hcMult, tilingData->hcMult, tilingData->hcEps);
ReduceSumARAPerf(reduceLocal, mixes02ReduceLocal, curRowFactor, tilingData->hcMult, tilingData->hcMult);
Adds(reduceLocal, reduceLocal, tilingData->hcEps, curRowFactor * tilingData->hcMult);
PipeBarrier<PIPE_V>();
DivABABrcInline(combFragLocal, mixes02ReduceLocal, reduceLocal, curRowFactor, tilingData->hcMult,
tilingData->hcMult);
for (int64_t iter = 0; iter < tilingData->iterTimes - 1; iter++) {
LastDimReduceSumPerf(reduceLocal, combFragLocal,
curRowFactor * tilingData->hcMult, tilingData->hcMult);
Adds(reduceLocal, reduceLocal, tilingData->hcEps,
curRowFactor * tilingData->hcMult);
PipeBarrier<PIPE_V>();
DivABLastDimBrcInline<float, true>(combFragLocal, combFragLocal,
reduceLocal, hcBrcbLocal1, curRowFactor * tilingData->hcMult,
tilingData->hcMult);
ReduceSumARAPerf(reduceLocal, combFragLocal, curRowFactor,
tilingData->hcMult, tilingData->hcMult);
Adds(reduceLocal, reduceLocal, tilingData->hcEps,
curRowFactor * tilingData->hcMult);
PipeBarrier<PIPE_V>();
DivABABrcInline(combFragLocal, combFragLocal, reduceLocal,
curRowFactor, tilingData->hcMult, tilingData->hcMult);
}
mixesQue2.FreeTensor(mixes2Local);
squareSumQue.template FreeTensor(squareSumOutLocal);
combFragQue.EnQue(combFragLocal);
combFragLocal = combFragQue.DeQue<float>();
CopyOut(combFragLocal,
combFragGm[stage2BlockIdx * tilingData->rowOfFormerBlock *
tilingData->hcMult * tilingData->hcMult +
rowOuterIdx * tilingData->stage2RowFactor * tilingData->hcMult *
tilingData->hcMult], curRowFactor * tilingData->hcMult,
tilingData->hcMult);
combFragQue.FreeTensor(combFragLocal);
}
}
}
private:
TPipe *pipe;
const HcPreTilingData *tilingData;
GlobalTensor<float> mixesGm;
GlobalTensor<float> rsqrtGm;
GlobalTensor<float> hcScaleGm;
GlobalTensor<float> hcBaseGm;
GlobalTensor<float> workspaceGm;
GlobalTensor<T> xGm;
GlobalTensor<T> yGm;
GlobalTensor<float> postGm;
GlobalTensor<float> combFragGm;
TQue<QuePosition::VECIN, 1> mixesQue01;
TQue<QuePosition::VECIN, 1> mixesQue2;
TQue<QuePosition::VECIN, 1> xQue;
TQue<QuePosition::VECOUT, 1> yQue;
TQue<QuePosition::VECOUT, 1> postQue;
TQue<QuePosition::VECOUT, 1> combFragQue;
TQue<QuePosition::VECIN, 1> squareSumQue;
TBuf<QuePosition::VECCALC> hcBaseBuf0;
TBuf<QuePosition::VECCALC> hcBaseBuf1;
TBuf<QuePosition::VECCALC> hcBaseBuf2;
TBuf<QuePosition::VECCALC> rowBrcbBuf0;
TBuf<QuePosition::VECCALC> hcBrcbBuf1;
TBuf<QuePosition::VECCALC> reduceBuf;
TBuf<QuePosition::VECCALC> rsqrtBuf;
TBuf<QuePosition::VECCALC> squareReduceBuf;
TBuf<QuePosition::VECCALC> mixes01ReduceBuf;
TBuf<QuePosition::VECCALC> mixes02ReduceBuf;
TBuf<QuePosition::VECCALC> xCastBuf;
TBuf<QuePosition::VECCALC> yCastBuf;
TBuf<QuePosition::VECCALC> maskPatternBuf;
LocalTensor<float> mixes01Local;
LocalTensor<float> mixes2Local;
LocalTensor<float> rsqrtLocal;
LocalTensor<T> xLocal;
LocalTensor<T> yLocal;
LocalTensor<float> postLocal;
LocalTensor<float> combFragLocal;
LocalTensor<float> hcBase0Local;
LocalTensor<float> hcBase1Local;
LocalTensor<float> hcBase2Local;
LocalTensor<float> rowBrcbLocal0;
LocalTensor<float> hcBrcbLocal1;
LocalTensor<float> reduceLocal;
LocalTensor<float> squareReduceLocal;
LocalTensor<float> mixes01ReduceLocal;
LocalTensor<float> mixes02ReduceLocal;
LocalTensor<float> xCastLocal;
LocalTensor<float> yCastLocal;
LocalTensor<float> squareSumOutLocal;
LocalTensor<uint32_t> maskPatternLocal;
};
} // namespace HcPreSinkhorn
#endif

View File

@@ -0,0 +1,428 @@
/**
* 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 hc_pre_m_k_split_core.h
* \brief
*/
#ifndef HC_PRE_M_K_SPLIT_CORE_ARACH35_H
#define HC_PRE_M_K_SPLIT_CORE_ARACH35_H
#include "kernel_operator.h"
#include "hc_pre_base_arch35.h"
#include "hc_pre_cube_compute_arch35.h"
namespace HcPreNs {
using namespace AscendC;
template <typename T>
class HcPreMKSplitCorePart1 {
public:
__aicore__ inline HcPreMKSplitCorePart1()
{}
__aicore__ inline void Init(
GM_ADDR x, GM_ADDR hcFn, GM_ADDR workspace, const HcPreTilingData* tilingDataPtr, TPipe* pipePtr)
{
pipe = pipePtr;
tilingData = tilingDataPtr;
xGm.SetGlobalBuffer((__gm__ T*)x);
hcFnGm.SetGlobalBuffer((__gm__ float*)hcFn);
mmGm.SetGlobalBuffer((__gm__ float*)workspace);
rmsGm.SetGlobalBuffer((__gm__ float*)workspace + tilingData->kBlockFactor * tilingData->bs * tilingData->hcMix);
TBuf<TPosition::A1> l1Buffer;
pipe->InitBuffer(l1Buffer, L1_ALLOC_SIZE);
xL1_ = l1Buffer.Get<float>();
wL1_ = l1Buffer.Get<float>()[L1_BUF_NUM * L1_BUF_OFFSET];
// InQue
pipe->InitBuffer(xQue, 2, tilingData->mUbSize * RoundUp<T>(tilingData->kUbSize) * sizeof(T));
// OutQue
pipe->InitBuffer(rmsQue, 2, RoundUp<float>(tilingData->mUbSize) * sizeof(float));
// Calc Buf
pipe->InitBuffer(castBuf, tilingData->mUbSize * (RoundUp<float>(tilingData->kUbSize) * sizeof(float) + BLOCK_SIZE));
pipe->InitBuffer(nd2NzBuf, CeilAlign(tilingData->mUbSize, C0_SIZE) * RoundUp<float>(tilingData->kUbSize) * sizeof(float) * DOUBLE_BUFFER);
if ASCEND_IS_AIC {
mmService_.Init();
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIC_AIV_FLAG);
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIC_AIV_FLAG + FLAG_ID_MAX);
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIC_AIV_FLAG);
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIC_AIV_FLAG + FLAG_ID_MAX);
}
xCastLocal = castBuf.Get<float>();
xNd2NzLocal = nd2NzBuf.Get<float>();
}
__aicore__ inline void Process()
{
int64_t curBlockIdx = GetBlockIdx();
int64_t totalBlockNum = GetBlockNum();
uint64_t mBlkDimIdx = curBlockIdx / tilingData->cubeBlockDimK;
uint64_t kBlkDimIdx = curBlockIdx % tilingData->cubeBlockDimK;
// todo 移到tiling计算
uint64_t mCnt = CeilDiv(tilingData->bs, tilingData->mL1Size);
uint64_t singleCoreMaxRound = CeilDiv(mCnt, tilingData->cubeBlockDimM);
uint64_t mainCoreCount = mCnt % tilingData->cubeBlockDimM;
uint64_t singleCoreRound = (mainCoreCount == 0 || curBlockIdx < mainCoreCount) ? singleCoreMaxRound : singleCoreMaxRound - 1;
uint64_t mGmOffset = 0;
uint64_t nd2NzBufSize = CeilAlign(tilingData->mUbSize, C0_SIZE) * RoundUp<float>(tilingData->kUbSize);
if ASCEND_IS_AIC {
mGmOffset = (curBlockIdx / tilingData->cubeBlockDimK) * singleCoreMaxRound * tilingData->mL1Size;
} else {
mGmOffset = ((curBlockIdx / 2) / tilingData->cubeBlockDimK) * singleCoreMaxRound * tilingData->mL1Size;
}
int64_t xGmBaseOffset = 0;
int64_t rmsGmBaseOffset = 0;
if ASCEND_IS_AIV {
int64_t aivCurBlockIdx = GetBlockIdx();
xGmBaseOffset = ((aivCurBlockIdx / 2) / tilingData->cubeBlockDimK) * singleCoreMaxRound * tilingData->mL1Size * tilingData->hcMult * tilingData->d +
((aivCurBlockIdx / 2) % tilingData->cubeBlockDimK) * tilingData->multCoreSplitKSize;
rmsGmBaseOffset = ((aivCurBlockIdx / 2) / tilingData->cubeBlockDimK) * singleCoreMaxRound * tilingData->mL1Size;
}
// todo 移到tiling计算
int64_t xSplitOffset = 0;
int64_t rmsSplitOffset = 0;
// m轴切分 按照0 0 1 1..分核
int64_t bufferIdx = 0;
int64_t curAicBlockIdx = 0;
if ASCEND_IS_AIC {
curAicBlockIdx = curBlockIdx;
} else {
curAicBlockIdx = curBlockIdx / 2;
}
if (curAicBlockIdx < tilingData->cubeBlockDimK * tilingData->cubeBlockDimM) {
if ASCEND_IS_AIV {
SetFlag<HardEvent::MTE3_V>(static_cast<event_t>(0));
SetFlag<HardEvent::MTE3_V>(static_cast<event_t>(1));
}
for (uint64_t roundIdx = 0; roundIdx < singleCoreRound; mGmOffset += tilingData->mL1Size, ++roundIdx)
{
uint64_t mL1RealSize = AscendC::Std::min(tilingData->bs - mGmOffset, (uint64_t)tilingData->mL1Size);
uint64_t kGmStartOffset = 0;
uint64_t kGmEndOffset = AscendC::Std::min(tilingData->multCoreSplitKSize,
tilingData->k - (curAicBlockIdx % tilingData->cubeBlockDimK) * tilingData->multCoreSplitKSize);
if ASCEND_IS_AIV {
if (GetBlockIdx() % 2 != 0) {
xSplitOffset = (mL1RealSize / 2) * tilingData->hcMult * tilingData->d;
rmsSplitOffset = mL1RealSize / 2;
}
rmsNormLocal = rmsQue.template AllocTensor<float>();
}
int64_t curRowFactor = 0;
for (int64_t kGmOffset = kGmStartOffset; kGmOffset < kGmEndOffset; kGmOffset += tilingData->kL1Size) {
uint64_t kL1RealSize = AscendC::Std::min(kGmEndOffset - kGmOffset, (uint64_t)tilingData->kL1Size);
if ASCEND_IS_AIC {
bool isFirstKL1 = kGmOffset == kGmStartOffset;
bool isLastKL1 = (kGmOffset + tilingData->kL1Size) >= kGmEndOffset;
mmService_.CopyInB1Nd2Nz(tilingData->hcMult * tilingData->d, kL1RealSize,
tilingData->hcMix, hcFnGm[kGmOffset + kBlkDimIdx * tilingData->multCoreSplitKSize],
wL1_[mmService_.GetBL1BufferId() * L1_BUF_OFFSET]);
CrossCoreWaitFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIV_AIC_FLAG + FLAG_ID_MAX);
CrossCoreWaitFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIV_AIC_FLAG);
uint64_t mL1AlignSize = Align(mL1RealSize, AscendC::BLOCK_CUBE);
uint64_t nL1AlignSize = Align((uint64_t)tilingData->hcMix, AscendC::BLOCK_CUBE);
mmService_.Process(tilingData->bs, tilingData->hcMix, mL1RealSize, (256 / AscendC::Std::max(mL1AlignSize, nL1AlignSize)) * 32,
isFirstKL1, isLastKL1, xL1_[aL1BufferID_ * L1_BUF_OFFSET], wL1_[mmService_.GetBL1BufferId() * L1_BUF_OFFSET]);
if (isLastKL1) {
mmService_.CopyOut(mmGm[mBlkDimIdx * tilingData->mL1Size * singleCoreMaxRound * tilingData->hcMix + kBlkDimIdx * tilingData->bs * tilingData->hcMix + roundIdx * tilingData->mL1Size]);
}
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIC_AIV_FLAG); // 写出ub搬出,cv流水同步比较复杂,暂不讨论
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIC_AIV_FLAG + FLAG_ID_MAX);
} else {
CrossCoreWaitFlag<SYNC_MODE4, PIPE_MTE3>(SYNC_AIC_AIV_FLAG);
int64_t rowFactor = mL1RealSize / 2;
int64_t tailRowFactor = mL1RealSize - rowFactor;
curRowFactor = rowFactor;
int64_t mL1SizeAlign = CeilAlign(mL1RealSize, AscendC::BLOCK_CUBE);
if (curBlockIdx % 2 == 1) {
curRowFactor = tailRowFactor;
}
uint64_t cvLoopKSize = CeilDiv(kL1RealSize, tilingData->kUbSize);
uint64_t kReminderSize = kL1RealSize - (cvLoopKSize - 1) * tilingData->kUbSize;
float coeff = 1 / static_cast<float>(tilingData->hcMult * tilingData->d);
for (int64_t cvLoopIdx = 0; cvLoopIdx < cvLoopKSize; cvLoopIdx++) {
uint64_t kRealSize = (cvLoopIdx == cvLoopKSize - 1) ? kReminderSize : tilingData->kUbSize;
xLocal = xQue.template AllocTensor<T>();
CopyIn(xGm[xGmBaseOffset + xSplitOffset + roundIdx * tilingData->mL1Size * tilingData->hcMult * tilingData->d + kGmOffset + cvLoopIdx * tilingData->kUbSize],
xLocal, curRowFactor, kRealSize, tilingData->hcMult * tilingData->d - kRealSize);
xQue.template EnQue(xLocal);
xLocal = xQue.template DeQue<T>();
if (kGmOffset == kGmStartOffset && cvLoopIdx == 0) {
VFProcessCastAndInvRmsPart1<T, false>(rmsNormLocal, xCastLocal, xLocal, coeff, curRowFactor, kRealSize);
} else {
VFProcessCastAndInvRmsPart1<T, true>(rmsNormLocal, xCastLocal, xLocal, coeff, curRowFactor, kRealSize);
}
xQue.template FreeTensor(xLocal);
WaitFlag<HardEvent::MTE3_V>(static_cast<event_t>(bufferIdx & 1));
VFTransND2NZ(xNd2NzLocal[nd2NzBufSize * (bufferIdx & 1)], xCastLocal, curRowFactor, kRealSize);
SetFlag<HardEvent::V_MTE3>(static_cast<event_t>(bufferIdx & 1));
WaitFlag<HardEvent::V_MTE3>(static_cast<event_t>(bufferIdx & 1));
if (curBlockIdx % 2 == 0) {
DataCopyParams dataCopyXParams;
dataCopyXParams.blockCount = CeilDiv(kRealSize, C0_SIZE);
dataCopyXParams.blockLen = curRowFactor * C0_SIZE * sizeof(float) / BLOCK_SIZE;
dataCopyXParams.srcStride = CeilAlign(curRowFactor, C0_SIZE) - curRowFactor;
dataCopyXParams.dstStride = CeilAlign(mL1RealSize, 16) - curRowFactor;
CopyToL1(xNd2NzLocal[nd2NzBufSize * (bufferIdx & 1)], xL1_[(aL1BufferID_ * L1_BUF_OFFSET) + cvLoopIdx * tilingData->kUbSize * mL1SizeAlign], dataCopyXParams);
} else {
DataCopyParams dataCopyXParams;
dataCopyXParams.blockCount = CeilDiv(kRealSize, C0_SIZE);
dataCopyXParams.blockLen = curRowFactor * C0_SIZE * sizeof(float) / BLOCK_SIZE;
dataCopyXParams.srcStride = CeilAlign(curRowFactor, C0_SIZE) - curRowFactor;
dataCopyXParams.dstStride = CeilAlign(mL1RealSize, 16) - curRowFactor;
CopyToL1(xNd2NzLocal[nd2NzBufSize * (bufferIdx & 1)], xL1_[(aL1BufferID_ * L1_BUF_OFFSET) + rowFactor * (BLOCK_SIZE / sizeof(float)) + cvLoopIdx * tilingData->kUbSize * mL1SizeAlign], dataCopyXParams);
}
SetFlag<HardEvent::MTE3_V>(static_cast<event_t>(bufferIdx & 1));
bufferIdx++;
}
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE3>(SYNC_AIV_AIC_FLAG);
}
aL1BufferID_ ^= 1;
}
if ASCEND_IS_AIV {
int64_t kBaseOffset = (GetBlockIdx() / 2) % tilingData->kBlockFactor * tilingData->bs;
rmsQue.template EnQue(rmsNormLocal);
rmsNormLocal = rmsQue.template DeQue<float>();
CopyOut(rmsNormLocal, rmsGm[kBaseOffset + rmsGmBaseOffset + rmsSplitOffset + roundIdx * tilingData->mL1Size], 1, curRowFactor);
rmsQue.template FreeTensor(rmsNormLocal);
}
}
if ASCEND_IS_AIV {
WaitFlag<HardEvent::MTE3_V>(static_cast<event_t>(0));
WaitFlag<HardEvent::MTE3_V>(static_cast<event_t>(1));
}
}
SyncAll<false>();
}
private:
TPipe* pipe;
const HcPreTilingData* tilingData;
// (M, K) * (N, K)
GlobalTensor<T> xGm;
GlobalTensor<float> hcFnGm;
GlobalTensor<float> mmGm;
GlobalTensor<float> rmsGm;
TQue<QuePosition::VECIN, 1> xQue;
TQue<QuePosition::VECOUT, 1> rmsQue;
TBuf<QuePosition::VECCALC> castBuf;
TBuf<QuePosition::VECCALC> nd2NzBuf;
LocalTensor<T> xLocal;
LocalTensor<float> mmXLocal;
LocalTensor<float> rmsNormLocal;
LocalTensor<float> xCastLocal;
LocalTensor<float> xNd2NzLocal;
HcPreCubeCompute mmService_;
LocalTensor<float> xL1_;
LocalTensor<float> wL1_;
static constexpr uint64_t SYNC_AIV_AIC_FLAG = 8;
static constexpr uint64_t SYNC_AIC_AIV_FLAG = 9;
static constexpr uint64_t SYNC_AIV_AIC_PRE_POST_FLAG = 10;
static constexpr uint64_t SYNC_AIC_AIV_PRE_POST_FLAG = 11;
static constexpr uint64_t FLAG_ID_MAX = 16;
uint64_t cvLoopIdx_ = 0;
uint8_t aL1BufferID_ = 0;
};
template <typename T>
class HcPreMKSplitCorePart2 {
public:
__aicore__ inline HcPreMKSplitCorePart2()
{}
__aicore__ inline void Init(
GM_ADDR x, GM_ADDR hcScale, GM_ADDR hcBase, GM_ADDR y, GM_ADDR post,
GM_ADDR combFrag, GM_ADDR workspace, const HcPreTilingData* tilingDataPtr, TPipe* pipePtr)
{
pipe = pipePtr;
tilingData = tilingDataPtr;
xGm.SetGlobalBuffer((__gm__ T*)x);
hcScaleGm.SetGlobalBuffer((__gm__ float*)hcScale);
hcBaseGm.SetGlobalBuffer((__gm__ float*)hcBase);
yGm.SetGlobalBuffer((__gm__ T*)y);
postGm.SetGlobalBuffer((__gm__ float*)post);
combFragGm.SetGlobalBuffer((__gm__ float*)combFrag);
mmGm.SetGlobalBuffer((__gm__ float*)workspace);
rmsGm.SetGlobalBuffer((__gm__ float*)workspace + tilingData->kBlockFactor * tilingData->bs * tilingData->hcMix);
// InQue
pipe->InitBuffer(
xQue, 2, tilingData->stage2RowFactor * tilingData->hcMult * RoundUp<T>(tilingData->dFactor) * sizeof(T));
int64_t rmsAndmmQueSize = tilingData->kBlockFactor * RoundUp<float>(tilingData->stage2RowFactor) * sizeof(float) +
tilingData->kBlockFactor * tilingData->stage2RowFactor * RoundUp<float>(tilingData->hcMix) * sizeof(float);
pipe->InitBuffer(rmsAndmmQue, 2, rmsAndmmQueSize);
// OutQue
pipe->InitBuffer(
yQue, 2, tilingData->stage2RowFactor * RoundUp<T>(tilingData->dFactor) * sizeof(T));
pipe->InitBuffer(postQue, 2, tilingData->stage2RowFactor * tilingData->hcMultAlign * sizeof(float));
pipe->InitBuffer(combFragQue, DOUBLE_BUFFER,
tilingData->stage2RowFactor * tilingData->hcMult * tilingData->hcMult * sizeof(float));
// TBuf
pipe->InitBuffer(hcBaseBuf0, tilingData->hcMultAlign * sizeof(float));
pipe->InitBuffer(hcBaseBuf1, tilingData->hcMultAlign * sizeof(float));
pipe->InitBuffer(hcBaseBuf2, tilingData->hcMult * tilingData->hcMultAlign * sizeof(float));
pipe->InitBuffer(mixesBuf, tilingData->stage2RowFactor * RoundUp<float>(tilingData->hcMix) * sizeof(float));
hcBase0Local = hcBaseBuf0.Get<float>();
hcBase1Local = hcBaseBuf1.Get<float>();
hcBase2Local = hcBaseBuf2.Get<float>();
mixesLocal = mixesBuf.Get<float>();
}
__aicore__ inline void Process()
{
if ASCEND_IS_AIV {
int64_t stage1UsedCoreNum = tilingData->cubeBlockDimK;
int64_t curBlockIdx = GetBlockIdx();
int64_t stage2UsedCoreNum = tilingData->secondUsedCoreNum;
if (curBlockIdx >= stage2UsedCoreNum) {
return;
}
int64_t rowOuterLoop =
(curBlockIdx == stage2UsedCoreNum - 1) ? tilingData->rowLoopOfTailBlock : tilingData->rowLoopOfFormerBlock;
int64_t tailRowFactor = (curBlockIdx == stage2UsedCoreNum - 1) ? tilingData->tailRowFactorOfTailBlock :
tilingData->tailRowFactorOfFormerBlock;
CopyIn(hcBaseGm, hcBase0Local, 1, tilingData->hcMult);
CopyIn(hcBaseGm[tilingData->hcMult], hcBase1Local, 1, tilingData->hcMult);
CopyIn(hcBaseGm[tilingData->hcMult * 2], hcBase2Local, 1, tilingData->hcMult * tilingData->hcMult);
event_t eventId = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V));
SetFlag<HardEvent::MTE2_V>(eventId);
WaitFlag<HardEvent::MTE2_V>(eventId);
int64_t mmGmBaseOffset = curBlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMix;
int64_t rmsGmBaseOffset = curBlockIdx * tilingData->rowOfFormerBlock;
int64_t xGmBaseOffset = curBlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMult * tilingData->d;
int64_t mmLocalSize = stage1UsedCoreNum * tilingData->stage2RowFactor * RoundUp<float>(tilingData->hcMix);
for (int64_t rowOuterIdx = 0; rowOuterIdx < rowOuterLoop; rowOuterIdx++) {
int64_t curRowFactor = (rowOuterIdx == rowOuterLoop - 1) ? tailRowFactor : tilingData->stage2RowFactor;
rmsAndmmLocal = rmsAndmmQue.AllocTensor<float>();
CopyInWithLoopMode(
mmGm[mmGmBaseOffset + rowOuterIdx * tilingData->stage2RowFactor * tilingData->hcMix], rmsAndmmLocal, tilingData->kBlockFactor, curRowFactor, tilingData->hcMix, tilingData->bs * tilingData->hcMix);
CopyIn(
rmsGm[rmsGmBaseOffset + rowOuterIdx * tilingData->stage2RowFactor],
rmsAndmmLocal[mmLocalSize], tilingData->kBlockFactor, curRowFactor, tilingData->bs - curRowFactor);
rmsAndmmQue.EnQue(rmsAndmmLocal);
rmsAndmmLocal = rmsAndmmQue.DeQue<float>();
VFProcessInvRmsPart3WithGroupReduce(mixesLocal, rmsAndmmLocal, rmsAndmmLocal[mmLocalSize], tilingData->normEps, tilingData->kBlockFactor, curRowFactor, tilingData->hcMix);
VFProcessPre(
mixesLocal, mixesLocal, hcBase0Local, hcScaleGm.GetValue(0), tilingData->hcEps,
curRowFactor, tilingData->hcMult, tilingData->hcMix);
for (int64_t dLoopIdx = 0; dLoopIdx < tilingData->dLoop; dLoopIdx++) {
int64_t curDFactor =
(dLoopIdx == tilingData->dLoop - 1) ? tilingData->tailDFactor : tilingData->dFactor;
xLocal = xQue.template AllocTensor<T>();
CopyIn(
xGm[xGmBaseOffset + rowOuterIdx * tilingData->stage2RowFactor * tilingData->hcMult * tilingData->d +
dLoopIdx * tilingData->dFactor],
xLocal, curRowFactor * tilingData->hcMult, curDFactor, tilingData->d - curDFactor);
xQue.template EnQue(xLocal);
xLocal = xQue.template DeQue<T>();
yLocal = yQue.template AllocTensor<T>();
VFProcessY(yLocal, mixesLocal, xLocal, curRowFactor, tilingData->hcMult, curDFactor, tilingData->hcMix);
xQue.template FreeTensor(xLocal);
yQue.template EnQue(yLocal);
yLocal = yQue.template DeQue<T>();
CopyOut(yLocal, yGm[curBlockIdx * tilingData->rowOfFormerBlock * tilingData->d + rowOuterIdx * tilingData->stage2RowFactor * tilingData->d + dLoopIdx * tilingData->dFactor], curRowFactor, curDFactor, tilingData->d - curDFactor);
yQue.template FreeTensor(yLocal);
}
// post
postLocal = postQue.AllocTensor<float>();
VFProcessPost(
postLocal, mixesLocal[tilingData->hcMult], hcBase1Local,
hcScaleGm.GetValue(1), tilingData->hcEps, curRowFactor, tilingData->hcMult, tilingData->hcMix);
postQue.EnQue(postLocal);
postLocal = postQue.DeQue<float>();
CopyOut(postLocal, postGm[curBlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMult + rowOuterIdx * tilingData->stage2RowFactor * tilingData->hcMult], curRowFactor, tilingData->hcMult);
postQue.FreeTensor(postLocal);
// combFrag
combFragLocal = combFragQue.AllocTensor<float>();
VFProcessCombFragPacked(
combFragLocal, mixesLocal[tilingData->hcMult * 2], hcBase2Local, hcScaleGm.GetValue(2), tilingData->hcEps,
tilingData->iterTimes - 1, curRowFactor, tilingData->hcMult, tilingData->hcMix);
rmsAndmmQue.FreeTensor(rmsAndmmLocal);
combFragQue.EnQue(combFragLocal);
combFragLocal = combFragQue.DeQue<float>();
int64_t combLen = tilingData->hcMult * tilingData->hcMult;
int64_t combOutOffset = (curBlockIdx * tilingData->rowOfFormerBlock +
rowOuterIdx * tilingData->stage2RowFactor) * combLen;
CopyOut(combFragLocal, combFragGm[combOutOffset], curRowFactor, combLen);
combFragQue.FreeTensor(combFragLocal);
}
}
}
private:
TPipe* pipe;
const HcPreTilingData* tilingData;
GlobalTensor<float> hcScaleGm;
GlobalTensor<float> hcBaseGm;
GlobalTensor<T> xGm;
GlobalTensor<T> yGm;
GlobalTensor<float> postGm;
GlobalTensor<float> combFragGm;
GlobalTensor<float> mmGm;
GlobalTensor<float> rmsGm;
TQue<QuePosition::VECIN, 1> rmsAndmmQue;
TQue<QuePosition::VECIN, 1> xQue;
TQue<QuePosition::VECOUT, 1> yQue;
TQue<QuePosition::VECOUT, 1> postQue;
TQue<QuePosition::VECOUT, 1> combFragQue;
TBuf<QuePosition::VECCALC> mixesBuf;
TBuf<QuePosition::VECCALC> hcBaseBuf0;
TBuf<QuePosition::VECCALC> hcBaseBuf1;
TBuf<QuePosition::VECCALC> hcBaseBuf2;
LocalTensor<float> mixesLocal;
LocalTensor<float> rmsAndmmLocal;
LocalTensor<T> xLocal;
LocalTensor<T> yLocal;
LocalTensor<float> postLocal;
LocalTensor<float> combFragLocal;
LocalTensor<float> hcBase0Local;
LocalTensor<float> hcBase1Local;
LocalTensor<float> hcBase2Local;
};
} // namespace HCPreSinkhorn
#endif

View File

@@ -0,0 +1,411 @@
/**
* 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 hc_pre_m_k_split_core.h
* \brief
*/
#ifndef HC_PRE_M_SPLIT_CORE_H
#define HC_PRE_M_SPLIT_CORE_H
#include "kernel_operator.h"
#include "hc_pre_base_arch35.h"
#include "hc_pre_cube_compute_arch35.h"
namespace HcPreNs {
using namespace AscendC;
template <typename T>
class HcPreMSplitCorePart1 {
public:
__aicore__ inline HcPreMSplitCorePart1()
{}
__aicore__ inline void Init(
GM_ADDR x, GM_ADDR hcFn, GM_ADDR hcScale, GM_ADDR hcBase,
GM_ADDR y, GM_ADDR post, GM_ADDR combFrag, const HcPreTilingData* tilingDataPtr, TPipe* pipePtr)
{
pipe = pipePtr;
tilingData = tilingDataPtr;
xGm.SetGlobalBuffer((__gm__ T*)x);
hcFnGm.SetGlobalBuffer((__gm__ float*)hcFn);
yGm.SetGlobalBuffer((__gm__ T*)y);
hcScaleGm.SetGlobalBuffer((__gm__ float*)hcScale);
hcBaseGm.SetGlobalBuffer((__gm__ float*)hcBase);
postGm.SetGlobalBuffer((__gm__ float*)post);
combFragGm.SetGlobalBuffer((__gm__ float*)combFrag);
TBuf<TPosition::A1> l1Buffer;
pipe->InitBuffer(l1Buffer, L1_ALLOC_SIZE);
xL1_ = l1Buffer.Get<float>();
wL1_ = l1Buffer.Get<float>()[L1_BUF_NUM * L1_BUF_OFFSET];
pipe->InitBufPool(tbufPool0, tilingData->bufferPool0Size);
tbufPool0.InitBuffer(mmXBuf, CeilDiv(tilingData->mL1Size, 2) * RoundUp<float>(tilingData->hcMix) * sizeof(float));
mmXLocal = mmXBuf.Get<float>();
if ASCEND_IS_AIC {
mmService_.Init();
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIC_AIV_FLAG);
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIC_AIV_FLAG + FLAG_ID_MAX);
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIC_AIV_FLAG);
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIC_AIV_FLAG + FLAG_ID_MAX);
} else {
tbufPool0.InitBuffer(rmsNormBuf, RoundUp<float>(CeilDiv(tilingData->mL1Size, 2)) * sizeof(float));
tbufPool0.InitBufPool(tbufPool1, tilingData->bufferPool1Size);
tbufPool0.InitBuffer(hcBaseBuf0, tilingData->hcMultAlign * sizeof(float));
tbufPool0.InitBuffer(hcBaseBuf1, tilingData->hcMultAlign * sizeof(float));
tbufPool0.InitBuffer(hcBaseBuf2, tilingData->hcMult * tilingData->hcMultAlign * sizeof(float));
hcBase0Local = hcBaseBuf0.Get<float>();
hcBase1Local = hcBaseBuf1.Get<float>();
hcBase2Local = hcBaseBuf2.Get<float>();
}
}
__aicore__ inline void Process()
{
int64_t curBlockIdx = GetBlockIdx();
int64_t logicalBlockIdx = curBlockIdx;
if ASCEND_IS_AIV {
logicalBlockIdx = curBlockIdx / 2;
}
if (logicalBlockIdx >= tilingData->cubeBlockDimM) {
return;
}
if ASCEND_IS_AIV {
CopyIn(hcBaseGm, hcBase0Local, 1, tilingData->hcMult);
CopyIn(hcBaseGm[tilingData->hcMult], hcBase1Local, 1, tilingData->hcMult);
CopyIn(hcBaseGm[tilingData->hcMult * 2], hcBase2Local, 1, tilingData->hcMult * tilingData->hcMult);
event_t eventId = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V));
SetFlag<HardEvent::MTE2_V>(eventId);
WaitFlag<HardEvent::MTE2_V>(eventId);
}
int64_t totalBlockNum = GetBlockNum();
uint64_t mBlkDimIdx = curBlockIdx % tilingData->cubeBlockDimM;
uint64_t kBlkDimIdx = curBlockIdx % tilingData->cubeBlockDimK;
// todo 移到tiling计算
uint64_t mCnt = CeilDiv(tilingData->bs, tilingData->mL1Size);
uint64_t singleCoreMaxRound = CeilDiv(mCnt, tilingData->cubeBlockDimM);
uint64_t mainCoreCount = mCnt % tilingData->cubeBlockDimM;
uint64_t singleCoreRound = (mainCoreCount == 0 || logicalBlockIdx < mainCoreCount) ? singleCoreMaxRound : singleCoreMaxRound - 1;
uint64_t mGmOffset = 0;
if ASCEND_IS_AIC {
if (mainCoreCount == 0 || curBlockIdx <= mainCoreCount) {
mGmOffset = curBlockIdx * singleCoreMaxRound * tilingData->mL1Size;
} else {
mGmOffset = (mainCoreCount * singleCoreMaxRound + (curBlockIdx - mainCoreCount) * (singleCoreMaxRound - 1)) * tilingData->mL1Size;
}
} else {
if (mainCoreCount == 0 || (curBlockIdx / 2) <= mainCoreCount) {
mGmOffset = curBlockIdx / 2 * singleCoreMaxRound * tilingData->mL1Size;
} else {
mGmOffset = (mainCoreCount * singleCoreMaxRound + (curBlockIdx / 2 - mainCoreCount) * (singleCoreMaxRound - 1)) * tilingData->mL1Size;
}
}
int64_t xGmBaseOffset = 0;
int64_t yGmBaseOffset = 0;
int64_t postGmBaseOffset = 0;
int64_t combFragGmBaseOffset = 0;
if ASCEND_IS_AIV {
xGmBaseOffset = mGmOffset * tilingData->hcMult * tilingData->d;
yGmBaseOffset = mGmOffset * tilingData->d;
postGmBaseOffset = mGmOffset * tilingData->hcMult;
combFragGmBaseOffset = mGmOffset * tilingData->hcMult * tilingData->hcMult;
SetFlag<HardEvent::MTE3_MTE2>(static_cast<event_t>(0));
}
uint64_t cvLoopKSize = tilingData->kL1Size / tilingData->kUbSize;
int64_t xSplitOffset = 0;
int64_t ySplitOffset = 0;
int64_t postSplitOffset = 0;
int64_t combFragSplitOffset = 0;
int64_t xOutSplitOffset = 0;
// m轴切分 按照0 0 1 1..分核
for (uint64_t roundIdx = 0; roundIdx < singleCoreRound; mGmOffset += tilingData->mL1Size, ++roundIdx)
{
uint64_t mL1RealSize = AscendC::Std::min(tilingData->bs - mGmOffset, (uint64_t)tilingData->mL1Size);
uint64_t kGmStartOffset = 0;
uint64_t kGmEndOffset = tilingData->multCoreSplitKSize;
uint64_t nd2NzBufSize = CeilAlign(tilingData->mUbSize, C0_SIZE) * RoundUp<float>(tilingData->kUbSize);
if ASCEND_IS_AIV {
tbufPool1.Reset();
tbufPool1.InitBuffer(xQue, 2, tilingData->mUbSize * RoundUp<T>(tilingData->kUbSize) * sizeof(T));
tbufPool1.InitBuffer(castBuf, tilingData->mUbSize * (RoundUp<float>(tilingData->kUbSize) * sizeof(float) + BLOCK_SIZE));
tbufPool1.InitBuffer(nd2NzBuf, nd2NzBufSize * sizeof(float) * DOUBLE_BUFFER);
xCastLocal = castBuf.Get<float>();
xNd2NzLocal = nd2NzBuf.Get<float>();
rmsNormLocal = rmsNormBuf.Get<float>();
WaitFlag<HardEvent::MTE3_MTE2>(static_cast<event_t>(0));
if (GetBlockIdx() % 2 != 0) {
xSplitOffset = CeilDiv(mL1RealSize, 2) * tilingData->hcMult * tilingData->d;
xOutSplitOffset = CeilDiv(mL1RealSize, 2) * tilingData->hcMult * tilingData->d;
ySplitOffset = CeilDiv(mL1RealSize, 2) * tilingData->d;
postSplitOffset = CeilDiv(mL1RealSize, 2) * tilingData->hcMult;
combFragSplitOffset = CeilDiv(mL1RealSize, 2) * tilingData->hcMult * tilingData->hcMult;
}
}
// k轴切分(kCoreDim=1)
int64_t bufferIdx = 0;
if ASCEND_IS_AIV {
SetFlag<HardEvent::MTE3_V>(static_cast<event_t>(0));
SetFlag<HardEvent::MTE3_V>(static_cast<event_t>(1));
}
for (int64_t kGmOffset = kGmStartOffset; kGmOffset < kGmEndOffset; kGmOffset += tilingData->kL1Size) {
if ASCEND_IS_AIC {
bool isFirstKL1 = kGmOffset == kGmStartOffset;
bool isLastKL1 = (kGmOffset + tilingData->kL1Size) >= kGmEndOffset;
uint64_t kL1RealSize = AscendC::Std::min(kGmEndOffset - kGmOffset, (uint64_t)tilingData->kL1Size);
mmService_.CopyInB1Nd2Nz(tilingData->multCoreSplitKSize, kL1RealSize,
tilingData->hcMix, hcFnGm[kGmOffset],
wL1_[mmService_.GetBL1BufferId() * L1_BUF_OFFSET]);
CrossCoreWaitFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIV_AIC_FLAG + FLAG_ID_MAX);
CrossCoreWaitFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIV_AIC_FLAG);
uint64_t mL1AlignSize = Align(mL1RealSize, AscendC::BLOCK_CUBE);
uint64_t nL1AlignSize = Align((uint64_t)tilingData->hcMix, AscendC::BLOCK_CUBE);
mmService_.Process(tilingData->bs, tilingData->hcMix, mL1RealSize, (256 / AscendC::Std::max(mL1AlignSize, nL1AlignSize)) * 32,
isFirstKL1, isLastKL1, xL1_[aL1BufferID_ * L1_BUF_OFFSET], wL1_[mmService_.GetBL1BufferId() * L1_BUF_OFFSET]);
if (isLastKL1) {
mmService_.CopyOut(mmXLocal);
CrossCoreSetFlag<SYNC_MODE4, PIPE_FIX>(SYNC_AIC_AIV_PRE_POST_FLAG);
CrossCoreSetFlag<SYNC_MODE4, PIPE_FIX>(SYNC_AIC_AIV_PRE_POST_FLAG + FLAG_ID_MAX);
}
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIC_AIV_FLAG); // 写出ub搬出,cv流水同步比较复杂,暂不讨论
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE1>(SYNC_AIC_AIV_FLAG + FLAG_ID_MAX);
} else {
CrossCoreWaitFlag<SYNC_MODE4, PIPE_MTE3>(SYNC_AIC_AIV_FLAG);
// Even cores take the first half with CeilDiv, and odd cores take the second half.
// This must match the Sinkhorn/output row split; otherwise odd mL1RealSize shifts the
// second half by one row.
int64_t rowFactor = CeilDiv(mL1RealSize, 2);
int64_t tailRowFactor = mL1RealSize - rowFactor;
int64_t curRowFactor = rowFactor;
int64_t mL1SizeAlign = CeilAlign(mL1RealSize, AscendC::BLOCK_CUBE);
if (curBlockIdx % 2 == 1) {
curRowFactor = tailRowFactor;
}
float coeff = 1 / static_cast<float>(tilingData->hcMult * tilingData->d);
for (int64_t cvLoopIdx = 0; cvLoopIdx < cvLoopKSize; cvLoopIdx++) {
uint64_t kRealSize = kGmOffset + tilingData->kUbSize >= kGmEndOffset ? kGmEndOffset - kGmOffset : tilingData->kUbSize;
xLocal = xQue.template AllocTensor<T>();
CopyIn(xGm[xGmBaseOffset + xSplitOffset + roundIdx * tilingData->mL1Size * tilingData->hcMult * tilingData->d + kGmOffset + cvLoopIdx * tilingData->kUbSize],
xLocal, curRowFactor, tilingData->kUbSize, tilingData->hcMult * tilingData->d - tilingData->kUbSize);
xQue.template EnQue(xLocal);
xLocal = xQue.template DeQue<T>();
if (kGmOffset == kGmStartOffset && cvLoopIdx == 0) {
VFProcessCastAndInvRmsPart1<T, false>(rmsNormLocal, xCastLocal, xLocal, coeff, curRowFactor, tilingData->kUbSize);
} else {
VFProcessCastAndInvRmsPart1<T, true>(rmsNormLocal, xCastLocal, xLocal, coeff, curRowFactor, tilingData->kUbSize);
}
xQue.template FreeTensor(xLocal);
WaitFlag<HardEvent::MTE3_V>(static_cast<event_t>(bufferIdx & 1));
VFTransND2NZ(xNd2NzLocal[nd2NzBufSize * (bufferIdx & 1)], xCastLocal, curRowFactor, tilingData->kUbSize);
SetFlag<HardEvent::V_MTE3>(static_cast<event_t>(bufferIdx & 1));
WaitFlag<HardEvent::V_MTE3>(static_cast<event_t>(bufferIdx & 1));
if (curBlockIdx % 2 == 0) {
DataCopyParams dataCopyXParams;
dataCopyXParams.blockCount = CeilDiv(tilingData->kUbSize, C0_SIZE);
dataCopyXParams.blockLen = curRowFactor * C0_SIZE * sizeof(float) / BLOCK_SIZE;
dataCopyXParams.srcStride = CeilAlign(curRowFactor, C0_SIZE) - curRowFactor;
dataCopyXParams.dstStride = CeilAlign(mL1RealSize, 16) - curRowFactor;
CopyToL1(xNd2NzLocal[nd2NzBufSize * (bufferIdx & 1)], xL1_[(aL1BufferID_ * L1_BUF_OFFSET) + cvLoopIdx * tilingData->kUbSize * mL1SizeAlign], dataCopyXParams);
} else {
DataCopyParams dataCopyXParams;
dataCopyXParams.blockCount = CeilDiv(tilingData->kUbSize, C0_SIZE);
dataCopyXParams.blockLen = curRowFactor * C0_SIZE * sizeof(float) / BLOCK_SIZE;
dataCopyXParams.srcStride = CeilAlign(curRowFactor, C0_SIZE) - curRowFactor;
dataCopyXParams.dstStride = CeilAlign(mL1RealSize, 16) - curRowFactor;
CopyToL1(xNd2NzLocal[nd2NzBufSize * (bufferIdx & 1)], xL1_[(aL1BufferID_ * L1_BUF_OFFSET) + rowFactor * (BLOCK_SIZE / sizeof(float)) + cvLoopIdx * tilingData->kUbSize * mL1SizeAlign], dataCopyXParams);
}
SetFlag<HardEvent::MTE3_V>(static_cast<event_t>(bufferIdx & 1));
bufferIdx++;
}
CrossCoreSetFlag<SYNC_MODE4, PIPE_MTE3>(SYNC_AIV_AIC_FLAG);
}
aL1BufferID_ ^= 1;
}
if ASCEND_IS_AIV {
WaitFlag<HardEvent::MTE3_V>(static_cast<event_t>(0));
WaitFlag<HardEvent::MTE3_V>(static_cast<event_t>(1));
CrossCoreWaitFlag<SYNC_MODE4, PIPE_V>(SYNC_AIC_AIV_PRE_POST_FLAG);
// mm计算结果存入mmXLocal,mmXLocal每轮循环需要累加;
tbufPool1.Reset();
tbufPool1.InitBuffer(xQue, 2, tilingData->rowInnerFactor * tilingData->hcMult * RoundUp<T>(tilingData->dFactor) * sizeof(T));
tbufPool1.InitBuffer(
yQue, 2, tilingData->rowInnerFactor * RoundUp<T>(tilingData->dFactor) * sizeof(T));
tbufPool1.InitBuffer(postQue, 2, tilingData->rowInnerFactor * tilingData->hcMultAlign * sizeof(float));
tbufPool1.InitBuffer(combFragQue, DOUBLE_BUFFER,
tilingData->rowInnerFactor * tilingData->hcMult * tilingData->hcMult * sizeof(float));
// TBuf
tbufPool1.InitBuffer(mixesBuf, tilingData->rowInnerFactor * RoundUp<float>(tilingData->hcMix) * sizeof(float));
mixesLocal = mixesBuf.Get<float>();
SetWaitFlag<HardEvent::V_MTE2>(HardEvent::V_MTE2);
// m内层循环
int64_t currentRow = mL1RealSize / 2;
if (mL1RealSize % 2 == 1 && curBlockIdx % 2 == 0) {
// m不整除时偶数核多处理一行
currentRow += 1;
}
for (int64_t innerRowIdx = 0; innerRowIdx < currentRow; innerRowIdx += tilingData->rowInnerFactor) {
int64_t currentInnerRowFactor = innerRowIdx + tilingData->rowInnerFactor >= currentRow ? currentRow - innerRowIdx :
tilingData->rowInnerFactor;
VFProcessInvRmsPart3(mixesLocal, mmXLocal[innerRowIdx * tilingData->hcMix], rmsNormLocal[innerRowIdx],
tilingData->normEps, currentInnerRowFactor, tilingData->hcMix);
VFProcessPre(
mixesLocal, mixesLocal, hcBase0Local, hcScaleGm.GetValue(0), tilingData->hcEps,
currentInnerRowFactor, tilingData->hcMult, tilingData->hcMix);
for (int64_t dLoopIdx = 0; dLoopIdx < tilingData->dLoop; dLoopIdx++)
{
int64_t curDFactor =
(dLoopIdx == tilingData->dLoop - 1) ? tilingData->tailDFactor : tilingData->dFactor;
xLocal = xQue.template AllocTensor<T>();
CopyIn(
xGm[xGmBaseOffset + xOutSplitOffset + roundIdx * tilingData->mL1Size * tilingData->hcMult * tilingData->d +
innerRowIdx * tilingData->hcMult * tilingData->d + dLoopIdx * tilingData->dFactor],
xLocal, currentInnerRowFactor * tilingData->hcMult, curDFactor, tilingData->d - curDFactor);
xQue.template EnQue(xLocal);
xLocal = xQue.template DeQue<T>();
yLocal = yQue.template AllocTensor<T>();
VFProcessY(yLocal, mixesLocal, xLocal, currentInnerRowFactor, tilingData->hcMult, curDFactor, tilingData->hcMix);
xQue.template FreeTensor(xLocal);
yQue.template EnQue(yLocal);
yLocal = yQue.template DeQue<T>();
CopyOut(yLocal, yGm[yGmBaseOffset + ySplitOffset + roundIdx * tilingData->mL1Size * tilingData->d + innerRowIdx * tilingData->d + dLoopIdx * tilingData->dFactor],
currentInnerRowFactor, curDFactor, tilingData->d - curDFactor);
yQue.template FreeTensor(yLocal);
}
// post
postLocal = postQue.AllocTensor<float>();
VFProcessPost(
postLocal, mixesLocal[tilingData->hcMult], hcBase1Local,
hcScaleGm.GetValue(1), tilingData->hcEps, currentInnerRowFactor, tilingData->hcMult, tilingData->hcMix);
postQue.EnQue(postLocal);
postLocal = postQue.DeQue<float>();
CopyOut(postLocal, postGm[postGmBaseOffset + postSplitOffset + roundIdx * tilingData->mL1Size * tilingData->hcMult + innerRowIdx * tilingData->hcMult], currentInnerRowFactor, tilingData->hcMult);
postQue.FreeTensor(postLocal);
// combFrag
combFragLocal = combFragQue.AllocTensor<float>();
VFProcessCombFragPacked(
combFragLocal, mixesLocal[tilingData->hcMult * 2], hcBase2Local, hcScaleGm.GetValue(2), tilingData->hcEps,
tilingData->iterTimes - 1, currentInnerRowFactor, tilingData->hcMult, tilingData->hcMix);
combFragQue.EnQue(combFragLocal);
combFragLocal = combFragQue.DeQue<float>();
CopyOut(combFragLocal, combFragGm[combFragGmBaseOffset + combFragSplitOffset + roundIdx * tilingData->mL1Size * tilingData->hcMult * tilingData->hcMult + innerRowIdx * tilingData->hcMult * tilingData->hcMult],
currentInnerRowFactor, tilingData->hcMult * tilingData->hcMult);
combFragQue.FreeTensor(combFragLocal);
}
SetFlag<HardEvent::MTE3_MTE2>(static_cast<event_t>(0));
}
}
if ASCEND_IS_AIV {
WaitFlag<HardEvent::MTE3_MTE2>(static_cast<event_t>(0));
} else {
mmService_.End();
}
}
private:
TPipe *pipe;
const HcPreTilingData *tilingData;
// (M, K) * (N, K)
GlobalTensor<T> xGm;
GlobalTensor<float> hcFnGm;
GlobalTensor<float> workspaceGm;
GlobalTensor<T> yGm;
GlobalTensor<float> invRmsGm;
GlobalTensor<float> hcScaleGm;
GlobalTensor<float> hcBaseGm;
GlobalTensor<float> postGm;
GlobalTensor<float> combFragGm;
TQue<QuePosition::VECIN, 1> xQue;
TQue<QuePosition::VECOUT, 1> yQue;
TQue<QuePosition::VECOUT, 1> postQue;
TQue<QuePosition::VECOUT, 1> combFragQue;
TBuf<QuePosition::VECCALC> castBuf;
TBuf<QuePosition::VECCALC> nd2NzBuf;
TQue<QuePosition::VECIN, 1> squareSumQue;
TBuf<QuePosition::VECCALC> hcBaseBuf0;
TBuf<QuePosition::VECCALC> hcBaseBuf1;
TBuf<QuePosition::VECCALC> hcBaseBuf2;
TBuf<QuePosition::VECCALC> rowBrcbBuf0;
TBuf<QuePosition::VECCALC> hcBrcbBuf1;
TBuf<QuePosition::VECCALC> reduceBuf;
TBuf<QuePosition::VECCALC> rsqrtBuf;
TBuf<QuePosition::VECCALC> squareReduceBuf;
TBuf<QuePosition::VECCALC> mixes01ReduceBuf;
TBuf<QuePosition::VECCALC> xCastBuf;
TBuf<QuePosition::VECCALC> yCastBuf;
TBuf<QuePosition::VECCALC> mixesBuf;
TBuf<QuePosition::VECCALC> rmsNormBuf;
TBuf<QuePosition::VECCALC> mmXBuf;
LocalTensor<T> xLocal;
LocalTensor<T> yLocal;
LocalTensor<float> mmXLocal;
LocalTensor<float> rmsNormLocal;
LocalTensor<float> xCastLocal;
LocalTensor<float> xNd2NzLocal;
LocalTensor<float> mixesLocal;
LocalTensor<float> rmsAndmmLocal;
LocalTensor<float> postLocal;
LocalTensor<float> combFragLocal;
LocalTensor<float> hcBase0Local;
LocalTensor<float> hcBase1Local;
LocalTensor<float> hcBase2Local;
HcPreCubeCompute mmService_;
LocalTensor<float> xL1_;
LocalTensor<float> wL1_;
static constexpr uint64_t SYNC_AIV_AIC_FLAG = 8;
static constexpr uint64_t SYNC_AIC_AIV_FLAG = 9;
static constexpr uint64_t SYNC_AIC_AIV_PRE_POST_FLAG = 10;
static constexpr uint64_t FLAG_ID_MAX = 16;
uint64_t cvLoopIdx_ = 0;
uint8_t aL1BufferID_{0};
TBufPool<QuePosition::VECCALC, 12> tbufPool0;
TBufPool<QuePosition::VECCALC, 12> tbufPool1;
};
} // namespace HCPreSinkhorn
#endif