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,70 @@
/**
* Copyright (c) 2025-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 scatter_nd_update_common.h
* \brief ScatterNdUpdateV2 公共定义和工具函数
*/
#ifndef SCATTER_ND_UPDATE_V2_COMMON_H
#define SCATTER_ND_UPDATE_V2_COMMON_H
#include "kernel_operator.h"
namespace ScatterNdUpdateV2 {
using namespace AscendC;
// 公共常量定义
constexpr uint64_t DOUBLE_BUFFER = 1;
constexpr uint64_t SORT_RES_NUM = 2;
constexpr uint64_t SORT_TMP_NUM = 3;
constexpr uint64_t ALIGNED_BLOCK_NUM = 32;
constexpr uint64_t ALIGN_NUM = 8; // 32 字节对齐 = 8 个 int32
constexpr uint64_t ALIGNED_SIZE = 512;
// 公共同步函数
__aicore__ inline void PipeMte2ToS()
{
event_t eventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE2_S));
SetFlag<HardEvent::MTE2_S>(eventID);
WaitFlag<HardEvent::MTE2_S>(eventID);
}
__aicore__ inline void PipeMte3ToS()
{
event_t eventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_S));
SetFlag<HardEvent::MTE3_S>(eventID);
WaitFlag<HardEvent::MTE3_S>(eventID);
}
__aicore__ inline void PipeVToMte3()
{
event_t eventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3));
SetFlag<HardEvent::V_MTE3>(eventID);
WaitFlag<HardEvent::V_MTE3>(eventID);
}
// 计算 block 分布参数
__aicore__ inline void CalcBlockDistribution(
uint64_t blockIdx, uint64_t frontNum, uint64_t frontRow, uint64_t tailRow,
uint64_t& computeRow, uint64_t& start)
{
if (blockIdx >= frontNum) {
computeRow = tailRow;
start = frontNum * frontRow + (blockIdx - frontNum) * computeRow;
} else {
computeRow = frontRow;
start = blockIdx * computeRow;
}
}
} // namespace ScatterNdUpdateV2
#endif // SCATTER_ND_UPDATE_V2_COMMON_H

View File

