176
csrc/moe/causal_conv1d/op_kernel/arch35/causal_conv1d_regbase.h
Normal file
176
csrc/moe/causal_conv1d/op_kernel/arch35/causal_conv1d_regbase.h
Normal file
@@ -0,0 +1,176 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 causal_conv1d_regbase.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef CAUSAL_CONV1D_REGBASE_H
|
||||
#define CAUSAL_CONV1D_REGBASE_H
|
||||
|
||||
namespace NsCausalConv1d {
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::MicroAPI;
|
||||
|
||||
constexpr uint16_t V_LENGTH = VECTOR_REG_WIDTH / sizeof(float);
|
||||
|
||||
constexpr CastTrait castTraitB16ToB32 = {
|
||||
RegLayout::ZERO, SatMode::UNKNOWN, MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
template <typename T, bool hasActivation>
|
||||
__aicore__ inline void ComputeFnRollingOutputRegbase(LocalTensor<T> ring, LocalTensor<float> currF,
|
||||
LocalTensor<float> state0F, LocalTensor<float> weightF, uint32_t dataCount)
|
||||
{
|
||||
__ubuf__ T* ringAddr = (__ubuf__ T*)ring.GetPhyAddr();
|
||||
__ubuf__ float* currFAddr = (__ubuf__ float*)currF.GetPhyAddr();
|
||||
__ubuf__ float* state0FAddr = (__ubuf__ float*)state0F.GetPhyAddr();
|
||||
__ubuf__ float* weightFAddr = (__ubuf__ float*)weightF.GetPhyAddr();
|
||||
|
||||
uint16_t colLoopTimes = static_cast<uint16_t>(Ceil(dataCount, V_LENGTH));
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<T> ring;
|
||||
RegTensor<float> currF;
|
||||
RegTensor<float> state0F;
|
||||
RegTensor<float> weightF;
|
||||
RegTensor<float> tmp;
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t j = 0; j < colLoopTimes; j++) {
|
||||
pregLoop = UpdateMask<float>(dataCount);
|
||||
DataCopy<T, LoadDist::DIST_UNPACK_B16>(ring, ringAddr + j * V_LENGTH);
|
||||
DataCopy(state0F, state0FAddr + j * V_LENGTH);
|
||||
DataCopy(weightF, weightFAddr + j * V_LENGTH);
|
||||
Cast<float, T, castTraitB16ToB32>(currF, ring, pregLoop);
|
||||
Mul(currF, currF, weightF, pregLoop);
|
||||
Add(state0F, state0F, currF, pregLoop);
|
||||
if constexpr (hasActivation) {
|
||||
Muls(tmp, state0F, -1.0f, pregLoop);
|
||||
Exp(tmp, tmp, pregLoop);
|
||||
Adds(tmp, tmp, 1.0f, pregLoop);
|
||||
Div(currF, state0F, tmp, pregLoop);
|
||||
DataCopy(currFAddr + j * V_LENGTH, currF, pregLoop);
|
||||
} else {
|
||||
DataCopy(state0FAddr + j * V_LENGTH, state0F, pregLoop);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
static __simd_vf__ inline void AdvanceFnLocalPartialsWidthTwo(__ubuf__ T* ringAddr, __ubuf__ float* weight0FAddr,
|
||||
__ubuf__ float* state0FAddr, uint32_t dataCount, uint16_t colLoopTimes)
|
||||
{
|
||||
RegTensor<T> ring;
|
||||
RegTensor<float> currF;
|
||||
RegTensor<float> weight0F;
|
||||
RegTensor<float> state0F;
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t j = 0; j < colLoopTimes; j++) {
|
||||
pregLoop = UpdateMask<float>(dataCount);
|
||||
DataCopy<T, LoadDist::DIST_UNPACK_B16>(ring, ringAddr + j * V_LENGTH);
|
||||
DataCopy(weight0F, weight0FAddr + j * V_LENGTH);
|
||||
Cast<float, T, castTraitB16ToB32>(currF, ring, pregLoop);
|
||||
Mul(state0F, currF, weight0F, pregLoop);
|
||||
DataCopy(state0FAddr + j * V_LENGTH, state0F, pregLoop);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
static __simd_vf__ inline void AdvanceFnLocalPartialsWidthThree(__ubuf__ T* ringAddr, __ubuf__ float* weight0FAddr,
|
||||
__ubuf__ float* weight1FAddr, __ubuf__ float* state0FAddr, __ubuf__ float* state1FAddr, uint32_t dataCount,
|
||||
uint16_t colLoopTimes)
|
||||
{
|
||||
RegTensor<T> ring;
|
||||
RegTensor<float> currF;
|
||||
RegTensor<float> weight0F;
|
||||
RegTensor<float> weight1F;
|
||||
RegTensor<float> state0F;
|
||||
RegTensor<float> state1F;
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t j = 0; j < colLoopTimes; j++) {
|
||||
pregLoop = UpdateMask<float>(dataCount);
|
||||
DataCopy<T, LoadDist::DIST_UNPACK_B16>(ring, ringAddr + j * V_LENGTH);
|
||||
DataCopy(state1F, state1FAddr + j * V_LENGTH);
|
||||
Cast<float, T, castTraitB16ToB32>(currF, ring, pregLoop);
|
||||
DataCopy(weight1F, weight1FAddr + j * V_LENGTH);
|
||||
Mul(state0F, currF, weight1F, pregLoop);
|
||||
DataCopy(weight0F, weight0FAddr + j * V_LENGTH);
|
||||
Add(state0F, state0F, state1F, pregLoop);
|
||||
Mul(state1F, currF, weight0F, pregLoop);
|
||||
DataCopy(state0FAddr + j * V_LENGTH, state0F, pregLoop);
|
||||
DataCopy(state1FAddr + j * V_LENGTH, state1F, pregLoop);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
static __simd_vf__ inline void AdvanceFnLocalPartialsWidthFour(__ubuf__ T* ringAddr, __ubuf__ float* weight0FAddr,
|
||||
__ubuf__ float* weight1FAddr, __ubuf__ float* weight2FAddr, __ubuf__ float* state0FAddr, __ubuf__ float* state1FAddr,
|
||||
__ubuf__ float* state2FAddr, uint32_t dataCount, uint16_t colLoopTimes)
|
||||
{
|
||||
RegTensor<T> ring;
|
||||
RegTensor<float> currF;
|
||||
RegTensor<float> weight0F;
|
||||
RegTensor<float> weight1F;
|
||||
RegTensor<float> weight2F;
|
||||
RegTensor<float> state0F;
|
||||
RegTensor<float> state1F;
|
||||
RegTensor<float> state2F;
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t j = 0; j < colLoopTimes; j++) {
|
||||
pregLoop = UpdateMask<float>(dataCount);
|
||||
DataCopy<T, LoadDist::DIST_UNPACK_B16>(ring, ringAddr + j * V_LENGTH);
|
||||
DataCopy(state1F, state1FAddr + j * V_LENGTH);
|
||||
DataCopy(state2F, state2FAddr + j * V_LENGTH);
|
||||
Cast<float, T, castTraitB16ToB32>(currF, ring, pregLoop);
|
||||
DataCopy(weight2F, weight2FAddr + j * V_LENGTH);
|
||||
Mul(state0F, currF, weight2F, pregLoop);
|
||||
DataCopy(weight1F, weight1FAddr + j * V_LENGTH);
|
||||
Add(state0F, state0F, state1F, pregLoop);
|
||||
Mul(state1F, currF, weight1F, pregLoop);
|
||||
DataCopy(weight0F, weight0FAddr + j * V_LENGTH);
|
||||
Add(state1F, state1F, state2F, pregLoop);
|
||||
Mul(state2F, currF, weight0F, pregLoop);
|
||||
DataCopy(state0FAddr + j * V_LENGTH, state0F, pregLoop);
|
||||
DataCopy(state1FAddr + j * V_LENGTH, state1F, pregLoop);
|
||||
DataCopy(state2FAddr + j * V_LENGTH, state2F, pregLoop);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, int32_t kTemplateWidth>
|
||||
__aicore__ inline void AdvanceFnLocalPartialsRegbase(LocalTensor<T> ring, LocalTensor<float> weightF,
|
||||
LocalTensor<float> state0F, LocalTensor<float> state1F, LocalTensor<float> state2F, uint32_t dataCount,
|
||||
uint32_t weightStep)
|
||||
{
|
||||
uint16_t colLoopTimes = static_cast<uint16_t>(Ceil(dataCount, V_LENGTH));
|
||||
|
||||
__ubuf__ T* ringAddr = (__ubuf__ T*)ring.GetPhyAddr();
|
||||
__ubuf__ float* weight0FAddr = (__ubuf__ float*)weightF.GetPhyAddr();
|
||||
__ubuf__ float* state0FAddr = (__ubuf__ float*)state0F.GetPhyAddr();
|
||||
if constexpr (kTemplateWidth == 2) {
|
||||
AscendC::VF_CALL<AdvanceFnLocalPartialsWidthTwo<T>>(ringAddr, weight0FAddr, state0FAddr, dataCount, colLoopTimes);
|
||||
} else if constexpr (kTemplateWidth == 3) {
|
||||
__ubuf__ float* weight1FAddr = weight0FAddr + weightStep;
|
||||
__ubuf__ float* state1FAddr = (__ubuf__ float*)state1F.GetPhyAddr();
|
||||
AscendC::VF_CALL<AdvanceFnLocalPartialsWidthThree<T>>(ringAddr, weight0FAddr, weight1FAddr, state0FAddr,
|
||||
state1FAddr, dataCount, colLoopTimes);
|
||||
} else if constexpr (kTemplateWidth == 4) {
|
||||
__ubuf__ float* weight1FAddr = weight0FAddr + weightStep;
|
||||
__ubuf__ float* weight2FAddr = weight1FAddr + weightStep;
|
||||
__ubuf__ float* state1FAddr = (__ubuf__ float*)state1F.GetPhyAddr();
|
||||
__ubuf__ float* state2FAddr = (__ubuf__ float*)state2F.GetPhyAddr();
|
||||
AscendC::VF_CALL<AdvanceFnLocalPartialsWidthFour<T>>(ringAddr, weight0FAddr, weight1FAddr, weight2FAddr,
|
||||
state0FAddr, state1FAddr, state2FAddr, dataCount, colLoopTimes);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace NsCausalConv1d
|
||||
|
||||
#endif // CAUSAL_CONV1D_REGBASE_H
|
||||
57
csrc/moe/causal_conv1d/op_kernel/causal_conv1d.cpp
Normal file
57
csrc/moe/causal_conv1d/op_kernel/causal_conv1d.cpp
Normal file
@@ -0,0 +1,57 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 causal_conv1d.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "causal_conv1d_fn.h"
|
||||
#include "causal_conv1d_update.h"
|
||||
|
||||
namespace {
|
||||
|
||||
template <typename T, uint32_t runModeKey, uint32_t widthKey, uint32_t fnPlanKey>
|
||||
__aicore__ inline void RunCausalConv1d(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates,
|
||||
GM_ADDR queryStartLoc, GM_ADDR cacheIndices, GM_ADDR initialStateMode,
|
||||
GM_ADDR numAcceptedTokens, GM_ADDR y, GM_ADDR workspace,
|
||||
const CausalConv1dTilingData *tilingData)
|
||||
{
|
||||
if constexpr (runModeKey == CAUSAL_CONV1D_TPL_RUN_MODE_FN) {
|
||||
NsCausalConv1d::RunCausalConv1dFn<T, widthKey, fnPlanKey>(
|
||||
x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y, workspace,
|
||||
tilingData);
|
||||
} else {
|
||||
NsCausalConv1d::RunCausalConv1dUpdate<T>(
|
||||
x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y, workspace,
|
||||
tilingData);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
template <uint32_t runModeKey, uint32_t widthKey, uint32_t fnPlanKey>
|
||||
__global__ __aicore__ void causal_conv1d(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates,
|
||||
GM_ADDR queryStartLoc, GM_ADDR cacheIndices, GM_ADDR initialStateMode,
|
||||
GM_ADDR numAcceptedTokens, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
|
||||
{
|
||||
REGISTER_TILING_DEFAULT(CausalConv1dTilingData);
|
||||
GET_TILING_DATA(tilingData, tiling);
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIV_1_0);
|
||||
GM_ADDR userWorkspace = workspace;
|
||||
if (workspace != nullptr) {
|
||||
userWorkspace = AscendC::GetUserWorkspace(workspace);
|
||||
}
|
||||
|
||||
RunCausalConv1d<DTYPE_X, runModeKey, widthKey, fnPlanKey>(
|
||||
x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y,
|
||||
userWorkspace, &tilingData);
|
||||
}
|
||||
1088
csrc/moe/causal_conv1d/op_kernel/causal_conv1d.h
Normal file
1088
csrc/moe/causal_conv1d/op_kernel/causal_conv1d.h
Normal file
File diff suppressed because it is too large
Load Diff
66
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_common.h
Normal file
66
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_common.h
Normal file
@@ -0,0 +1,66 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 causal_conv1d_common.h
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_COMMON_H
|
||||
#define CAUSAL_CONV1D_COMMON_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
|
||||
namespace NsCausalConv1dCommon {
|
||||
|
||||
constexpr int32_t MAX_WIDTH = 4;
|
||||
constexpr int32_t MAX_BLOCK_DIM = 4096;
|
||||
constexpr int32_t RING_SLOTS = 5;
|
||||
|
||||
__aicore__ inline int32_t SlotCurr(int32_t t)
|
||||
{
|
||||
return (t + 3) % RING_SLOTS;
|
||||
}
|
||||
|
||||
__aicore__ inline int32_t SlotHist(int32_t t, int32_t i)
|
||||
{
|
||||
return (t + 3 - i) % RING_SLOTS;
|
||||
}
|
||||
|
||||
__aicore__ inline int32_t SlotPrefetch(int32_t t)
|
||||
{
|
||||
return (t + 4) % RING_SLOTS;
|
||||
}
|
||||
|
||||
struct CalcBufLayout {
|
||||
AscendC::LocalTensor<float> weightF;
|
||||
AscendC::LocalTensor<float> biasF;
|
||||
AscendC::LocalTensor<float> accF;
|
||||
AscendC::LocalTensor<float> tmpF;
|
||||
AscendC::LocalTensor<float> currF;
|
||||
|
||||
__aicore__ inline CalcBufLayout() = default;
|
||||
|
||||
__aicore__ static inline CalcBufLayout FromCalcBuf(AscendC::TBuf<AscendC::QuePosition::VECCALC> &calcBuf)
|
||||
{
|
||||
CalcBufLayout layout;
|
||||
AscendC::LocalTensor<float> calc = calcBuf.template Get<float>();
|
||||
layout.weightF = calc;
|
||||
layout.biasF = calc[MAX_WIDTH * MAX_BLOCK_DIM];
|
||||
layout.accF = layout.biasF[MAX_BLOCK_DIM];
|
||||
layout.tmpF = layout.accF[MAX_BLOCK_DIM];
|
||||
layout.currF = layout.tmpF[MAX_BLOCK_DIM];
|
||||
return layout;
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace NsCausalConv1dCommon
|
||||
|
||||
#endif // CAUSAL_CONV1D_COMMON_H
|
||||
93
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_fn.h
Normal file
93
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_fn.h
Normal file
@@ -0,0 +1,93 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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.
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_FN_H
|
||||
#define CAUSAL_CONV1D_FN_H
|
||||
|
||||
#include "causal_conv1d.h"
|
||||
|
||||
namespace NsCausalConv1d {
|
||||
|
||||
template <typename T, uint32_t widthKey, uint32_t fnPlanKey>
|
||||
class CausalConv1dFn : public CausalConv1d<T, CAUSAL_CONV1D_TPL_RUN_MODE_FN, widthKey, fnPlanKey> {
|
||||
public:
|
||||
__aicore__ inline void Init(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates, GM_ADDR queryStartLoc,
|
||||
GM_ADDR cacheIndices, GM_ADDR initialStateMode, GM_ADDR numAcceptedTokens, GM_ADDR y,
|
||||
GM_ADDR workspace, const CausalConv1dTilingData *tilingData)
|
||||
{
|
||||
(void)numAcceptedTokens;
|
||||
this->ResetRuntimeState(tilingData);
|
||||
this->xGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(x));
|
||||
this->weightGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(weight));
|
||||
this->biasGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(bias));
|
||||
this->convStatesGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(convStates));
|
||||
if (tilingData->hasQueryStartLoc != 0) {
|
||||
if (tilingData->queryStartLocUseInt64 != 0) {
|
||||
this->queryStartLocGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(queryStartLoc));
|
||||
} else {
|
||||
this->queryStartLocGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(queryStartLoc));
|
||||
}
|
||||
}
|
||||
if (tilingData->hasCacheIndices != 0) {
|
||||
if (tilingData->cacheIndicesUseInt64 != 0) {
|
||||
this->cacheIndicesGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(cacheIndices));
|
||||
} else {
|
||||
this->cacheIndicesGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(cacheIndices));
|
||||
}
|
||||
}
|
||||
if (tilingData->hasInitialStateMode != 0) {
|
||||
if (tilingData->initialStateModeDtype == 2) {
|
||||
this->initialStateModeGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(initialStateMode));
|
||||
} else if (tilingData->initialStateModeDtype == 1) {
|
||||
this->initialStateModeGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(initialStateMode));
|
||||
} else {
|
||||
this->initialStateModeGmBool.SetGlobalBuffer(reinterpret_cast<__gm__ bool *>(initialStateMode));
|
||||
}
|
||||
}
|
||||
this->yGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(y));
|
||||
if (tilingData->hasInitStateWorkspace != 0) {
|
||||
const uint64_t syncElems =
|
||||
static_cast<uint64_t>(GetBlockNum()) * INIT_STATE_SYNCALL_NEED_SIZE;
|
||||
const uint64_t syncBytes = syncElems * sizeof(int32_t);
|
||||
const uint64_t workspaceElems =
|
||||
static_cast<uint64_t>(tilingData->batch) *
|
||||
static_cast<uint64_t>(tilingData->width - 1) *
|
||||
static_cast<uint64_t>(tilingData->dim);
|
||||
this->initStateSyncGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(workspace), syncElems);
|
||||
auto *workspaceBytes = reinterpret_cast<__gm__ uint8_t *>(workspace);
|
||||
this->initStateWorkspaceGm_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(workspaceBytes + syncBytes),
|
||||
workspaceElems);
|
||||
}
|
||||
this->InitSharedBuffersAndEvents();
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
this->ProcessVarlenTokenTiled();
|
||||
this->ReleaseEvents();
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T, uint32_t widthKey, uint32_t fnPlanKey>
|
||||
__aicore__ inline void RunCausalConv1dFn(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates,
|
||||
GM_ADDR queryStartLoc, GM_ADDR cacheIndices, GM_ADDR initialStateMode,
|
||||
GM_ADDR numAcceptedTokens, GM_ADDR y, GM_ADDR workspace,
|
||||
const CausalConv1dTilingData *tilingData)
|
||||
{
|
||||
CausalConv1dFn<T, widthKey, fnPlanKey> op;
|
||||
op.Init(x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y, workspace,
|
||||
tilingData);
|
||||
op.Process();
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
#endif
|
||||
306
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_fn_tasks.h
Normal file
306
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_fn_tasks.h
Normal file
@@ -0,0 +1,306 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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.
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_FN_TASKS_H
|
||||
#define CAUSAL_CONV1D_FN_TASKS_H
|
||||
|
||||
struct FnDirectBlockTask {
|
||||
bool valid = false;
|
||||
int32_t tokenTileId = 0;
|
||||
int32_t baseDimIdx = 0;
|
||||
int32_t tokenStart = 0;
|
||||
int32_t tokenEnd = 0;
|
||||
int32_t channelStart = 0;
|
||||
int32_t baseDimSize = 0;
|
||||
};
|
||||
|
||||
__aicore__ inline FnDirectBlockTask ResolveFnDirectBlockTask(int32_t blockIdx, int32_t tokenBlockCnt, int32_t tokenBlockSize,
|
||||
int32_t cuSeqlen, int32_t baseDimCnt, int32_t baseDim,
|
||||
int32_t dim)
|
||||
{
|
||||
FnDirectBlockTask task;
|
||||
if (blockIdx < 0 || tokenBlockCnt <= 0 || tokenBlockSize <= 0 || cuSeqlen <= 0 || baseDimCnt <= 0 || baseDim <= 0 ||
|
||||
dim <= 0) {
|
||||
return task;
|
||||
}
|
||||
|
||||
const int64_t phase1Grid = static_cast<int64_t>(tokenBlockCnt) * baseDimCnt;
|
||||
if (phase1Grid <= 0 || static_cast<int64_t>(blockIdx) >= phase1Grid) {
|
||||
return task;
|
||||
}
|
||||
|
||||
task.tokenTileId = blockIdx / baseDimCnt;
|
||||
task.baseDimIdx = blockIdx % baseDimCnt;
|
||||
task.channelStart = task.baseDimIdx * baseDim;
|
||||
if (task.channelStart >= dim) {
|
||||
return task;
|
||||
}
|
||||
|
||||
task.baseDimSize = (task.channelStart + baseDim <= dim) ? baseDim : (dim - task.channelStart);
|
||||
task.tokenStart = task.tokenTileId * tokenBlockSize;
|
||||
if (task.tokenStart >= cuSeqlen) {
|
||||
return task;
|
||||
}
|
||||
|
||||
const int32_t tokenEndRaw = task.tokenStart + tokenBlockSize;
|
||||
task.tokenEnd = (tokenEndRaw <= cuSeqlen) ? tokenEndRaw : cuSeqlen;
|
||||
if (task.baseDimSize <= 0 || task.tokenEnd <= task.tokenStart) {
|
||||
return {};
|
||||
}
|
||||
|
||||
task.valid = true;
|
||||
return task;
|
||||
}
|
||||
|
||||
__aicore__ inline bool IsFnInitStateSnapshotOwnerBlock(const FnDirectBlockTask &task)
|
||||
{
|
||||
return task.valid && task.tokenTileId == 0;
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline int32_t CAUSAL_CONV1D_CLASS::FindVarlenSeqByToken(int32_t tokenIdx) const
|
||||
{
|
||||
int32_t left = 0;
|
||||
int32_t right = static_cast<int32_t>(tilingData_->batch);
|
||||
while (left < right) {
|
||||
const int32_t mid = left + ((right - left) >> 1);
|
||||
const int32_t endVal = ReadQueryStartLocValue(mid + 1);
|
||||
if (tokenIdx < endVal) {
|
||||
right = mid;
|
||||
} else {
|
||||
left = mid + 1;
|
||||
}
|
||||
}
|
||||
return left;
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline bool CAUSAL_CONV1D_CLASS::ResolveExplicitTokenTileSeqRange(int32_t tokenTileId, int32_t &startSeq,
|
||||
int32_t &endSeq) const
|
||||
{
|
||||
if (!HasExplicitFnTokenSeqRanges() || tokenTileId < 0 || tokenTileId >= tilingData_->explicitTokenSeqRangeCount) {
|
||||
return false;
|
||||
}
|
||||
startSeq = static_cast<int32_t>(tilingData_->tokenTileStartSeq[tokenTileId]);
|
||||
endSeq = static_cast<int32_t>(tilingData_->tokenTileEndSeq[tokenTileId]);
|
||||
return (startSeq >= 0) && (endSeq >= startSeq);
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline void CAUSAL_CONV1D_CLASS::InitRingSeqSplit(int32_t seq, int32_t cacheIdx, bool hasInit,
|
||||
int32_t seqStart, int32_t tileStart, int32_t tileLen,
|
||||
int32_t channelStart, int32_t baseDim, int32_t dim)
|
||||
{
|
||||
const int32_t stateLen = tilingData_->stateLen;
|
||||
const int32_t width = static_cast<int32_t>(tilingData_->width);
|
||||
const int32_t historyCount = width - 1;
|
||||
const int32_t ringStart = MAX_WIDTH - width;
|
||||
const int32_t historyStartTok = tileStart - historyCount;
|
||||
LocalTensor<T> ring = inBuf.Get<T>();
|
||||
bool hasGmHistoryCopy = false;
|
||||
bool hasVectorInit = false;
|
||||
const int64_t stateBaseOffset = static_cast<int64_t>(cacheIdx) * stateLen * dim + channelStart;
|
||||
int64_t xHistoryOffset = static_cast<int64_t>(historyStartTok) * dim + channelStart;
|
||||
|
||||
for (int32_t i = 0; i < ringStart; ++i) {
|
||||
Duplicate(ring[i * MAX_BLOCK_DIM], static_cast<T>(0), baseDim);
|
||||
hasVectorInit = true;
|
||||
}
|
||||
|
||||
for (int32_t i = 0, srcTok = historyStartTok; i < historyCount; ++i, ++srcTok, xHistoryOffset += dim) {
|
||||
LocalTensor<T> histSlot = ring[(ringStart + i) * MAX_BLOCK_DIM];
|
||||
if (srcTok >= seqStart) {
|
||||
DataCopy(histSlot, xGm[xHistoryOffset], baseDim);
|
||||
hasGmHistoryCopy = true;
|
||||
} else if (hasInit) {
|
||||
const int32_t statePos = srcTok - seqStart + historyCount;
|
||||
const int64_t stateOffset = stateBaseOffset + static_cast<int64_t>(statePos) * dim;
|
||||
if (tilingData_->hasInitStateWorkspace != 0) {
|
||||
const int64_t snapshotOffset =
|
||||
(static_cast<int64_t>(seq) * historyCount + statePos) * dim + channelStart;
|
||||
DataCopy(histSlot, initStateWorkspaceGm_[snapshotOffset], baseDim);
|
||||
} else {
|
||||
DataCopy(histSlot, convStatesGm[stateOffset], baseDim);
|
||||
}
|
||||
hasGmHistoryCopy = true;
|
||||
} else {
|
||||
Duplicate(histSlot, static_cast<T>(0), baseDim);
|
||||
hasVectorInit = true;
|
||||
}
|
||||
}
|
||||
|
||||
if (hasGmHistoryCopy) {
|
||||
SetFlag<HardEvent::MTE2_V>(stateMte2ToVEvent_);
|
||||
WaitFlag<HardEvent::MTE2_V>(stateMte2ToVEvent_);
|
||||
}
|
||||
if (hasVectorInit) {
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
if (tileLen > 0) {
|
||||
const int32_t slot0 = SlotCurr(0);
|
||||
const int64_t xOffset = static_cast<int64_t>(tileStart) * dim + channelStart;
|
||||
DataCopy(ring[slot0 * MAX_BLOCK_DIM], xGm[xOffset], baseDim);
|
||||
SetFlag<HardEvent::MTE2_V>(inputMte2ToVEvent_[slot0]);
|
||||
}
|
||||
|
||||
if (tileLen > 1) {
|
||||
SetFlag<HardEvent::V_MTE2>(inputVToMte2Event_);
|
||||
}
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline void CAUSAL_CONV1D_CLASS::ProcessFnChunk(int32_t seq, int32_t cacheIdx, bool hasInit,
|
||||
int32_t seqStart, int32_t seqLen, int32_t chunkStart,
|
||||
int32_t chunkLen, int32_t channelStart, int32_t baseDim,
|
||||
int32_t dim)
|
||||
{
|
||||
LoadWeightAndBias(channelStart, baseDim);
|
||||
InitRingSeqSplit(seq, cacheIdx, hasInit, seqStart, chunkStart, chunkLen, channelStart, baseDim, dim);
|
||||
|
||||
RunSeq(chunkStart, chunkLen, channelStart, baseDim, dim);
|
||||
|
||||
MaybeWriteBackSeqSplitTailChunk(chunkStart, chunkLen, seqStart, seqLen, cacheIdx, channelStart, baseDim, dim);
|
||||
DrainTaskMte3();
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline void
|
||||
CAUSAL_CONV1D_CLASS::MaybeWriteBackSeqSplitTailChunk(int32_t chunkStart, int32_t chunkLen, int32_t seqStart,
|
||||
int32_t seqLen, int32_t cacheIdx, int32_t channelStart,
|
||||
int32_t baseDim, int32_t dim)
|
||||
{
|
||||
if (chunkStart + chunkLen != seqStart + seqLen) {
|
||||
return;
|
||||
}
|
||||
|
||||
DrainTaskMte3();
|
||||
WriteBackState(cacheIdx, chunkLen, channelStart, baseDim, dim);
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline void CAUSAL_CONV1D_CLASS::PrefetchInitStatesToWorkspace(int32_t channelStart, int32_t baseDimSize)
|
||||
{
|
||||
if (tilingData_->hasInitStateWorkspace == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
const int32_t dim = tilingData_->dim;
|
||||
const int32_t historyCount = static_cast<int32_t>(tilingData_->width - 1);
|
||||
const int32_t batch = tilingData_->batch;
|
||||
const bool hasCacheIndices = (tilingData_->hasCacheIndices != 0);
|
||||
const bool hasInitialStateMode = (tilingData_->hasInitialStateMode != 0);
|
||||
LocalTensor<T> tmpBuf = inBuf.Get<T>()[0 * MAX_BLOCK_DIM];
|
||||
|
||||
for (int32_t seq = 0; seq < batch; ++seq) {
|
||||
if (!ResolveSeqHasInit(seq, hasInitialStateMode)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
int32_t cacheIdx = 0;
|
||||
if (!ResolveSeqCacheIndex(seq, hasCacheIndices, cacheIdx)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const int64_t stateBaseOffset = static_cast<int64_t>(cacheIdx) * tilingData_->stateLen * dim + channelStart;
|
||||
const int64_t snapshotBaseOffset = static_cast<int64_t>(seq) * historyCount * dim + channelStart;
|
||||
for (int32_t statePos = 0; statePos < historyCount; ++statePos) {
|
||||
const int64_t stateOffset = stateBaseOffset + static_cast<int64_t>(statePos) * dim;
|
||||
const int64_t snapshotOffset = snapshotBaseOffset + static_cast<int64_t>(statePos) * dim;
|
||||
DataCopy(tmpBuf, convStatesGm[stateOffset], baseDimSize);
|
||||
SetFlag<HardEvent::MTE2_MTE3>(initSnapshotMte2ToMte3Event_);
|
||||
WaitFlag<HardEvent::MTE2_MTE3>(initSnapshotMte2ToMte3Event_);
|
||||
DataCopy(initStateWorkspaceGm_[snapshotOffset], tmpBuf, baseDimSize);
|
||||
SetFlag<HardEvent::MTE3_MTE2>(initSnapshotMte3ToMte2Event_);
|
||||
WaitFlag<HardEvent::MTE3_MTE2>(initSnapshotMte3ToMte2Event_);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline void CAUSAL_CONV1D_CLASS::ProcessVarlenTokenTiled()
|
||||
{
|
||||
const int32_t dim = tilingData_->dim;
|
||||
const int32_t batch = tilingData_->batch;
|
||||
const int32_t seqLen = tilingData_->seqLen;
|
||||
const int32_t cuSeqlen = tilingData_->cuSeqlen;
|
||||
const int32_t baseDim = static_cast<int32_t>(tilingData_->baseDim);
|
||||
const int32_t baseDimCnt = static_cast<int32_t>(tilingData_->baseDimCnt);
|
||||
const int32_t tokenBlockSize = static_cast<int32_t>(tilingData_->tokenBlockSize);
|
||||
const int32_t tokenBlockCnt = static_cast<int32_t>(tilingData_->tokenBlockCnt);
|
||||
const bool hasCacheIndices = (tilingData_->hasCacheIndices != 0);
|
||||
const bool hasInitialStateMode = (tilingData_->hasInitialStateMode != 0);
|
||||
const bool isVarlenMode = (tilingData_->inputMode == 0);
|
||||
|
||||
const int32_t blockIdx = static_cast<int32_t>(GetBlockIdx());
|
||||
const auto blockTask = ResolveFnDirectBlockTask(blockIdx, tokenBlockCnt, tokenBlockSize, cuSeqlen, baseDimCnt,
|
||||
baseDim, dim);
|
||||
if (tilingData_->hasInitStateWorkspace != 0) {
|
||||
if (IsFnInitStateSnapshotOwnerBlock(blockTask)) {
|
||||
PrefetchInitStatesToWorkspace(blockTask.channelStart, blockTask.baseDimSize);
|
||||
}
|
||||
SyncAll();
|
||||
}
|
||||
if (!blockTask.valid) {
|
||||
return;
|
||||
}
|
||||
|
||||
int32_t seq = 0;
|
||||
int32_t seqUpperBound = batch;
|
||||
if (isVarlenMode) {
|
||||
if (!ResolveExplicitTokenTileSeqRange(blockTask.tokenTileId, seq, seqUpperBound)) {
|
||||
seq = FindVarlenSeqByToken(blockTask.tokenStart);
|
||||
}
|
||||
} else {
|
||||
seq = (seqLen > 0) ? (blockTask.tokenStart / seqLen) : 0;
|
||||
}
|
||||
|
||||
int32_t cursor = blockTask.tokenStart;
|
||||
while (cursor < blockTask.tokenEnd && seq < seqUpperBound) {
|
||||
int32_t seqStart = 0;
|
||||
int32_t curSeqLen = 0;
|
||||
if (!ResolveSeqTaskWindow(seq, tilingData_->inputMode, seqLen, seqStart, curSeqLen)) {
|
||||
++seq;
|
||||
continue;
|
||||
}
|
||||
const int32_t curSeqEnd = seqStart + curSeqLen;
|
||||
if (cursor < seqStart) {
|
||||
cursor = seqStart;
|
||||
}
|
||||
if (cursor >= curSeqEnd) {
|
||||
++seq;
|
||||
continue;
|
||||
}
|
||||
|
||||
const int32_t tileEnd = (blockTask.tokenEnd <= curSeqEnd) ? blockTask.tokenEnd : curSeqEnd;
|
||||
const int32_t tileLen = tileEnd - cursor;
|
||||
if (tileLen <= 0) {
|
||||
++seq;
|
||||
continue;
|
||||
}
|
||||
|
||||
int32_t cacheIdx = 0;
|
||||
if (!ResolveSeqCacheIndex(seq, hasCacheIndices, cacheIdx)) {
|
||||
cursor = tileEnd;
|
||||
++seq;
|
||||
continue;
|
||||
}
|
||||
|
||||
const bool hasInit = ResolveSeqHasInit(seq, hasInitialStateMode);
|
||||
ProcessFnChunk(seq, cacheIdx, hasInit, seqStart, curSeqLen, cursor, tileLen, blockTask.channelStart,
|
||||
blockTask.baseDimSize, dim);
|
||||
|
||||
cursor = tileEnd;
|
||||
++seq;
|
||||
}
|
||||
}
|
||||
|
||||
#endif
|
||||
68
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_tiling_data.h
Normal file
68
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_tiling_data.h
Normal file
@@ -0,0 +1,68 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 causal_conv1d_tiling_data.h
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_TILING_DATA_H_
|
||||
#define CAUSAL_CONV1D_TILING_DATA_H_
|
||||
|
||||
#include <cstdint>
|
||||
|
||||
enum FnExecutionPlan : int64_t {
|
||||
FN_EXECUTION_PLAN_INVALID = 0,
|
||||
FN_EXECUTION_PLAN_CUTBS = 1,
|
||||
FN_EXECUTION_PLAN_CUTBSD = 2,
|
||||
};
|
||||
|
||||
inline constexpr int64_t ResolveFnExecutionPlan(int64_t baseDimCnt)
|
||||
{
|
||||
return (baseDimCnt <= 0) ? FN_EXECUTION_PLAN_INVALID
|
||||
: (baseDimCnt <= 1) ? FN_EXECUTION_PLAN_CUTBS
|
||||
: FN_EXECUTION_PLAN_CUTBSD;
|
||||
}
|
||||
|
||||
|
||||
struct CausalConv1dTilingData {
|
||||
int64_t dim;
|
||||
int64_t cuSeqlen;
|
||||
int64_t seqLen;
|
||||
int64_t inputMode;
|
||||
|
||||
int64_t width;
|
||||
|
||||
int64_t stateLen;
|
||||
int64_t numCacheLines;
|
||||
int64_t batch;
|
||||
int64_t activationMode;
|
||||
int64_t padSlotId;
|
||||
int64_t hasBias;
|
||||
int64_t baseDim;
|
||||
int64_t baseDimCnt;
|
||||
int64_t hasNumAcceptedTokens;
|
||||
int64_t hasQueryStartLoc;
|
||||
int64_t hasCacheIndices;
|
||||
int64_t hasInitialStateMode;
|
||||
int64_t queryStartLocUseInt64;
|
||||
int64_t cacheIndicesStride;
|
||||
int64_t cacheIndicesUseInt64;
|
||||
int64_t initialStateModeDtype;
|
||||
int64_t numAcceptedTokensUseInt64;
|
||||
int64_t tokenBlockSize;
|
||||
int64_t tokenBlockCnt;
|
||||
int64_t hasExplicitTokenSeqRanges;
|
||||
int64_t explicitTokenSeqRangeCount;
|
||||
int64_t tokenTileStartSeq[128];
|
||||
int64_t tokenTileEndSeq[128];
|
||||
int64_t hasInitStateWorkspace;
|
||||
};
|
||||
#endif // CAUSAL_CONV1D_TILING_DATA_H_
|
||||
66
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_tiling_key.h
Normal file
66
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_tiling_key.h
Normal file
@@ -0,0 +1,66 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 causal_conv1d_tiling_key.h
|
||||
* \brief causal_conv1d tiling key declare
|
||||
*/
|
||||
|
||||
#ifndef __CAUSAL_CONV1D_TILING_KEY_H__
|
||||
#define __CAUSAL_CONV1D_TILING_KEY_H__
|
||||
|
||||
#include "causal_conv1d_tiling_data.h"
|
||||
#include "ascendc/host_api/tiling/template_argument.h"
|
||||
|
||||
#define CAUSAL_CONV1D_TPL_RUN_MODE_FN 0
|
||||
#define CAUSAL_CONV1D_TPL_RUN_MODE_UPDATE 1
|
||||
#define CAUSAL_CONV1D_TPL_WIDTH_RUNTIME 0
|
||||
#define CAUSAL_CONV1D_TPL_WIDTH_2 1
|
||||
#define CAUSAL_CONV1D_TPL_WIDTH_3 2
|
||||
#define CAUSAL_CONV1D_TPL_WIDTH_4 3
|
||||
#define CAUSAL_CONV1D_TPL_FN_PLAN_INVALID 0
|
||||
#define CAUSAL_CONV1D_TPL_FN_PLAN_CUTBS 1
|
||||
#define CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD 2
|
||||
ASCENDC_TPL_ARGS_DECL(CausalConv1d,
|
||||
ASCENDC_TPL_UINT_DECL(runModeKey, 1, ASCENDC_TPL_UI_LIST, CAUSAL_CONV1D_TPL_RUN_MODE_FN,
|
||||
CAUSAL_CONV1D_TPL_RUN_MODE_UPDATE),
|
||||
ASCENDC_TPL_UINT_DECL(widthKey, 2, ASCENDC_TPL_UI_LIST, CAUSAL_CONV1D_TPL_WIDTH_RUNTIME,
|
||||
CAUSAL_CONV1D_TPL_WIDTH_2, CAUSAL_CONV1D_TPL_WIDTH_3,
|
||||
CAUSAL_CONV1D_TPL_WIDTH_4),
|
||||
ASCENDC_TPL_UINT_DECL(fnPlanKey, 2, ASCENDC_TPL_UI_LIST, CAUSAL_CONV1D_TPL_FN_PLAN_INVALID,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBS, CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD));
|
||||
|
||||
#define CAUSAL_CONV1D_TPL_SEL_ENTRY(RUN_MODE, WIDTH, FN_PLAN) \
|
||||
ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(runModeKey, ASCENDC_TPL_UI_LIST, RUN_MODE), \
|
||||
ASCENDC_TPL_UINT_SEL(widthKey, ASCENDC_TPL_UI_LIST, WIDTH), \
|
||||
ASCENDC_TPL_UINT_SEL(fnPlanKey, ASCENDC_TPL_UI_LIST, FN_PLAN), \
|
||||
ASCENDC_TPL_TILING_STRUCT_SEL(CausalConv1dTilingData))
|
||||
|
||||
// Keep entries in encoded tiling-key order: real-device sub-kernel dispatch is sensitive to declaration order.
|
||||
ASCENDC_TPL_SEL(
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_UPDATE, CAUSAL_CONV1D_TPL_WIDTH_RUNTIME,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_INVALID),
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_WIDTH_2,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBS),
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_WIDTH_3,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBS),
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_WIDTH_4,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBS),
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_WIDTH_2,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD),
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_WIDTH_3,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD),
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_WIDTH_4,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD));
|
||||
|
||||
#undef CAUSAL_CONV1D_TPL_SEL_ENTRY
|
||||
|
||||
#endif // __CAUSAL_CONV1D_TILING_KEY_H__
|
||||
91
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_update.h
Normal file
91
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_update.h
Normal file
@@ -0,0 +1,91 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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.
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_UPDATE_H
|
||||
#define CAUSAL_CONV1D_UPDATE_H
|
||||
|
||||
#include "causal_conv1d.h"
|
||||
|
||||
namespace NsCausalConv1d {
|
||||
|
||||
template <typename T>
|
||||
class CausalConv1dUpdate
|
||||
: public CausalConv1d<T, CAUSAL_CONV1D_TPL_RUN_MODE_UPDATE, CAUSAL_CONV1D_TPL_WIDTH_RUNTIME,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_INVALID> {
|
||||
public:
|
||||
__aicore__ inline void Init(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates, GM_ADDR queryStartLoc,
|
||||
GM_ADDR cacheIndices, GM_ADDR, GM_ADDR numAcceptedTokens, GM_ADDR y, GM_ADDR workspace,
|
||||
const CausalConv1dTilingData *tilingData)
|
||||
{
|
||||
(void)workspace;
|
||||
this->ResetRuntimeState(tilingData);
|
||||
this->xGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(x));
|
||||
this->weightGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(weight));
|
||||
this->biasGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(bias));
|
||||
this->convStatesGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(convStates));
|
||||
if (tilingData->hasQueryStartLoc != 0) {
|
||||
if (tilingData->queryStartLocUseInt64 != 0) {
|
||||
this->queryStartLocGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(queryStartLoc));
|
||||
} else {
|
||||
this->queryStartLocGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(queryStartLoc));
|
||||
}
|
||||
}
|
||||
if (tilingData->hasCacheIndices != 0) {
|
||||
if (tilingData->cacheIndicesUseInt64 != 0) {
|
||||
this->cacheIndicesGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(cacheIndices));
|
||||
} else {
|
||||
this->cacheIndicesGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(cacheIndices));
|
||||
}
|
||||
}
|
||||
if (tilingData->hasNumAcceptedTokens != 0) {
|
||||
if (tilingData->numAcceptedTokensUseInt64 != 0) {
|
||||
this->numAcceptedTokensGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(numAcceptedTokens));
|
||||
} else {
|
||||
this->numAcceptedTokensGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(numAcceptedTokens));
|
||||
}
|
||||
}
|
||||
this->yGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(y));
|
||||
this->InitSharedBuffersAndEvents();
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
const CausalConv1dTilingData *tilingData = this->GetTilingData();
|
||||
const int32_t dim = tilingData->dim;
|
||||
const int32_t baseDimCnt = static_cast<int32_t>(tilingData->baseDimCnt);
|
||||
const int32_t width = static_cast<int32_t>(tilingData->width);
|
||||
const int32_t baseDim = static_cast<int32_t>(tilingData->baseDim);
|
||||
if (baseDim <= 0 || baseDimCnt <= 0 || baseDim > MAX_BLOCK_DIM || width < 2 || width > MAX_WIDTH || dim <= 0 ||
|
||||
tilingData->batch <= 0) {
|
||||
this->ReleaseEvents();
|
||||
return;
|
||||
}
|
||||
|
||||
this->ProcessDefault();
|
||||
this->ReleaseEvents();
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void RunCausalConv1dUpdate(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates,
|
||||
GM_ADDR queryStartLoc, GM_ADDR cacheIndices, GM_ADDR initialStateMode,
|
||||
GM_ADDR numAcceptedTokens, GM_ADDR y, GM_ADDR workspace,
|
||||
const CausalConv1dTilingData *tilingData)
|
||||
{
|
||||
CausalConv1dUpdate<T> op;
|
||||
op.Init(x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y, workspace,
|
||||
tilingData);
|
||||
op.Process();
|
||||
}
|
||||
|
||||
} // namespace NsCausalConv1d
|
||||
|
||||
#endif // CAUSAL_CONV1D_UPDATE_H
|
||||
Reference in New Issue
Block a user