@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
249
csrc/moe/scatter_nd_update_v2/op_kernel/scatter_nd_update_v2.h
Normal file
249
csrc/moe/scatter_nd_update_v2/op_kernel/scatter_nd_update_v2.h
Normal 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
|
||||
Reference in New Issue
Block a user