@@ -0,0 +1,175 @@
/**
* Copyright (c) 2025-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 scatter_nd_update_large_index.h
* \brief LargeIndex Kernel (index > 2^31-1)
*/
#include "kernel_operator.h"
#include "kernel_tiling/kernel_tiling.h"
#include "scatter_nd_update_common.h"
namespace ScatterNdUpdateV2 {
template<typename T>
class LargeIndexKernel {
public:
__aicore__ inline LargeIndexKernel() = delete;
__aicore__ inline LargeIndexKernel(
GM_ADDR indices, GM_ADDR updates, GM_ADDR output,
const ScatterNdUpdateV2TilingData& tiling, TPipe& pipe)
{
InitParams(tiling);
InitBuffers(pipe);
SetGmAddr(indices, updates, output, tiling);
}
__aicore__ inline void InitParams(const ScatterNdUpdateV2TilingData& tiling)
{
blockIdx_ = GetBlockIdx();
CalcBlockDistribution(blockIdx_, tiling.scatterTiling.frontNum, tiling.scatterTiling.frontRow,
tiling.scatterTiling.tailRow, computeRow_, start_);
end_ = start_ + computeRow_;
startInt64_ = static_cast<int64_t>(start_);
endInt64_ = static_cast<int64_t>(end_);
indexDim_ = tiling.linearIndexTiling.indexDim;
blockLength_ = tiling.linearIndexTiling.blockLength;
blockNum_ = tiling.linearIndexTiling.blockNum;
blockRemainLength_ = tiling.linearIndexTiling.blockRemainLength;
scatterLength_ = tiling.scatterTiling.scatterLength;
ubLengthForUpdates_ = tiling.scatterTiling.ubLengthForUpdates;
scatterTileNum_ = tiling.scatterTiling.scatterTileNum;
scatterTileLength_ = tiling.scatterTiling.scatterTileLength;
scatterTileTail_ = tiling.scatterTiling.scatterTileTail;
for (uint64_t i = 0; i < indexDim_; ++i) {
indicesMask_[i] = tiling.linearIndexTiling.indicesMask[i];
}
}
__aicore__ inline void InitBuffers(TPipe& pipe)
{
uint64_t indicesInt64Size = ((blockLength_ * indexDim_ * 2) + ALIGN_NUM - 1) & ~(ALIGN_NUM - 1);
uint64_t updateBufBytes = (ubLengthForUpdates_ * sizeof(T) + 31) & ~31ULL;
pipe.InitBuffer(indicesBuf, indicesInt64Size * sizeof(int));
pipe.InitBuffer(updateBuf, updateBufBytes);
indicesInt64Local = indicesBuf.Get<int>().ReinterpretCast<int64_t>();
updateLocal = updateBuf.Get<T>();
}
__aicore__ inline void SetGmAddr(GM_ADDR indices, GM_ADDR updates, GM_ADDR output,
const ScatterNdUpdateV2TilingData& tiling)
{
indicesGmInt64_.SetGlobalBuffer((__gm__ int64_t*)indices);
updatesGm_.SetGlobalBuffer((__gm__ T*)updates);
outputGm_.SetGlobalBuffer((__gm__ T*)output);
}
__aicore__ inline void Process()
{
for (uint64_t blockIdx = 0; blockIdx < blockNum_; ++blockIdx) {
ProcessOneBlock(blockIdx, false);
}
if (blockRemainLength_ != 0) {
ProcessOneBlock(blockNum_, true);
}
}
__aicore__ inline void ProcessOneBlock(uint64_t blockIdx, bool isTail)
{
uint64_t copyRow = isTail ? blockRemainLength_ : blockLength_;
CopyInInt64(blockIdx, isTail);
for (uint64_t i = 0; i < copyRow; ++i) {
int64_t linearIndex = ComputeLinearIndex(i);
if (linearIndex >= startInt64_ && linearIndex < endInt64_) {
ScatterUpdate(i, linearIndex);
}
}
}
__aicore__ inline void CopyInInt64(uint64_t blockIdx, bool isTail)
{
uint64_t indicesOffset = blockIdx * blockLength_ * indexDim_;
uint64_t copyRow = isTail ? blockRemainLength_ : blockLength_;
DataCopyExtParams copyParams{1, static_cast<uint32_t>(copyRow * indexDim_ * sizeof(int64_t)), 0, 0, 0};
DataCopyPadExtParams<int64_t> padParams{true, 0, 0, 0};
DataCopyPad(indicesInt64Local, indicesGmInt64_[indicesOffset], copyParams, padParams);
PipeMte2ToS();
}
__aicore__ inline int64_t ComputeLinearIndex(uint64_t rowIdx)
{
int64_t linearIndex = 0;
for (uint64_t dim = 0; dim < indexDim_; ++dim) {
int64_t idxValue = indicesInt64Local.GetValue(rowIdx * indexDim_ + dim);
int64_t stride = static_cast<int64_t>(indicesMask_[dim]);
linearIndex += idxValue * stride;
}
return linearIndex;
}
__aicore__ inline void ScatterUpdate(uint64_t rowIdx, int64_t linearIndex)
{
for (uint64_t tileIdx = 0; tileIdx < scatterTileNum_; ++tileIdx) {
uint64_t tileLength = (tileIdx == scatterTileNum_ - 1) ? scatterTileTail_ : scatterTileLength_;
uint64_t gmOffset = rowIdx * scatterLength_ + tileIdx * scatterTileLength_;
DataCopyExtParams updateCopyParams{1, static_cast<uint32_t>(tileLength * sizeof(T)), 0, 0, 0};
DataCopyPadExtParams<T> padParams{true, 0, 0, 0};
DataCopyPad(updateLocal, updatesGm_[gmOffset], updateCopyParams, padParams);
PipeMte2ToS();
uint64_t outOffset = static_cast<uint64_t>(linearIndex) + tileIdx * scatterTileLength_;
DataCopyExtParams outParams{1, static_cast<uint32_t>(tileLength * sizeof(T)), 0, 0, 0};
DataCopyPad(outputGm_[outOffset], updateLocal, outParams);
PipeMte3ToS();
}
}
private:
GlobalTensor<int64_t> indicesGmInt64_;
GlobalTensor<T> updatesGm_;
GlobalTensor<T> outputGm_;
TBuf<TPosition::VECCALC> indicesBuf;
TBuf<TPosition::VECCALC> updateBuf;
LocalTensor<int64_t> indicesInt64Local;
LocalTensor<T> updateLocal;
uint64_t blockIdx_;
uint64_t computeRow_;
uint64_t start_;
uint64_t end_;
int64_t startInt64_;
int64_t endInt64_;
uint64_t indexDim_;
uint64_t blockLength_;
uint64_t blockNum_;
uint64_t blockRemainLength_;
uint64_t indicesMask_[8];
uint64_t scatterLength_;
uint64_t ubLengthForUpdates_;
uint64_t scatterTileNum_;
uint64_t scatterTileLength_;
uint64_t scatterTileTail_;
};
} // namespace ScatterNdUpdateV2

View File

