init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -0,0 +1,48 @@
/**
* This program is free software, you can redistribute it and/or modify it.
* Copyright (c) 2026 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_v310.cpp
* \brief
*/
#include "causal_conv1d_v310.h"
namespace {
// NOTE:
// Dtype is provided via AscendC compile macros (e.g. DTYPE_X / ORIG_DTYPE_X), so tiling key does not need to carry dtype.
template <typename T>
__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, const CausalConv1dTilingData *tilingData)
{
NsCausalConv1d::CausalConv1dV310<T> op;
op.Init(x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y,
tilingData);
op.Process();
}
} // namespace
template <uint32_t schMode>
__global__ __aicore__ void causal_conv1d_v310(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_WITH_STRUCT(CausalConv1dTilingData, tilingData, tiling);
RunCausalConv1d<half>(x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens,
y, &tilingData);
}

View File

@@ -0,0 +1,661 @@
/**
* 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_v310.h
* \brief CausalConv1D (prefill/extend) AscendC kernel implementation.
*
*/
#ifndef CAUSAL_CONV1D_V310_H
#define CAUSAL_CONV1D_V310_H
#include "kernel_operator.h"
#include "kernel_tiling/kernel_tiling.h"
#include "causal_conv1d_v310_tiling_data.h"
#include "causal_conv1d_v310_tiling_key.h"
#include "causal_conv1d_v310_common.h"
namespace NsCausalConv1d {
using namespace AscendC;
using namespace NsCausalConv1dCommon;
inline constexpr int64_t INT32_MAX_VALUE = 2147483647LL;
template <typename T>
class CausalConv1dV310 {
public:
__aicore__ inline CausalConv1dV310() = default;
__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,
const CausalConv1dTilingData *tilingData);
__aicore__ inline void Process();
private:
__aicore__ inline void LoadWeightAndBias(int32_t c0, int32_t dimTileSize);
__aicore__ inline void InitRing(int32_t cacheIdx, bool hasInit, int32_t stateTokenOffset, int32_t start,
int32_t len, int32_t c0, int32_t dimTileSize, int32_t dim);
__aicore__ inline void RunSeq(int32_t start, int32_t len, int32_t c0, int32_t dimTileSize, int32_t dim);
__aicore__ inline void WriteBackState(int32_t cacheIdx, int32_t len, int32_t c0, int32_t dimTileSize, int32_t dim);
__aicore__ inline void WriteBackStateSpec(int32_t cacheIdx, bool hasInit, int32_t stateTokenOffset, int32_t start,
int32_t len, int32_t c0, int32_t dimTileSize, int32_t dim);
__aicore__ inline void AllocEvents();
__aicore__ inline void ReleaseEvents();
__aicore__ inline int32_t ReadQueryStartLocValue(int32_t index) const;
__aicore__ inline int64_t ReadCacheIndexValue(int32_t seq) const;
__aicore__ inline bool ReadInitialStateModeValue(int32_t seq) const;
__aicore__ inline int32_t ReadNumAcceptedTokensValue(int32_t seq) const;
private:
TPipe pipe;
TBuf<QuePosition::VECIN> inBuf;
TBuf<QuePosition::VECOUT> outBuf;
TBuf<QuePosition::VECCALC> calcBuf;
TEventID weightBiasMte2ToVEvent_;
TEventID stateMte2ToVEvent_;
TEventID inputMte2ToVEvent_[RING_SLOTS];
TEventID inputVToMte2Event_;
TEventID outMte3ToVEvent_[2];
TEventID outVToMte3Event_[2];
TEventID stateWritebackMte3ToVEvent_;
TEventID stateWritebackMte3ToMte2Event_;
TEventID stateShiftMte2ToMte3Event_;
TEventID stateShiftMte3ToMte2Event_;
TEventID stateShiftVToMte3Event_;
TEventID specWritebackMte2ToMte3Event_[2];
TEventID specWritebackMte3ToMte2Event_[2];
GlobalTensor<T> xGm;
GlobalTensor<T> weightGm;
GlobalTensor<T> biasGm;
GlobalTensor<T> convStatesGm;
GlobalTensor<int32_t> queryStartLocGmInt32;
GlobalTensor<int64_t> queryStartLocGmInt64;
GlobalTensor<int32_t> cacheIndicesGmInt32;
GlobalTensor<int64_t> cacheIndicesGmInt64;
GlobalTensor<bool> initialStateModeGmBool;
GlobalTensor<int32_t> initialStateModeGmInt32;
GlobalTensor<int64_t> initialStateModeGmInt64;
GlobalTensor<int32_t> numAcceptedTokensGmInt32;
GlobalTensor<int64_t> numAcceptedTokensGmInt64;
GlobalTensor<T> yGm;
const CausalConv1dTilingData *tilingData_{nullptr};
bool weightCacheValid_{false};
int32_t cachedC0_{-1};
int32_t cachedDimTileSize_{-1};
};
template <typename T>
__aicore__ inline void CausalConv1dV310<T>::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,
const CausalConv1dTilingData *tilingData)
{
tilingData_ = tilingData;
weightCacheValid_ = false;
cachedC0_ = -1;
cachedDimTileSize_ = -1;
xGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(x));
weightGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(weight));
if (tilingData_->hasBias != 0) {
biasGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(bias));
}
convStatesGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(convStates));
if (tilingData_->hasQueryStartLoc != 0) {
if (tilingData_->queryStartLocUseInt64 != 0) {
queryStartLocGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(queryStartLoc));
} else {
queryStartLocGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(queryStartLoc));
}
}
if (tilingData_->hasCacheIndices != 0) {
if (tilingData_->cacheIndicesUseInt64 != 0) {
cacheIndicesGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(cacheIndices));
} else {
cacheIndicesGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(cacheIndices));
}
}
if (tilingData_->hasInitialStateMode != 0) {
if (tilingData_->initialStateModeDtype == 2) {
initialStateModeGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(initialStateMode));
} else if (tilingData_->initialStateModeDtype == 1) {
initialStateModeGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(initialStateMode));
} else {
initialStateModeGmBool.SetGlobalBuffer(reinterpret_cast<__gm__ bool *>(initialStateMode));
}
}
if (tilingData_->hasNumAcceptedTokens != 0) {
if (tilingData_->numAcceptedTokensUseInt64 != 0) {
numAcceptedTokensGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(numAcceptedTokens));
} else {
numAcceptedTokensGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(numAcceptedTokens));
}
}
yGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(y));
pipe.InitBuffer(inBuf, RING_SLOTS * MAX_BLOCK_DIM * sizeof(T));
pipe.InitBuffer(outBuf, 2 * MAX_BLOCK_DIM * sizeof(T));
pipe.InitBuffer(calcBuf, (MAX_WIDTH + 3) * MAX_BLOCK_DIM * sizeof(float));
AllocEvents();
}
template <typename T>
__aicore__ inline void CausalConv1dV310<T>::AllocEvents()
{
weightBiasMte2ToVEvent_ = GetTPipePtr()->AllocEventID<HardEvent::MTE2_V>();
stateMte2ToVEvent_ = GetTPipePtr()->AllocEventID<HardEvent::MTE2_V>();
for (int32_t i = 0; i < RING_SLOTS; ++i) {
inputMte2ToVEvent_[i] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_V>();
}
inputVToMte2Event_ = GetTPipePtr()->AllocEventID<HardEvent::V_MTE2>();
outMte3ToVEvent_[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>();
outMte3ToVEvent_[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>();
outVToMte3Event_[0] = GetTPipePtr()->AllocEventID<HardEvent::V_MTE3>();
outVToMte3Event_[1] = GetTPipePtr()->AllocEventID<HardEvent::V_MTE3>();
stateWritebackMte3ToVEvent_ = GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>();
stateWritebackMte3ToMte2Event_ = GetTPipePtr()->AllocEventID<HardEvent::MTE3_MTE2>();
stateShiftMte2ToMte3Event_ = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE3>();
stateShiftMte3ToMte2Event_ = GetTPipePtr()->AllocEventID<HardEvent::MTE3_MTE2>();
stateShiftVToMte3Event_ = GetTPipePtr()->AllocEventID<HardEvent::V_MTE3>();
specWritebackMte2ToMte3Event_[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE3>();
specWritebackMte2ToMte3Event_[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE3>();
specWritebackMte3ToMte2Event_[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE3_MTE2>();
specWritebackMte3ToMte2Event_[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE3_MTE2>();
}
template <typename T>
__aicore__ inline void CausalConv1dV310<T>::ReleaseEvents()
{
GetTPipePtr()->ReleaseEventID<HardEvent::MTE2_V>(weightBiasMte2ToVEvent_);
GetTPipePtr()->ReleaseEventID<HardEvent::MTE2_V>(stateMte2ToVEvent_);
for (int32_t i = 0; i < RING_SLOTS; ++i) {
GetTPipePtr()->ReleaseEventID<HardEvent::MTE2_V>(inputMte2ToVEvent_[i]);
}
GetTPipePtr()->ReleaseEventID<HardEvent::V_MTE2>(inputVToMte2Event_);
GetTPipePtr()->ReleaseEventID<HardEvent::MTE3_V>(outMte3ToVEvent_[0]);
GetTPipePtr()->ReleaseEventID<HardEvent::MTE3_V>(outMte3ToVEvent_[1]);
GetTPipePtr()->ReleaseEventID<HardEvent::V_MTE3>(outVToMte3Event_[0]);
GetTPipePtr()->ReleaseEventID<HardEvent::V_MTE3>(outVToMte3Event_[1]);
GetTPipePtr()->ReleaseEventID<HardEvent::MTE3_V>(stateWritebackMte3ToVEvent_);
GetTPipePtr()->ReleaseEventID<HardEvent::MTE3_MTE2>(stateWritebackMte3ToMte2Event_);
GetTPipePtr()->ReleaseEventID<HardEvent::MTE2_MTE3>(stateShiftMte2ToMte3Event_);
GetTPipePtr()->ReleaseEventID<HardEvent::MTE3_MTE2>(stateShiftMte3ToMte2Event_);
GetTPipePtr()->ReleaseEventID<HardEvent::V_MTE3>(stateShiftVToMte3Event_);
GetTPipePtr()->ReleaseEventID<HardEvent::MTE2_MTE3>(specWritebackMte2ToMte3Event_[0]);
GetTPipePtr()->ReleaseEventID<HardEvent::MTE2_MTE3>(specWritebackMte2ToMte3Event_[1]);
GetTPipePtr()->ReleaseEventID<HardEvent::MTE3_MTE2>(specWritebackMte3ToMte2Event_[0]);
GetTPipePtr()->ReleaseEventID<HardEvent::MTE3_MTE2>(specWritebackMte3ToMte2Event_[1]);
}
template <typename T>
__aicore__ inline int32_t CausalConv1dV310<T>::ReadQueryStartLocValue(int32_t index) const
{
if (tilingData_->queryStartLocUseInt64 != 0) {
const int64_t value = queryStartLocGmInt64.GetValue(index);
if (value < 0 || value > INT32_MAX_VALUE) {
return -1;
}
return static_cast<int32_t>(value);
}
return queryStartLocGmInt32.GetValue(index);
}
template <typename T>
__aicore__ inline int64_t CausalConv1dV310<T>::ReadCacheIndexValue(int32_t seq) const
{
const int32_t offset = seq * static_cast<int32_t>(tilingData_->cacheIndicesStride);
if (tilingData_->cacheIndicesUseInt64 != 0) {
return cacheIndicesGmInt64.GetValue(offset);
}
return static_cast<int64_t>(cacheIndicesGmInt32.GetValue(offset));
}
template <typename T>
__aicore__ inline bool CausalConv1dV310<T>::ReadInitialStateModeValue(int32_t seq) const
{
if (tilingData_->initialStateModeDtype == 2) {
return initialStateModeGmInt64.GetValue(seq) != 0;
}
if (tilingData_->initialStateModeDtype == 1) {
return initialStateModeGmInt32.GetValue(seq) != 0;
}
return initialStateModeGmBool.GetValue(seq);
}
template <typename T>
__aicore__ inline int32_t CausalConv1dV310<T>::ReadNumAcceptedTokensValue(int32_t seq) const
{
if (tilingData_->numAcceptedTokensUseInt64 != 0) {
const int64_t value = numAcceptedTokensGmInt64.GetValue(seq);
if (value <= 0) {
return 0;
}
if (value > INT32_MAX_VALUE) {
return static_cast<int32_t>(INT32_MAX_VALUE);
}
return static_cast<int32_t>(value);
}
return numAcceptedTokensGmInt32.GetValue(seq);
}
template <typename T>
__aicore__ inline void CausalConv1dV310<T>::LoadWeightAndBias(int32_t c0, int32_t dimTileSize)
{
const int32_t dim = tilingData_->dim;
const int32_t width = static_cast<int32_t>(tilingData_->width);
const int32_t jStart = MAX_WIDTH - width;
LocalTensor<float> calc = calcBuf.Get<float>();
LocalTensor<float> weightF = calc;
LocalTensor<float> biasF = weightF[MAX_WIDTH * MAX_BLOCK_DIM];
const bool hasBias = (tilingData_->hasBias != 0);
for (int32_t j = 0; j < width; ++j) {
const int32_t jDst = jStart + j;
const int64_t weightOffset = static_cast<int64_t>(j) * dim + c0;
if constexpr (std::is_same<T, float>::value) {
DataCopy(weightF[jDst * MAX_BLOCK_DIM], weightGm[weightOffset], dimTileSize);
} else {
DataCopy(weightF.ReinterpretCast<T>()[jDst * MAX_BLOCK_DIM * 2 + MAX_BLOCK_DIM], weightGm[weightOffset],
dimTileSize);
}
}
if (hasBias) {
if constexpr (std::is_same<T, float>::value) {
DataCopy(biasF, biasGm[c0], dimTileSize);
} else {
DataCopy(biasF.ReinterpretCast<T>()[MAX_BLOCK_DIM], biasGm[c0], dimTileSize);
}
}
SetFlag<HardEvent::MTE2_V>(weightBiasMte2ToVEvent_);
WaitFlag<HardEvent::MTE2_V>(weightBiasMte2ToVEvent_);
if constexpr (!std::is_same<T, float>::value) {
for (int32_t j = 0; j < width; ++j) {
const int32_t jDst = jStart + j;
Cast(weightF[jDst * MAX_BLOCK_DIM], weightF.ReinterpretCast<T>()[jDst * MAX_BLOCK_DIM * 2 + MAX_BLOCK_DIM],
RoundMode::CAST_NONE, dimTileSize);
}
if (hasBias) {
Cast(biasF, biasF.ReinterpretCast<T>()[MAX_BLOCK_DIM], RoundMode::CAST_NONE, dimTileSize);
}
}
if (!hasBias) {
Duplicate(biasF, 0.0f, dimTileSize);
}
}
template <typename T>
__aicore__ inline void CausalConv1dV310<T>::InitRing(int32_t cacheIdx, bool hasInit, int32_t stateTokenOffset,
int32_t start, int32_t len, int32_t c0, int32_t dimTileSize,
int32_t dim)
{
const int32_t stateLen = tilingData_->stateLen;
const int32_t width = static_cast<int32_t>(tilingData_->width);
const int32_t ringStart = MAX_WIDTH - width;
LocalTensor<T> ring = inBuf.Get<T>();
for (int32_t i = 0; i < ringStart; ++i) {
Duplicate(ring[i * MAX_BLOCK_DIM], static_cast<T>(0), dimTileSize);
}
if (ringStart > 0) {
PipeBarrier<PIPE_V>();
}
if (hasInit) {
for (int32_t i = 0; i < (width - 1); ++i) {
const int32_t pos = stateTokenOffset + i;
const int64_t stateOffset =
static_cast<int64_t>(cacheIdx) * stateLen * dim + static_cast<int64_t>(pos) * dim + c0;
DataCopy(ring[(ringStart + i) * MAX_BLOCK_DIM], convStatesGm[stateOffset], dimTileSize);
}
SetFlag<HardEvent::MTE2_V>(stateMte2ToVEvent_);
WaitFlag<HardEvent::MTE2_V>(stateMte2ToVEvent_);
} else {
for (int32_t i = 0; i < (width - 1); ++i) {
Duplicate(ring[(ringStart + i) * MAX_BLOCK_DIM], static_cast<T>(0), dimTileSize);
}
PipeBarrier<PIPE_V>();
}
if (len > 0) {
const int32_t slot0 = SlotCurr(0);
const int64_t xOffset = static_cast<int64_t>(start) * dim + c0;
DataCopy(ring[slot0 * MAX_BLOCK_DIM], xGm[xOffset], dimTileSize);
SetFlag<HardEvent::MTE2_V>(inputMte2ToVEvent_[slot0]);
}
if (len > 1) {
SetFlag<HardEvent::V_MTE2>(inputVToMte2Event_);
}
}
template <typename T>
__aicore__ inline void CausalConv1dV310<T>::RunSeq(int32_t start, int32_t len, int32_t c0, int32_t dimTileSize, int32_t dim)
{
const int32_t width = static_cast<int32_t>(tilingData_->width);
const int32_t jStart = MAX_WIDTH - width;
LocalTensor<float> calc = calcBuf.Get<float>();
LocalTensor<float> weightF = calc;
LocalTensor<float> biasF = weightF[MAX_WIDTH * MAX_BLOCK_DIM];
LocalTensor<float> accF = biasF[MAX_BLOCK_DIM];
LocalTensor<float> tmpF = accF[MAX_BLOCK_DIM];
LocalTensor<T> ring = inBuf.Get<T>();
LocalTensor<T> outT = outBuf.Get<T>();
const bool hasActivation = (tilingData_->activationMode != 0);
for (int32_t t = 0; t < len; ++t) {
const int32_t slotCurr = SlotCurr(t);
WaitFlag<HardEvent::MTE2_V>(inputMte2ToVEvent_[slotCurr]);
if (t + 1 < len) {
const int32_t slotNext = SlotPrefetch(t);
const int64_t xOffsetNext = static_cast<int64_t>(start + t + 1) * dim + c0;
WaitFlag<HardEvent::V_MTE2>(inputVToMte2Event_);
DataCopy(ring[slotNext * MAX_BLOCK_DIM], xGm[xOffsetNext], dimTileSize);
SetFlag<HardEvent::MTE2_V>(inputMte2ToVEvent_[slotNext]);
}
DataCopy(accF, biasF, dimTileSize);
PipeBarrier<PIPE_V>();
for (int32_t j = jStart; j < MAX_WIDTH; ++j) {
const int32_t tap = (MAX_WIDTH - 1) - j;
const int32_t slot = (tap == 0) ? slotCurr : SlotHist(t, tap);
Cast(tmpF, ring[slot * MAX_BLOCK_DIM], RoundMode::CAST_NONE, dimTileSize);
PipeBarrier<PIPE_V>();
MulAddDst(accF, tmpF, weightF[j * MAX_BLOCK_DIM], dimTileSize);
}
if (hasActivation) {
Silu(tmpF, accF, dimTileSize);
}
const int32_t outSlot = t & 1;
LocalTensor<T> outSlotT = outT[outSlot * MAX_BLOCK_DIM];
if (t >= 2) {
WaitFlag<HardEvent::MTE3_V>(outMte3ToVEvent_[outSlot]);
}
if constexpr (IsSameType<T, float>::value) {
if (hasActivation) {
DataCopy(outSlotT, tmpF, dimTileSize);
} else {
DataCopy(outSlotT, accF, dimTileSize);
}
} else {
if (hasActivation) {
Cast(outSlotT, tmpF, RoundMode::CAST_NONE, dimTileSize);
} else {
Cast(outSlotT, accF, RoundMode::CAST_RINT, dimTileSize);
}
}
SetFlag<HardEvent::V_MTE3>(outVToMte3Event_[outSlot]);
const int64_t outOffset = static_cast<int64_t>(start + t) * dim + c0;
WaitFlag<HardEvent::V_MTE3>(outVToMte3Event_[outSlot]);
DataCopy(yGm[outOffset], outSlotT, dimTileSize);
if (t + 2 < len) {
SetFlag<HardEvent::MTE3_V>(outMte3ToVEvent_[outSlot]);
}
if (t + 2 < len) {
SetFlag<HardEvent::V_MTE2>(inputVToMte2Event_);
}
}
}
template <typename T>
__aicore__ inline void CausalConv1dV310<T>::WriteBackState(int32_t cacheIdx, int32_t len, int32_t c0, int32_t dimTileSize,
int32_t dim)
{
const int32_t stateLen = tilingData_->stateLen;
const int32_t width = static_cast<int32_t>(tilingData_->width);
if (len <= 0) {
return;
}
const int32_t expectedDimTileSize = (c0 + dimTileSize <= dim) ? dimTileSize : (dim - c0);
if (c0 < 0 || c0 >= dim || dimTileSize <= 0 || dimTileSize > expectedDimTileSize) {
// Invalid c0 or dimTileSize would cause wrong state writeback
return;
}
const int32_t lastT = len - 1;
LocalTensor<T> ring = inBuf.Get<T>();
const int32_t lastSlot = SlotCurr(lastT);
const int64_t stateBaseOffset = static_cast<int64_t>(cacheIdx) * stateLen * dim + c0;
for (int32_t pos = 0; pos < (width - 1); ++pos) {
const int32_t tap = (width - 2) - pos;
const int32_t slot = RetreatRingSlot(lastSlot, tap);
const int64_t stateOffset = stateBaseOffset + static_cast<int64_t>(pos) * dim;
DataCopy(convStatesGm[stateOffset], ring[slot * MAX_BLOCK_DIM], dimTileSize);
}
}
template <typename T>
__aicore__ inline void CausalConv1dV310<T>::WriteBackStateSpec(int32_t cacheIdx, bool hasInit, int32_t stateTokenOffset,
int32_t start, int32_t len, int32_t c0, int32_t dimTileSize,
int32_t dim)
{
const int32_t width = static_cast<int32_t>(tilingData_->width);
const int32_t stateLen = tilingData_->stateLen;
if (len <= 0) {
return;
}
if (width != 4) {
WriteBackState(cacheIdx, len, c0, dimTileSize, dim);
return;
}
constexpr int32_t keep = MAX_WIDTH - 2;
const int32_t reqStateLen = keep + len;
if (reqStateLen > stateLen) {
WriteBackState(cacheIdx, len, c0, dimTileSize, dim);
return;
}
LocalTensor<T> ring = inBuf.Get<T>();
LocalTensor<T> buf0 = ring[0 * MAX_BLOCK_DIM];
LocalTensor<T> buf1 = ring[1 * MAX_BLOCK_DIM];
if (hasInit) {
const int32_t srcPos0 = stateTokenOffset + 1;
const int32_t srcPos1 = stateTokenOffset + 2;
const int64_t srcOffset0 =
static_cast<int64_t>(cacheIdx) * stateLen * dim + static_cast<int64_t>(srcPos0) * dim + c0;
const int64_t srcOffset1 =
static_cast<int64_t>(cacheIdx) * stateLen * dim + static_cast<int64_t>(srcPos1) * dim + c0;
DataCopy(buf0, convStatesGm[srcOffset0], dimTileSize);
DataCopy(buf1, convStatesGm[srcOffset1], dimTileSize);
SetFlag<HardEvent::MTE2_MTE3>(stateShiftMte2ToMte3Event_);
WaitFlag<HardEvent::MTE2_MTE3>(stateShiftMte2ToMte3Event_);
const int64_t dstOffset0 = static_cast<int64_t>(cacheIdx) * stateLen * dim + static_cast<int64_t>(0) * dim + c0;
const int64_t dstOffset1 = static_cast<int64_t>(cacheIdx) * stateLen * dim + static_cast<int64_t>(1) * dim + c0;
DataCopy(convStatesGm[dstOffset0], buf0, dimTileSize);
DataCopy(convStatesGm[dstOffset1], buf1, dimTileSize);
SetFlag<HardEvent::MTE3_MTE2>(stateShiftMte3ToMte2Event_);
WaitFlag<HardEvent::MTE3_MTE2>(stateShiftMte3ToMte2Event_);
} else {
Duplicate(buf0, static_cast<T>(0), dimTileSize);
SetFlag<HardEvent::V_MTE3>(stateShiftVToMte3Event_);
WaitFlag<HardEvent::V_MTE3>(stateShiftVToMte3Event_);
const int64_t dstOffset0 = static_cast<int64_t>(cacheIdx) * stateLen * dim + static_cast<int64_t>(0) * dim + c0;
const int64_t dstOffset1 = static_cast<int64_t>(cacheIdx) * stateLen * dim + static_cast<int64_t>(1) * dim + c0;
DataCopy(convStatesGm[dstOffset0], buf0, dimTileSize);
DataCopy(convStatesGm[dstOffset1], buf0, dimTileSize);
SetFlag<HardEvent::MTE3_MTE2>(stateShiftMte3ToMte2Event_);
WaitFlag<HardEvent::MTE3_MTE2>(stateShiftMte3ToMte2Event_);
}
const int64_t xOffset0 = static_cast<int64_t>(start) * dim + c0;
DataCopy(buf0, xGm[xOffset0], dimTileSize);
SetFlag<HardEvent::MTE2_MTE3>(specWritebackMte2ToMte3Event_[0]);
for (int32_t t = 0; t < len; ++t) {
const int32_t curr = t & 1;
const int32_t next = curr ^ 1;
LocalTensor<T> currBuf = (curr == 0) ? buf0 : buf1;
LocalTensor<T> nextBuf = (next == 0) ? buf0 : buf1;
WaitFlag<HardEvent::MTE2_MTE3>(specWritebackMte2ToMte3Event_[curr]);
if (t + 1 < len) {
const int64_t xOffsetNext = static_cast<int64_t>(start + t + 1) * dim + c0;
if (t > 0) {
WaitFlag<HardEvent::MTE3_MTE2>(specWritebackMte3ToMte2Event_[next]);
}
DataCopy(nextBuf, xGm[xOffsetNext], dimTileSize);
SetFlag<HardEvent::MTE2_MTE3>(specWritebackMte2ToMte3Event_[next]);
}
const int64_t dstOffset =
static_cast<int64_t>(cacheIdx) * stateLen * dim + static_cast<int64_t>(keep + t) * dim + c0;
DataCopy(convStatesGm[dstOffset], currBuf, dimTileSize);
SetFlag<HardEvent::MTE3_MTE2>(specWritebackMte3ToMte2Event_[curr]);
}
WaitFlag<HardEvent::MTE3_MTE2>(specWritebackMte3ToMte2Event_[0]);
if (len > 1) {
WaitFlag<HardEvent::MTE3_MTE2>(specWritebackMte3ToMte2Event_[1]);
}
}
template <typename T>
__aicore__ inline void CausalConv1dV310<T>::Process()
{
const int32_t dim = tilingData_->dim;
const int32_t batch = tilingData_->batch;
const int32_t inputMode = tilingData_->inputMode;
const int32_t seqLen = tilingData_->seqLen;
const int32_t dimTileSize = static_cast<int32_t>(tilingData_->dimTileSize);
const int32_t blocksPerSeq = static_cast<int32_t>(tilingData_->blocksPerSeq);
const int32_t width = static_cast<int32_t>(tilingData_->width);
const bool isSpecDecodingGlobal = (tilingData_->runMode == 1) && (tilingData_->hasNumAcceptedTokens != 0) &&
(width == 4);
const uint32_t blockIdx = GetBlockIdx();
const uint32_t blockNum = GetBlockNum();
if (dimTileSize <= 0 || blocksPerSeq <= 0 || dimTileSize > MAX_BLOCK_DIM || width < 2 || width > MAX_WIDTH) {
ReleaseEvents();
return;
}
const int64_t gridSize = static_cast<int64_t>(batch) * blocksPerSeq;
for (int64_t task = static_cast<int64_t>(blockIdx); task < gridSize; task += static_cast<int64_t>(blockNum)) {
const int32_t seq = static_cast<int32_t>(task / blocksPerSeq);
const int32_t dimBlockId = static_cast<int32_t>(task % blocksPerSeq);
const int32_t c0 = dimBlockId * dimTileSize;
if (c0 >= dim) {
continue;
}
const int32_t dimTileSizeActual = (c0 + dimTileSize <= dim) ? dimTileSize : (dim - c0);
int32_t start = 0;
int32_t len = 0;
if (inputMode == 0) {
const int32_t startVal = ReadQueryStartLocValue(seq);
const int32_t endVal = ReadQueryStartLocValue(seq + 1);
if (startVal < 0 || endVal < startVal || endVal > tilingData_->cuSeqlen) {
continue;
}
start = startVal;
len = endVal - startVal;
} else if (inputMode == 2) {
start = seq;
len = 1;
} else {
start = seq * seqLen;
len = seqLen;
}
if (len <= 0) {
continue;
}
int32_t cacheIdx = seq;
if (tilingData_->hasCacheIndices != 0) {
const int64_t cacheIdx64 = ReadCacheIndexValue(seq);
if (cacheIdx64 == tilingData_->padSlotId) {
continue;
}
if (cacheIdx64 < 0 || cacheIdx64 >= tilingData_->numCacheLines) {
continue;
}
cacheIdx = static_cast<int32_t>(cacheIdx64);
}
const bool hasInit = (tilingData_->hasInitialStateMode != 0)
? ReadInitialStateModeValue(seq)
: (tilingData_->runMode == 1);
int32_t stateTokenOffset = 0;
if (isSpecDecodingGlobal) {
int32_t accepted = ReadNumAcceptedTokensValue(seq);
stateTokenOffset = accepted - 1;
const int32_t maxOffset = static_cast<int32_t>(tilingData_->stateLen - (width - 1));
if (stateTokenOffset < 0) {
stateTokenOffset = 0;
} else if (stateTokenOffset > maxOffset) {
stateTokenOffset = maxOffset;
}
}
const bool weightCacheHit = weightCacheValid_ && (cachedC0_ == c0) && (cachedDimTileSize_ == dimTileSizeActual);
if (!weightCacheHit) {
LoadWeightAndBias(c0, dimTileSizeActual);
weightCacheValid_ = true;
cachedC0_ = c0;
cachedDimTileSize_ = dimTileSizeActual;
}
InitRing(cacheIdx, hasInit, stateTokenOffset, start, len, c0, dimTileSizeActual, dim);
RunSeq(start, len, c0, dimTileSizeActual, dim);
SetFlag<HardEvent::MTE3_V>(stateWritebackMte3ToVEvent_);
WaitFlag<HardEvent::MTE3_V>(stateWritebackMte3ToVEvent_);
SetFlag<HardEvent::MTE3_MTE2>(stateWritebackMte3ToMte2Event_);
WaitFlag<HardEvent::MTE3_MTE2>(stateWritebackMte3ToMte2Event_);
if (isSpecDecodingGlobal) {
WriteBackStateSpec(cacheIdx, hasInit, stateTokenOffset, start, len, c0, dimTileSizeActual, dim);
} else {
WriteBackState(cacheIdx, len, c0, dimTileSizeActual, dim);
}
SetFlag<HardEvent::MTE3_V>(stateWritebackMte3ToVEvent_);
WaitFlag<HardEvent::MTE3_V>(stateWritebackMte3ToVEvent_);
SetFlag<HardEvent::MTE3_MTE2>(stateWritebackMte3ToMte2Event_);
WaitFlag<HardEvent::MTE3_MTE2>(stateWritebackMte3ToMte2Event_);
PipeBarrier<PIPE_V>();
PipeBarrier<PIPE_MTE2>();
PipeBarrier<PIPE_MTE3>();
}
ReleaseEvents();
}
} // namespace NsCausalConv1d
#endif // CAUSAL_CONV1D_V310_H

View File

@@ -0,0 +1,51 @@
/**
* This program is free software, you can redistribute it and/or modify it.
* Copyright (c) 2026 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_v310_common.h
* \brief Common utilities and constants for CausalConv1D prefill kernel.
*/
#ifndef CAUSAL_CONV1D_V310_COMMON_H
#define CAUSAL_CONV1D_V310_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;
}
__aicore__ inline int32_t RetreatRingSlot(int32_t slot, int32_t delta)
{
int32_t prev = slot - delta;
return (prev >= 0) ? prev : (prev + RING_SLOTS);
}
} // namespace NsCausalConv1dCommon
#endif // CAUSAL_CONV1D_V310_COMMON_H

View File

@@ -0,0 +1,55 @@
/**
* This program is free software, you can redistribute it and/or modify it.
* Copyright (c) 2026 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_v310_tiling_data.h
* \brief tiling data struct
*/
#ifndef CAUSAL_CONV1D_V310_TILING_DATA_H_
#define CAUSAL_CONV1D_V310_TILING_DATA_H_
#include <cstdint>
struct CausalConv1dTilingData {
int64_t dim;
int64_t cuSeqlen;
int64_t seqLen;
int64_t inputMode;
int64_t runMode;
int64_t width;
int64_t stateLen;
int64_t numCacheLines;
int64_t batch;
int64_t activationMode;
int64_t padSlotId;
int64_t hasBias;
int64_t dimTileSize;
int64_t blocksPerSeq;
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;
};
#endif // CAUSAL_CONV1D_V310_TILING_DATA_H_

View File

@@ -0,0 +1,34 @@
/**
* This program is free software, you can redistribute it and/or modify it.
* Copyright (c) 2026 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_v310_tiling_key.h
* \brief causal_conv1d_v310 tiling key declare
*/
#ifndef __CAUSAL_CONV1D_V310_TILING_KEY_H__
#define __CAUSAL_CONV1D_V310_TILING_KEY_H__
#include "ascendc/host_api/tiling/template_argument.h"
#define CAUSAL_CONV1D_TPL_SCH_MODE_DEFAULT 0
ASCENDC_TPL_ARGS_DECL(CausalConv1dV310,
ASCENDC_TPL_UINT_DECL(
schMode, 1, ASCENDC_TPL_UI_LIST, CAUSAL_CONV1D_TPL_SCH_MODE_DEFAULT)
);
ASCENDC_TPL_SEL(
ASCENDC_TPL_ARGS_SEL(
ASCENDC_TPL_UINT_SEL(
schMode, ASCENDC_TPL_UI_LIST, CAUSAL_CONV1D_TPL_SCH_MODE_DEFAULT)));
#endif // __CAUSAL_CONV1D_V310_TILING_KEY_H__