89
csrc/moe/hc_pre/op_kernel/hc_pre.cpp
Normal file
89
csrc/moe/hc_pre/op_kernel/hc_pre.cpp
Normal 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
|
||||
}
|
||||
688
csrc/moe/hc_pre/op_kernel/hc_pre_base.h
Normal file
688
csrc/moe/hc_pre/op_kernel/hc_pre_base.h
Normal 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
|
||||
1147
csrc/moe/hc_pre/op_kernel/hc_pre_base_arch35.h
Normal file
1147
csrc/moe/hc_pre/op_kernel/hc_pre_base_arch35.h
Normal file
File diff suppressed because it is too large
Load Diff
481
csrc/moe/hc_pre/op_kernel/hc_pre_cube_compute.h
Normal file
481
csrc/moe/hc_pre/op_kernel/hc_pre_cube_compute.h
Normal 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
|
||||
241
csrc/moe/hc_pre/op_kernel/hc_pre_cube_compute_arch35.h
Normal file
241
csrc/moe/hc_pre/op_kernel/hc_pre_cube_compute_arch35.h
Normal 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
|
||||
564
csrc/moe/hc_pre/op_kernel/hc_pre_m_k_split_core.h
Normal file
564
csrc/moe/hc_pre/op_kernel/hc_pre_m_k_split_core.h
Normal 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
|
||||
428
csrc/moe/hc_pre/op_kernel/hc_pre_m_k_split_core_arch35.h
Normal file
428
csrc/moe/hc_pre/op_kernel/hc_pre_m_k_split_core_arch35.h
Normal 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
|
||||
411
csrc/moe/hc_pre/op_kernel/hc_pre_m_split_core_arch35.h
Normal file
411
csrc/moe/hc_pre/op_kernel/hc_pre_m_split_core_arch35.h
Normal 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
|
||||
Reference in New Issue
Block a user