@@ -0,0 +1,300 @@
/**
* Copyright (c) 2025-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 scatter_nd_update_linear_index.h
* \brief LinearIndex Kernel
*/
#include "kernel_operator.h"
#include "kernel_tiling/kernel_tiling.h"
#include "scatter_nd_update_common.h"
namespace ScatterNdUpdateV2 {
template<bool isSort, typename IndicesT = int>
class LinearIndexKernel {
public:
__aicore__ inline LinearIndexKernel() = delete;
__aicore__ inline LinearIndexKernel(
GM_ADDR indices, GM_ADDR workSpace, const ScatterNdUpdateV2TilingData& tiling, TPipe& pipe)
{
InitParams(tiling);
InitBuffers(pipe);
SetGmAddr(indices, workSpace, tiling);
}
__aicore__ inline void InitParams(const ScatterNdUpdateV2TilingData& tiling)
{
blockIdx_ = GetBlockIdx();
frontCoreNum_ = tiling.linearIndexTiling.frontCoreNum;
tailCoreNum_ = tiling.linearIndexTiling.tailCoreNum;
frontBlockNum_ = tiling.linearIndexTiling.frontBlockNum;
tailBlockNum_ = tiling.linearIndexTiling.tailBlockNum;
if (blockIdx_ >= frontCoreNum_) {
computeNum_ = tailBlockNum_;
} else {
computeNum_ = frontBlockNum_;
}
ubSize_ = tiling.linearIndexTiling.ubSize;
coreNum_ = tiling.linearIndexTiling.coreNum;
blockNum_ = tiling.linearIndexTiling.blockNum;
blockLength_ = tiling.linearIndexTiling.blockLength;
blockRemainLength_ = tiling.linearIndexTiling.blockRemainLength;
indexDim_ = tiling.linearIndexTiling.indexDim;
indicesMask_ = tiling.linearIndexTiling.indicesMask;
}
template<bool isInt64, bool needSort>
__aicore__ inline void InitBuffersUnified()
{
uint64_t offset = 0;
indicesLocal = allUbLocal[offset];
offset += blockLength_;
uint64_t indicesOffset = offset;
if constexpr (isInt64) {
indicesInt64Local = allUbLocal[offset].ReinterpretCast<int64_t>();
indicesOriginLocal = allUbLocal[offset];
offset += blockLength_ * indexDim_ * 2;
} else {
indicesOriginLocal = allUbLocal[offset];
offset += blockLength_ * indexDim_;
}
addTmpLocal = allUbLocal[offset];
offset += blockLength_;
rangeLocal = allUbLocal[offset];
offset += blockLength_;
if constexpr (isSort) {
resLocal = allUbLocal[indicesOffset].ReinterpretCast<float>();
indicesOffset += blockLength_ * 2;
posIdxLocal = allUbLocal[indicesOffset];
indicesOffset += blockLength_;
sortTmpLocal = allUbLocal[indicesOffset].ReinterpretCast<float>();
}
}
__aicore__ inline void InitBuffers(TPipe& pipe)
{
pipe.InitBuffer(allUbBuf, ubSize_);
allUbLocal = allUbBuf.Get<int>();
if constexpr (isSort) {
if constexpr (std::is_same_v<IndicesT, int64_t>) {
InitBuffersUnified<true, true>();
} else {
InitBuffersUnified<false, true>();
}
} else {
if constexpr (std::is_same_v<IndicesT, int64_t>) {
InitBuffersUnified<true, false>();
} else {
InitBuffersUnified<false, false>();
}
}
}
__aicore__ inline void SetGmAddr(GM_ADDR indices, GM_ADDR workSpace, const ScatterNdUpdateV2TilingData& tiling)
{
indiceAddrOffset_ =
blockIdx_ < tiling.linearIndexTiling.frontCoreNum ?
tiling.linearIndexTiling.frontBlockNum * blockLength_ * blockIdx_ :
tiling.linearIndexTiling.frontCoreNum * tiling.linearIndexTiling.frontBlockNum * blockLength_ +
(blockIdx_ - tiling.linearIndexTiling.frontCoreNum) * tiling.linearIndexTiling.tailBlockNum *
blockLength_;
sortedIndicesGm_.SetGlobalBuffer((__gm__ int*)workSpace + indiceAddrOffset_);
if constexpr (isSort) {
posIndicesGm_.SetGlobalBuffer((__gm__ int*)workSpace + tiling.linearIndexTiling.sortWorkspace + indiceAddrOffset_);
}
if constexpr (std::is_same_v<IndicesT, int64_t>) {
indicesGmInt64_.SetGlobalBuffer((__gm__ int64_t*)indices + indiceAddrOffset_ * indexDim_);
} else {
indicesGm_.SetGlobalBuffer((__gm__ int*)indices + indiceAddrOffset_ * indexDim_);
}
}
__aicore__ inline void Process()
{
if constexpr (isSort) {
for (uint64_t i = 0; i < computeNum_; i++) {
ProcessOneWithSort(i, false);
}
uint64_t lastActiveCore = (blockNum_ == 0) ? 0 :
(tailCoreNum_ == 0 ? frontCoreNum_ - 1 : frontCoreNum_ + tailCoreNum_ - 1);
if (blockIdx_ == lastActiveCore && blockRemainLength_ != 0) {
ProcessOneWithSort(computeNum_, true);
}
} else {
for (uint64_t i = 0; i < computeNum_; i++) {
ProcessOne(i, false);
}
uint64_t lastActiveCore = (blockNum_ == 0) ? 0 :
(tailCoreNum_ == 0 ? frontCoreNum_ - 1 : frontCoreNum_ + tailCoreNum_ - 1);
if (blockIdx_ == lastActiveCore && blockRemainLength_ != 0) {
ProcessOne(computeNum_, true);
}
}
}
__aicore__ inline void ProcessOne(uint64_t idx, bool isTail)
{
CopyIn(idx, isTail);
if constexpr (std::is_same_v<IndicesT, int64_t>) {
CastToInt32(idx, isTail);
}
Compute4LinearIndex(idx, isTail);
CopyOut(idx, isTail);
}
__aicore__ inline void ProcessOneWithSort(uint64_t idx, bool isTail)
{
CopyIn(idx, isTail);
if constexpr (std::is_same_v<IndicesT, int64_t>) {
CastToInt32(idx, isTail);
}
Compute4LinearIndex(idx, isTail);
ComputeForSort(idx, isTail);
CopyOut(idx, isTail);
}
__aicore__ inline void CopyIn(uint64_t process, bool isTail)
{
uint64_t indicesOffset = process * blockLength_ * indexDim_;
uint64_t copyRow = isTail ? blockRemainLength_ : blockLength_;
if constexpr (std::is_same_v<IndicesT, int64_t>) {
DataCopyExtParams copyParams{1, static_cast<uint32_t>(copyRow * indexDim_ * sizeof(int64_t)), 0, 0, 0};
DataCopyPadExtParams<int64_t> padParams{true, 0, 0, 0};
DataCopyPad(indicesInt64Local, indicesGmInt64_[indicesOffset], copyParams, padParams);
} else {
DataCopyExtParams copyParams{1, static_cast<uint32_t>(copyRow * indexDim_ * sizeof(int)), 0, 0, 0};
DataCopyPadExtParams<int> padParams{true, 0, 0, 0};
DataCopyPad(indicesOriginLocal, indicesGm_[indicesOffset], copyParams, padParams);
}
PipeMte2ToS();
}
__aicore__ inline void CastToInt32(uint64_t process, bool isTail)
{
uint64_t computeRow = isTail ? blockRemainLength_ : blockLength_;
uint64_t totalElements = computeRow * indexDim_;
Cast(indicesOriginLocal, indicesInt64Local, RoundMode::CAST_NONE, totalElements);
PipeBarrier<PIPE_V>();
}
__aicore__ inline void Compute4LinearIndex(uint64_t process, bool isTail)
{
uint64_t computeRow = isTail ? blockRemainLength_ : blockLength_;
int32_t malValue = indexDim_ * sizeof(int);
Duplicate<int>(indicesLocal, 0, computeRow);
CreateVecIndex(rangeLocal, (int)0, computeRow);
PipeBarrier<PIPE_V>();
Muls(rangeLocal, rangeLocal, malValue, computeRow);
PipeBarrier<PIPE_V>();
for (int i = 0; i < indexDim_; ++i) {
if (i != 0) {
Adds(rangeLocal, rangeLocal, (int)(sizeof(int)), computeRow);
PipeBarrier<PIPE_V>();
}
LocalTensor<uint32_t> rangeLocalCasted = rangeLocal.ReinterpretCast<uint32_t>();
Gather(addTmpLocal, indicesOriginLocal, rangeLocalCasted, (uint32_t)0, (uint32_t)computeRow);
PipeBarrier<PIPE_V>();
Muls(addTmpLocal, addTmpLocal, (int)indicesMask_[i], computeRow);
PipeBarrier<PIPE_V>();
Add(indicesLocal, indicesLocal, addTmpLocal, computeRow);
PipeBarrier<PIPE_V>();
}
if constexpr (!isSort) {
PipeVToMte3();
}
}
__aicore__ inline void ComputeForSort(uint64_t process, bool isTail)
{
LocalTensor<float> indicesLocalFp32 = indicesLocal.ReinterpretCast<float>();
uint64_t computeRow = isTail ? blockRemainLength_ : blockLength_;
uint64_t computeRowAligned = (computeRow + ALIGNED_BLOCK_NUM - 1) & ~(ALIGNED_BLOCK_NUM - 1);
uint64_t repeatTimes = computeRowAligned / ALIGNED_BLOCK_NUM;
uint64_t repeatId = computeRow / ALIGNED_BLOCK_NUM;
uint64_t repeatRemain = computeRow % ALIGNED_BLOCK_NUM;
int addValue = indiceAddrOffset_ + process * blockLength_;
Cast(indicesLocalFp32, indicesLocal, RoundMode::CAST_ROUND, computeRowAligned);
if (repeatRemain != 0) {
// 对齐处理:不足32的部分设为-1
Duplicate<int>(rangeLocal, -1, (uint32_t)ALIGNED_BLOCK_NUM);
PipeBarrier<PIPE_V>();
Cast(rangeLocal, indicesLocalFp32[ALIGNED_BLOCK_NUM * repeatId], RoundMode::CAST_ROUND, (uint32_t)repeatRemain);
PipeBarrier<PIPE_V>();
Cast(indicesLocalFp32[ALIGNED_BLOCK_NUM * repeatId], rangeLocal, RoundMode::CAST_ROUND, (uint32_t)ALIGNED_BLOCK_NUM);
PipeBarrier<PIPE_V>();
}
Duplicate<int>(posIdxLocal, -1, computeRowAligned);
PipeBarrier<PIPE_V>();
CreateVecIndex<int>(posIdxLocal, 0U, computeRow);
LocalTensor<uint32_t> posIdxULocal = posIdxLocal.ReinterpretCast<uint32_t>();
PipeBarrier<PIPE_V>();
Sort<float, true>(resLocal, indicesLocalFp32, posIdxULocal, sortTmpLocal, repeatTimes);
PipeBarrier<PIPE_V>();
Extract(indicesLocalFp32, posIdxULocal, resLocal, repeatTimes);
PipeBarrier<PIPE_V>();
Cast(indicesLocal, indicesLocalFp32, RoundMode::CAST_ROUND, computeRowAligned);
PipeBarrier<PIPE_V>();
Adds(posIdxLocal, posIdxLocal, addValue, computeRow);
PipeBarrier<PIPE_V>();
PipeVToMte3();
}
__aicore__ inline void CopyOut(uint64_t process, bool isTail)
{
uint64_t outOffset = process * blockLength_;
uint64_t copyRow = isTail ? blockRemainLength_ : blockLength_;
DataCopyExtParams copyParams{1, static_cast<uint32_t>(copyRow * sizeof(int)), 0, 0, 0};
DataCopyPad(sortedIndicesGm_[outOffset], indicesLocal, copyParams);
if constexpr (isSort) {
DataCopyPad(posIndicesGm_[outOffset], posIdxLocal, copyParams);
}
PipeMte3ToS();
}
private:
GlobalTensor<int> indicesGm_;
GlobalTensor<int64_t> indicesGmInt64_;
GlobalTensor<int> sortedIndicesGm_;
GlobalTensor<int> posIndicesGm_;
TBuf<TPosition::VECCALC> allUbBuf;
LocalTensor<int> allUbLocal;
LocalTensor<int> indicesLocal;
LocalTensor<int> indicesOriginLocal;
LocalTensor<int64_t> indicesInt64Local;
LocalTensor<int> addTmpLocal;
LocalTensor<int> rangeLocal;
LocalTensor<int> posIdxLocal;
LocalTensor<float> sortTmpLocal;
LocalTensor<float> resLocal;
uint64_t ubSize_;
uint64_t coreNum_;
uint64_t blockIdx_;
uint64_t indexDim_;
uint64_t computeNum_;
uint64_t blockNum_;
uint64_t blockLength_;
uint64_t blockRemainLength_;
uint64_t frontCoreNum_;
uint64_t tailCoreNum_;
uint64_t frontBlockNum_;
uint64_t tailBlockNum_;
const uint64_t* indicesMask_;
uint64_t indiceAddrOffset_;
};
} // namespace ScatterNdUpdateV2

View File

@@ -0,0 +1,152 @@
/**
* Copyright (c) 2025-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 scatter_nd_update_no_sort.h
* \brief Scatter Kernel (NoSort)
*/
#include "kernel_operator.h"
#include "kernel_tiling/kernel_tiling.h"
#include "scatter_nd_update_common.h"
namespace ScatterNdUpdateV2 {
constexpr uint64_t ALIGNED_SIZE_INDEX = 8;
template<typename T>
class ScatterNdUpdateV2KernelNoSort {
public:
__aicore__ inline ScatterNdUpdateV2KernelNoSort() = delete;
__aicore__ inline ScatterNdUpdateV2KernelNoSort(
GM_ADDR updates, GM_ADDR output, GM_ADDR workSpace, const ScatterNdUpdateV2TilingData& tiling, TPipe& pipe)
{
InitParam(tiling);
InitBuffers(pipe);
SetGmAddr(updates, output, workSpace, tiling);
}
__aicore__ inline void InitParam(const ScatterNdUpdateV2TilingData& tiling)
{
blockIdx_ = GetBlockIdx();
CalcBlockDistribution(blockIdx_, tiling.scatterTiling.frontNum, tiling.scatterTiling.frontRow,
tiling.scatterTiling.tailRow, computeRow_, start_);
end_ = start_ + computeRow_;
totalIndexRow_ = tiling.linearIndexTiling.blockNum * tiling.linearIndexTiling.blockLength
+ tiling.linearIndexTiling.blockRemainLength;
scatterLength_ = tiling.scatterTiling.scatterLength;
scatterAlignLength_ = tiling.scatterTiling.scatterAlignLength;
ubLengthForUpdates_ = tiling.scatterTiling.ubLengthForUpdates;
scatterTileNum_ = tiling.scatterTiling.scatterTileNum;
scatterTileLength_ = tiling.scatterTiling.scatterTileLength;
scatterTileTail_ = tiling.scatterTiling.scatterTileTail;
scatterTileAlignLength_ = tiling.scatterTiling.scatterTileAlignLength;
CalcIndexTileParams();
}
__aicore__ inline void CalcIndexTileParams()
{
uint64_t ubSizeBytes = ubLengthForUpdates_ * sizeof(T);
uint64_t updateSizeBytes = scatterTileLength_ * sizeof(T);
uint64_t remainBytes = ubSizeBytes - updateSizeBytes;
indexTileLength_ = (remainBytes / sizeof(int) / ALIGNED_SIZE_INDEX) * ALIGNED_SIZE_INDEX;
if (indexTileLength_ == 0) {
indexTileLength_ = ALIGNED_SIZE_INDEX;
}
indexTileNum_ = (totalIndexRow_ + indexTileLength_ - 1) / indexTileLength_;
indexTileTail_ = totalIndexRow_ - (indexTileNum_ - 1) * indexTileLength_;
}
__aicore__ inline void InitBuffers(TPipe& pipe)
{
pipe.InitBuffer(indexQue_, DOUBLE_BUFFER, indexTileLength_ * sizeof(int));
pipe.InitBuffer(updateQue_, DOUBLE_BUFFER, scatterTileLength_ * sizeof(T));
}
__aicore__ inline void SetGmAddr(GM_ADDR updates, GM_ADDR output, GM_ADDR workSpace, const ScatterNdUpdateV2TilingData& tiling)
{
linearIndicesGm_.SetGlobalBuffer((__gm__ int*)workSpace);
updatesGm_.SetGlobalBuffer((__gm__ T*)updates);
outputGm_.SetGlobalBuffer((__gm__ T*)output);
}
__aicore__ inline void Process()
{
for (uint64_t tileIdx = 0; tileIdx < indexTileNum_; ++tileIdx) {
uint64_t curTileLen = (tileIdx == indexTileNum_ - 1) ? indexTileTail_ : indexTileLength_;
uint64_t gmOffset = tileIdx * indexTileLength_;
LocalTensor<int> indexLocal = indexQue_.AllocTensor<int>();
DataCopyExtParams indexCopyParams{1, static_cast<uint32_t>(curTileLen * sizeof(int)), 0, 0, 0};
DataCopyPadExtParams<int> padParams{true, 0, 0, 0};
DataCopyPad(indexLocal, linearIndicesGm_[gmOffset], indexCopyParams, padParams);
indexQue_.EnQue(indexLocal);
PipeMte2ToS();
LocalTensor<int> indexData = indexQue_.DeQue<int>();
for (uint64_t i = 0; i < curTileLen; ++i) {
int64_t linearIndex = static_cast<int64_t>(indexData.GetValue(i));
if (linearIndex >= (int64_t)start_ && linearIndex < (int64_t)end_) {
uint64_t idx = gmOffset + i;
ProcessOneIndex(idx, linearIndex);
}
}
indexQue_.FreeTensor<int>(indexLocal);
}
}
__aicore__ inline void ProcessOneIndex(uint64_t idx, int64_t linearIndex)
{
for (uint64_t tileIdx = 0; tileIdx < scatterTileNum_; ++tileIdx) {
uint64_t tileLength = (tileIdx == scatterTileNum_ - 1) ? scatterTileTail_ : scatterTileLength_;
LocalTensor<T> updateLocal = updateQue_.AllocTensor<T>();
uint64_t gmOffset = idx * scatterLength_ + tileIdx * scatterTileLength_;
DataCopyExtParams updateCopyParams{1, static_cast<uint32_t>(tileLength * sizeof(T)), 0, 0, 0};
DataCopyPadExtParams<T> padParams{true, 0, 0, 0};
DataCopyPad(updateLocal, updatesGm_[gmOffset], updateCopyParams, padParams);
PipeMte2ToS();
uint64_t outOffset = linearIndex + tileIdx * scatterTileLength_;
DataCopyExtParams outParams{1, static_cast<uint32_t>(tileLength * sizeof(T)), 0, 0, 0};
DataCopyPad(outputGm_[outOffset], updateLocal, outParams);
PipeMte3ToS();
updateQue_.FreeTensor<T>(updateLocal);
}
}
private:
GlobalTensor<int> linearIndicesGm_;
GlobalTensor<T> updatesGm_;
GlobalTensor<T> outputGm_;
TQue<TPosition::VECOUT, DOUBLE_BUFFER> indexQue_;
TQue<TPosition::VECOUT, DOUBLE_BUFFER> updateQue_;
uint64_t blockIdx_;
uint64_t computeRow_;
uint64_t start_;
uint64_t end_;
uint64_t totalIndexRow_;
uint64_t scatterLength_;
uint64_t scatterAlignLength_;
uint64_t ubLengthForUpdates_;
uint64_t scatterTileNum_;
uint64_t scatterTileLength_;
uint64_t scatterTileTail_;
uint64_t scatterTileAlignLength_;
uint64_t indexTileLength_;
uint64_t indexTileNum_;
uint64_t indexTileTail_;
};
} // namespace ScatterNdUpdateV2

View File

@@ -0,0 +1,71 @@
/**
* Copyright (c) 2025-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 scatter_nd_update_v2.cpp
* \brief ScatterNdUpdateV2 算子入口
*/
#include "scatter_nd_update_v2.h"
#include "scatter_nd_update_linear_index.h"
#include "scatter_nd_update_no_sort.h"
#include "scatter_nd_update_large_index.h"
extern "C" __global__ __aicore__ void scatter_nd_update_v2(GM_ADDR varRef, GM_ADDR indices,
GM_ADDR updates, GM_ADDR output, GM_ADDR workSpace, GM_ADDR tiling) {
if (workSpace == nullptr) {
return;
}
GM_ADDR user = AscendC::GetUserWorkspace(workSpace);
if (user == nullptr) {
return;
}
GET_TILING_DATA(tilingData, tiling);
AscendC::TPipe tpipe;
#if (defined(DTYPE_VAR))
// tilingKey: indexType * 10 + sortFlag
// indexType: 1=int32, 2=int64(cast), 3=int64(large); sortFlag: 0=非排序, 1=排序
if (TILING_KEY_IS(11)) {
ScatterNdUpdateV2::LinearIndexKernel<true, int> op1(indices, workSpace, tilingData, tpipe);
op1.Process();
AscendC::SyncAll();
tpipe.Destroy();
AscendC::TPipe pipe;
ScatterNdUpdateV2::ScatterNdUpdateV2Kernel<DTYPE_VAR> op2(updates, output, workSpace, tilingData, pipe);
op2.Process();
} else if (TILING_KEY_IS(10)) {
ScatterNdUpdateV2::LinearIndexKernel<false, int> op1(indices, workSpace, tilingData, tpipe);
op1.Process();
AscendC::SyncAll();
tpipe.Destroy();
AscendC::TPipe pipe;
ScatterNdUpdateV2::ScatterNdUpdateV2KernelNoSort<DTYPE_VAR> op2(updates, output, workSpace, tilingData, pipe);
op2.Process();
} else if (TILING_KEY_IS(21)) {
ScatterNdUpdateV2::LinearIndexKernel<true, int64_t> op1(indices, workSpace, tilingData, tpipe);
op1.Process();
AscendC::SyncAll();
tpipe.Destroy();
AscendC::TPipe pipe;
ScatterNdUpdateV2::ScatterNdUpdateV2Kernel<DTYPE_VAR> op2(updates, output, workSpace, tilingData, pipe);
op2.Process();
} else if (TILING_KEY_IS(20)) {
ScatterNdUpdateV2::LinearIndexKernel<false, int64_t> op1(indices, workSpace, tilingData, tpipe);
op1.Process();
AscendC::SyncAll();
tpipe.Destroy();
AscendC::TPipe pipe;
ScatterNdUpdateV2::ScatterNdUpdateV2KernelNoSort<DTYPE_VAR> op2(updates, output, workSpace, tilingData, pipe);
op2.Process();
} else if (TILING_KEY_IS(30)) {
ScatterNdUpdateV2::LargeIndexKernel<DTYPE_VAR> op(indices, updates, output, tilingData, tpipe);
op.Process();
}
#endif
}

View File

@@ -0,0 +1,249 @@
/**
* Copyright (c) 2025-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 scatter_nd_update_v2.h
* \brief Scatter Kernel (Sort)
*/
#include "kernel_operator.h"
#include "kernel_tiling/kernel_tiling.h"
#include "scatter_nd_update_common.h"
namespace ScatterNdUpdateV2 {
template<typename T>
class ScatterNdUpdateV2Kernel {
public:
__aicore__ inline ScatterNdUpdateV2Kernel() = delete;
__aicore__ inline ScatterNdUpdateV2Kernel(
GM_ADDR updates, GM_ADDR output, GM_ADDR workSpace, const ScatterNdUpdateV2TilingData& tiling, TPipe& pipe)
{
InitParams(tiling);
InitBuffers(pipe);
SetGmAddr(updates, output, workSpace, tiling);
}
__aicore__ inline void InitParams(const ScatterNdUpdateV2TilingData& tiling)
{
blockIdx_ = GetBlockIdx();
CalcBlockDistribution(blockIdx_, tiling.scatterTiling.frontNum, tiling.scatterTiling.frontRow,
tiling.scatterTiling.tailRow, computeRow_, start_);
end_ = start_ + computeRow_;
blockNum_ = tiling.linearIndexTiling.blockNum;
blockLength_ = tiling.linearIndexTiling.blockLength;
blockRemainLength_ = tiling.linearIndexTiling.blockRemainLength;
coreNum_ = tiling.linearIndexTiling.coreNum;
scatterLength_ = tiling.scatterTiling.scatterLength;
scatterAlignLength_ = tiling.scatterTiling.scatterAlignLength;
ubLengthForUpdates_ = tiling.scatterTiling.ubLengthForUpdates;
formDim_ = tiling.scatterTiling.formDim;
copyRow_ = tiling.scatterTiling.copyRow;
scatterTileNum_ = tiling.scatterTiling.scatterTileNum;
scatterTileLength_ = tiling.scatterTiling.scatterTileLength;
scatterTileTail_ = tiling.scatterTiling.scatterTileTail;
scatterTileAlignLength_ = tiling.scatterTiling.scatterTileAlignLength;
}
__aicore__ inline void InitBuffers(TPipe& pipe)
{
pipe.InitBuffer(indiceQue_, DOUBLE_BUFFER, blockLength_ * sizeof(int));
pipe.InitBuffer(posIdxQue_, DOUBLE_BUFFER, blockLength_ * sizeof(int));
pipe.InitBuffer(updateQue_, DOUBLE_BUFFER, ubLengthForUpdates_ * sizeof(T));
}
__aicore__ inline void SetGmAddr(GM_ADDR updates, GM_ADDR output, GM_ADDR workSpace, const ScatterNdUpdateV2TilingData& tiling)
{
sortedIndicesGm_.SetGlobalBuffer((__gm__ int*)workSpace);
posIndicesGm_.SetGlobalBuffer((__gm__ int*)workSpace + tiling.linearIndexTiling.sortWorkspace);
updatesGm_.SetGlobalBuffer((__gm__ T*)updates);
outputGm_.SetGlobalBuffer((__gm__ T*)output);
}
__aicore__ inline void Process()
{
for (uint64_t i = 0; i < blockNum_; ++i) {
CopyIndicesIn(i, false);
Compute(i, false);
PipeMte3ToS();
}
if (blockRemainLength_ != 0) {
CopyIndicesIn(blockNum_, true);
Compute(blockNum_, true);
PipeMte3ToS();
}
}
__aicore__ inline void CopyIndicesIn(uint64_t process, bool isTail)
{
uint64_t copyNum = isTail ? blockRemainLength_ : blockLength_;
LocalTensor<int> indiceLocal = indiceQue_.AllocTensor<int>();
LocalTensor<int> posIdxLocal = posIdxQue_.AllocTensor<int>();
uint64_t indicesOffset = isTail ? (blockNum_ * blockLength_) : (process * blockLength_);
DataCopyParams indiceCopyParams{1, static_cast<uint16_t>(copyNum * sizeof(int)), 0, 0};
DataCopyPadParams padParams{true, 0, 0, 0};
DataCopyPad(indiceLocal, sortedIndicesGm_[indicesOffset], indiceCopyParams, padParams);
DataCopyPad(posIdxLocal, posIndicesGm_[indicesOffset], indiceCopyParams, padParams);
PipeMte2ToS();
PipeBarrier<PIPE_V>();
UpdateSearchParam(indiceLocal, isTail);
indiceQue_.EnQue<int>(indiceLocal);
posIdxQue_.EnQue<int>(posIdxLocal);
}
__aicore__ inline void CopyUpdateIn(LocalTensor<T> &updateLocal, uint64_t gmIdx, uint64_t ubIdx, uint64_t tileIdx, uint64_t tileLength)
{
uint64_t gmOffset = gmIdx * scatterLength_ + tileIdx * scatterTileLength_;
uint64_t ubOffset = ubIdx * scatterTileAlignLength_;
DataCopyExtParams updateCopyParams{1, static_cast<uint32_t>(tileLength * sizeof(T)), 0, 0, 0};
DataCopyPadExtParams<T> padParams{true, 0, 0, 0};
DataCopyPad(updateLocal[ubOffset], updatesGm_[gmOffset], updateCopyParams, padParams);
PipeMte2ToS();
}
// 降序数组:二分查找边界
__aicore__ inline int64_t findFirstLt(LocalTensor<int> &indiceLocal, int64_t target, bool isTail)
{
int64_t left = 0;
int64_t right = (isTail ? blockRemainLength_ : blockLength_) - 1;
int64_t res = isTail ? blockRemainLength_ : blockLength_;
while (left <= right) {
int64_t mid = left + (right - left) / 2;
int64_t value = indiceLocal.GetValue(mid);
if (value < target) {
res = mid;
right = mid - 1;
} else {
left = mid + 1;
}
}
return res;
}
__aicore__ inline int64_t findLastGe(LocalTensor<int> &indiceLocal, int64_t target, bool isTail)
{
int64_t left = 0;
int64_t right = (isTail ? blockRemainLength_ : blockLength_) - 1;
int64_t res = -1;
while (left <= right) {
int64_t mid = left + (right - left) / 2;
int64_t value = indiceLocal.GetValue(mid);
if (value >= target) {
res = mid;
left = mid + 1;
} else {
right = mid - 1;
}
}
return res;
}
__aicore__ inline void UpdateSearchParam(LocalTensor<int> &indiceLocal, bool isTail)
{
int64_t searchNum = isTail ? blockRemainLength_ : blockLength_;
leftBound_ = findFirstLt(indiceLocal, end_, isTail);
rightBound_ = findLastGe(indiceLocal, start_, isTail);
isValidBound_ = (leftBound_ < searchNum && rightBound_ != -1 && leftBound_ <= rightBound_);
}
__aicore__ inline void Compute(uint64_t process, bool isTail)
{
LocalTensor<int> indiceLocal = indiceQue_.DeQue<int>();
LocalTensor<int> posIdxLocal = posIdxQue_.DeQue<int>();
if (!isValidBound_) {
indiceQue_.FreeTensor<int>(indiceLocal);
posIdxQue_.FreeTensor<int>(posIdxLocal);
return;
}
for (uint64_t tileIdx = 0; tileIdx < scatterTileNum_; ++tileIdx) {
uint64_t tileLength = (tileIdx == scatterTileNum_ - 1) ? scatterTileTail_ : scatterTileLength_;
uint64_t inUbNum = 0;
LocalTensor<T> updateLocal;
lastProcessedIdx_ = -1;
for (int64_t i = rightBound_; i >= leftBound_; --i) {
if (inUbNum == 0) {
updateLocal = updateQue_.AllocTensor<T>();
}
int64_t posIdx = posIdxLocal.GetValue(i);
CopyUpdateIn(updateLocal, posIdx, inUbNum, tileIdx, tileLength);
inUbNum++;
if (inUbNum == copyRow_) {
updateQue_.EnQue<T>(updateLocal);
CopyOut(inUbNum, i, indiceLocal, posIdxLocal, tileIdx, tileLength);
inUbNum = 0;
}
if (i == leftBound_ && inUbNum != 0) {
updateQue_.EnQue<T>(updateLocal);
CopyOut(inUbNum, i, indiceLocal, posIdxLocal, tileIdx, tileLength);
}
}
}
indiceQue_.FreeTensor<int>(indiceLocal);
posIdxQue_.FreeTensor<int>(posIdxLocal);
}
__aicore__ inline void CopyOut(uint64_t inUbNum, int64_t curIdx, LocalTensor<int> &indiceLocal,
LocalTensor<int> &posIdxLocal, uint64_t tileIdx, uint64_t tileLength)
{
LocalTensor<T> updateLocal = updateQue_.DeQue<T>();
DataCopyExtParams outParams{1, static_cast<uint32_t>(tileLength * sizeof(T)), 0, 0, 0};
for (int64_t i = curIdx + inUbNum - 1; i >= curIdx; --i) {
int64_t curIdxValue = indiceLocal.GetValue(i);
if (curIdxValue == lastProcessedIdx_) continue;
lastProcessedIdx_ = curIdxValue;
uint64_t outOffset = curIdxValue + tileIdx * scatterTileLength_;
uint64_t updateOffset = (curIdx + inUbNum - 1 - i) * scatterTileAlignLength_;
DataCopyPad(outputGm_[outOffset], updateLocal[updateOffset], outParams);
}
PipeMte3ToS();
updateQue_.FreeTensor<T>(updateLocal);
}
private:
GlobalTensor<int> sortedIndicesGm_;
GlobalTensor<int> posIndicesGm_;
GlobalTensor<T> updatesGm_;
GlobalTensor<T> outputGm_;
TQue<TPosition::VECIN, DOUBLE_BUFFER> indiceQue_;
TQue<TPosition::VECIN, DOUBLE_BUFFER> posIdxQue_;
TQue<TPosition::VECOUT, DOUBLE_BUFFER> updateQue_;
uint64_t blockIdx_;
uint64_t computeRow_;
uint64_t start_;
uint64_t end_;
uint64_t blockNum_;
uint64_t blockLength_;
uint64_t blockRemainLength_;
uint64_t scatterLength_;
uint64_t scatterAlignLength_;
uint64_t ubLengthForUpdates_;
uint64_t formDim_;
uint64_t copyRow_;
uint64_t coreNum_;
uint64_t scatterTileNum_;
uint64_t scatterTileLength_;
uint64_t scatterTileTail_;
uint64_t scatterTileAlignLength_;
int64_t leftBound_;
int64_t rightBound_;
bool isValidBound_;
int64_t lastProcessedIdx_;
};
} // namespace ScatterNdUpdateV2