72
csrc/utils/inc/kernel/comm_args.h
Normal file
72
csrc/utils/inc/kernel/comm_args.h
Normal file
@@ -0,0 +1,72 @@
|
||||
#ifndef COMM_ARGS_H
|
||||
#define COMM_ARGS_H
|
||||
#include <cstdint>
|
||||
|
||||
#define FORCE_INLINE_AICORE __attribute__((always_inline)) inline __aicore__
|
||||
#include "kernel_operator.h"
|
||||
|
||||
namespace Moe {
|
||||
constexpr int CAM_MAX_RANK_SIZE = 384; // Maximum number of NPU cards supported by the communication library
|
||||
|
||||
constexpr int64_t IPC_BUFF_MAX_SIZE = 100 * 1024 * 1024;
|
||||
constexpr int64_t IPC_DATA_OFFSET = 2 * 1024 * 1024; // First 2MB as flag, then 100MB as data storage
|
||||
constexpr int64_t PING_PONG_SIZE = 2;
|
||||
constexpr int64_t UB_SINGLE_DMA_SIZE_MAX = 190 * 1024;
|
||||
constexpr int64_t SMALL_DATA_SIZE = 1 * 1024 * 1024;
|
||||
constexpr int64_t UB_SINGLE_PING_PONG_ADD_SIZE_MAX = UB_SINGLE_DMA_SIZE_MAX / 2;
|
||||
constexpr int UB_ALIGN_SIZE = 32;
|
||||
constexpr int64_t MAGIC_ALIGN_COUNT = UB_ALIGN_SIZE / sizeof(int32_t);
|
||||
|
||||
constexpr uint8_t COMM_NUM = 2; // Size of communication domain
|
||||
constexpr uint8_t COMM_EP_IDX = 0;
|
||||
constexpr uint8_t COMM_TP_IDX = 1;
|
||||
|
||||
constexpr int DFX_COUNT = 50;
|
||||
constexpr int64_t WAIT_SUCCESS = 112233445566;
|
||||
constexpr int64_t IPC_CHUNK_FLAG = 0; // Start offset for send recv, chunk flag region
|
||||
constexpr int64_t MAX_WAIT_ROUND_UNIT = 10 * 1000 * 1000; // Threshold for waiting to get Flag under normal conditions within the same SIO
|
||||
|
||||
constexpr static int32_t UB_HEAD_OFFSET = 96;
|
||||
constexpr static int32_t UB_MID_OFFSET = UB_HEAD_OFFSET + UB_SINGLE_PING_PONG_ADD_SIZE_MAX + UB_ALIGN_SIZE;
|
||||
constexpr static int64_t UB_FLAG_SIZE = 2 * 1024;
|
||||
constexpr static int64_t MAX_CORE_NUM = 48;
|
||||
constexpr static uint64_t STATE_WIN_OFFSET = 900 * 1024;
|
||||
constexpr static int64_t COMPARE_ALIGN_SIZE = 256;
|
||||
|
||||
constexpr static int64_t UB_SINGLE_TOTAL_SIZE_MAX = 192 * 1024;
|
||||
constexpr static int64_t START_OFFSET_FOR_SHARE = 512;
|
||||
|
||||
enum Op : int {
|
||||
COPYONLY = -1,
|
||||
ADD = 0,
|
||||
MUL = 1,
|
||||
MAX = 2,
|
||||
MIN = 3
|
||||
};
|
||||
|
||||
struct CommArgs {
|
||||
int rank = 0; // attr rank_id, global rank
|
||||
int localRank = -1;
|
||||
int rankSize = 0; // global rank size
|
||||
int localRankSize = -1; // This parameter refers to the number of cards interconnected in fullmesh
|
||||
uint32_t extraFlag = 0; // 32 bit map, the specific meaning of each bit is above in this file
|
||||
int testFlag = 0;
|
||||
GM_ADDR peerMems[CAM_MAX_RANK_SIZE] = {}; // Buffer obtained from initialization, all allreduce is the same parameter
|
||||
/**
|
||||
* @param sendCountMatrix One-dimensional array with a size of rankSize*rankSize
|
||||
* eg: The value of sendCountMatrix[1] corresponds to the [0][1] of the two-dimensional array, indicating the number of data that card 0 needs to send to card 1
|
||||
*/
|
||||
int64_t sendCountMatrix[CAM_MAX_RANK_SIZE * CAM_MAX_RANK_SIZE] = {}; // for all2allvc
|
||||
int64_t sendCounts[CAM_MAX_RANK_SIZE] = {}; // for all2allv
|
||||
int64_t sdispls[CAM_MAX_RANK_SIZE] = {}; // for all2allv
|
||||
int64_t recvCounts[CAM_MAX_RANK_SIZE] = {}; // for all2allv
|
||||
int64_t rdispls[CAM_MAX_RANK_SIZE] = {}; // for all2allv
|
||||
int64_t batchSize;
|
||||
int64_t hiddenSize;
|
||||
int64_t topk;
|
||||
int64_t sharedExpertRankNum;
|
||||
int64_t expertNumPerRank;
|
||||
int64_t dfx[DFX_COUNT] = {};
|
||||
};
|
||||
}
|
||||
#endif // COMM_ARGS_H
|
||||
68
csrc/utils/inc/kernel/data_copy.h
Normal file
68
csrc/utils/inc/kernel/data_copy.h
Normal file
@@ -0,0 +1,68 @@
|
||||
#ifndef CAM_DATACOPY_GM2GM_H
|
||||
#define CAM_DATACOPY_GM2GM_H
|
||||
#include <type_traits>
|
||||
#include "comm_args.h"
|
||||
|
||||
using namespace AscendC;
|
||||
using namespace Moe;
|
||||
|
||||
template <typename T>
|
||||
FORCE_INLINE_AICORE void SetAtomicOpType(int op)
|
||||
{
|
||||
switch (op) {
|
||||
case ADD:
|
||||
AscendC::SetAtomicAdd<T>();
|
||||
break;
|
||||
case MUL:
|
||||
// Ignore setting the atomic register when performing mul
|
||||
break;
|
||||
case MAX:
|
||||
AscendC::SetAtomicMax<T>();
|
||||
break;
|
||||
case MIN:
|
||||
AscendC::SetAtomicMin<T>();
|
||||
break;
|
||||
default:
|
||||
AscendC::SetAtomicNone();
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
FORCE_INLINE_AICORE void CpUB2GM(__gm__ T *gmAddr, __ubuf__ T *ubAddr, uint32_t size)
|
||||
{
|
||||
LocalTensor<uint8_t> ubTensor;
|
||||
GlobalTensor<uint8_t> gmTensor;
|
||||
DataCopyExtParams dataCopyParams(1, size, 0, 0, 0);
|
||||
ubTensor.address_.logicPos = static_cast<uint8_t>(TPosition::VECIN);
|
||||
ubTensor.address_.bufferAddr = reinterpret_cast<uint64_t>(ubAddr);
|
||||
gmTensor.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t *>(gmAddr));
|
||||
DataCopyPad(gmTensor, ubTensor, dataCopyParams);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
FORCE_INLINE_AICORE void CpGM2UB(__ubuf__ T *ubAddr, __gm__ T *gmAddr, uint32_t size)
|
||||
{
|
||||
LocalTensor<uint8_t> ubTensor;
|
||||
GlobalTensor<uint8_t> gmTensor;
|
||||
DataCopyExtParams dataCopyParams(1, size, 0, 0, 0);
|
||||
ubTensor.address_.logicPos = static_cast<uint8_t>(TPosition::VECIN);
|
||||
ubTensor.address_.bufferAddr = reinterpret_cast<uint64_t>(ubAddr);
|
||||
gmTensor.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t *>(gmAddr));
|
||||
DataCopyPadExtParams<uint8_t> padParams;
|
||||
DataCopyPad(ubTensor, gmTensor, dataCopyParams, padParams);
|
||||
}
|
||||
|
||||
template<typename T>
|
||||
FORCE_INLINE_AICORE void CopyUB2UB(__ubuf__ T *dst, __ubuf__ T *src, const uint32_t calCount)
|
||||
{
|
||||
LocalTensor<T> srcTensor;
|
||||
LocalTensor<T> dstTensor;
|
||||
TBuffAddr srcAddr, dstAddr;
|
||||
srcAddr.bufferAddr = reinterpret_cast<uint64_t>(src);
|
||||
dstAddr.bufferAddr = reinterpret_cast<uint64_t>(dst);
|
||||
srcTensor.SetAddr(srcAddr);
|
||||
dstTensor.SetAddr(dstAddr);
|
||||
DataCopy(dstTensor, srcTensor, calCount);
|
||||
}
|
||||
|
||||
#endif // CAM_DATACOPY_GM2GM_H
|
||||
121
csrc/utils/inc/kernel/dropmask.h
Normal file
121
csrc/utils/inc/kernel/dropmask.h
Normal file
@@ -0,0 +1,121 @@
|
||||
/**
|
||||
* Copyright (c) 2024 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.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 dropmask.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef DROPMASK_H
|
||||
#define DROPMASK_H
|
||||
|
||||
#include "util.h"
|
||||
|
||||
using AscendC::DROPOUT_MODE_BIT_MISALIGN;
|
||||
using AscendC::DropOutShapeInfo;
|
||||
using AscendC::DropOut;
|
||||
|
||||
struct DropMaskInfo {
|
||||
// for compute dropout mask offset
|
||||
// 参数按B N G S1 S2全部切分设置进行偏移计算,没有切分的轴对应的参数设置为合适的0或者原始值
|
||||
int64_t n2G; // n2 * g
|
||||
int64_t gSize; // g
|
||||
int64_t s1Size; // s1
|
||||
int64_t s2Size; // s2
|
||||
int64_t gOutIdx; // g out index
|
||||
int64_t bSSOffset; // boidx * s1 * s2 ===bSSOffset
|
||||
int64_t n2OutIdx; // n out index
|
||||
int64_t s1OutIdx; // s1 out index ===s1oIdx
|
||||
int64_t s1InnerIdx; // s1 inner index, 配比 ===loopIdx
|
||||
int64_t s1BaseSize; // S1基本块大小
|
||||
int64_t splitS1BaseSize; // s1 split size ===vec1S1BaseSize
|
||||
int64_t s2StartIdx; // s2 start index
|
||||
int64_t s2Idx; // s2 index =====s2LoopCount
|
||||
int64_t s2BaseNratioSize; // s2的配比长度: s2BaseSize(S2基本块大小) * nRatio
|
||||
|
||||
// for copy in dropout mask
|
||||
uint32_t s1CopySize;
|
||||
uint32_t s2CopySize;
|
||||
int64_t s2TotalSize;
|
||||
|
||||
// for compute dropout mask
|
||||
uint32_t firstAxis;
|
||||
uint32_t lstAxis;
|
||||
uint32_t maskLstAxis;
|
||||
int64_t vecCoreOffset = 0;
|
||||
float keepProb;
|
||||
|
||||
bool boolMode;
|
||||
};
|
||||
|
||||
template <bool hasDrop>
|
||||
__aicore__ inline int64_t ComputeDropOffset(DropMaskInfo &dropMaskInfo)
|
||||
{
|
||||
if constexpr (hasDrop == true) {
|
||||
// boidx * n2 * g* s1 * s2
|
||||
int64_t bOffset = dropMaskInfo.bSSOffset * dropMaskInfo.n2G;
|
||||
// n2oIdx * g * s1 *s2
|
||||
int64_t n2Offset = dropMaskInfo.n2OutIdx * dropMaskInfo.gSize * dropMaskInfo.s1Size * dropMaskInfo.s2Size;
|
||||
// goIdx * s1 * s2
|
||||
int64_t gOffset = dropMaskInfo.gOutIdx * dropMaskInfo.s1Size * dropMaskInfo.s2Size;
|
||||
// s1oIdx * s1BaseSize * s2Size + s1innerindex * vec1S1BaseSize * s2Size
|
||||
int64_t s1Offset = (dropMaskInfo.s1OutIdx * dropMaskInfo.s1BaseSize + dropMaskInfo.vecCoreOffset +
|
||||
dropMaskInfo.s1InnerIdx * dropMaskInfo.splitS1BaseSize) * dropMaskInfo.s2Size;
|
||||
// s2StartIdx + s2index * s2BaseNratioSize
|
||||
int64_t s2Offset = dropMaskInfo.s2StartIdx + dropMaskInfo.s2Idx * dropMaskInfo.s2BaseNratioSize;
|
||||
return bOffset + n2Offset + gOffset + s1Offset + s2Offset;
|
||||
} else {
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
|
||||
template <bool hasDrop>
|
||||
__aicore__ inline void CopyInDropMask(LocalTensor<uint8_t>&dstTensor, GlobalTensor<uint8_t>& srcBoolTensor,
|
||||
GlobalTensor<uint8_t>& srcByteTensor, DropMaskInfo &dropMaskInfo, int64_t alignedSize = blockBytes)
|
||||
{
|
||||
if constexpr (hasDrop == true) {
|
||||
int64_t dropMaskOffset = ComputeDropOffset<hasDrop>(dropMaskInfo);
|
||||
if (unlikely(dropMaskInfo.boolMode)) {
|
||||
BoolCopyIn(dstTensor, srcBoolTensor, dropMaskOffset,
|
||||
dropMaskInfo.s1CopySize, dropMaskInfo.s2CopySize, dropMaskInfo.s2TotalSize, alignedSize);
|
||||
} else {
|
||||
Bit2Int8CopyIn(dstTensor, srcByteTensor, dropMaskOffset, 1,
|
||||
dropMaskInfo.s1CopySize, dropMaskInfo.s2CopySize, dropMaskInfo.s2TotalSize, alignedSize);
|
||||
}
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, bool hasDrop>
|
||||
__aicore__ inline void ComputeDropMask(LocalTensor<T>& dstTensor, LocalTensor<T>& srcTensor,
|
||||
LocalTensor<uint8_t>& dropoutBuffer, LocalTensor<uint8_t>& tmpDropBuffer, DropMaskInfo &dropMaskInfo)
|
||||
{
|
||||
if constexpr (hasDrop == true) {
|
||||
DropOutShapeInfo dropOutShapeInfo;
|
||||
dropOutShapeInfo.firstAxis = dropMaskInfo.firstAxis;
|
||||
dropOutShapeInfo.srcLastAxis = dropMaskInfo.lstAxis;
|
||||
|
||||
if (unlikely(dropMaskInfo.boolMode)) {
|
||||
dropOutShapeInfo.maskLastAxis = CeilDiv(dropMaskInfo.maskLstAxis, blockBytes) * blockBytes;
|
||||
DropOut(dstTensor, srcTensor, dropoutBuffer, tmpDropBuffer, dropMaskInfo.keepProb, dropOutShapeInfo);
|
||||
} else {
|
||||
dropOutShapeInfo.maskLastAxis = CeilDiv(dropMaskInfo.maskLstAxis / byteBitRatio, blockBytes) * blockBytes;
|
||||
if (likely(dropMaskInfo.lstAxis / byteBitRatio % blockBytes == 0)) {
|
||||
DropOut(dstTensor, srcTensor, dropoutBuffer, tmpDropBuffer, dropMaskInfo.keepProb, dropOutShapeInfo);
|
||||
} else {
|
||||
DropOut<T, false, DROPOUT_MODE_BIT_MISALIGN>(dstTensor, srcTensor, dropoutBuffer, tmpDropBuffer,
|
||||
dropMaskInfo.keepProb, dropOutShapeInfo);
|
||||
}
|
||||
}
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
#endif // DROPMASK_H
|
||||
365
csrc/utils/inc/kernel/moe_distribute_base.h
Executable file
365
csrc/utils/inc/kernel/moe_distribute_base.h
Executable file
@@ -0,0 +1,365 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* 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 moe_distribute_base.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef MOE_DISTRIBUTE_BASE_H
|
||||
#define MOE_DISTRIBUTE_BASE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
|
||||
constexpr uint32_t LOCAL_NOTIFY_MAX_NUM = 64;
|
||||
constexpr uint32_t LOCAL_STREAM_MAX_NUM = 19U;
|
||||
constexpr uint32_t AICPU_OP_NOTIFY_MAX_NUM = 2;
|
||||
constexpr uint32_t AICPU_MAX_RANK_NUM = 128 * 1024;
|
||||
constexpr uint32_t TIME_CYCLE = 50; // 系统cycle数转换成时间的基准单位,固定为50
|
||||
|
||||
struct HcclSignalInfo {
|
||||
uint64_t resId; // 在代表event时为eventid,notify时为notifyid
|
||||
uint64_t addr;
|
||||
uint32_t devId;
|
||||
uint32_t tsId;
|
||||
uint32_t rankId;
|
||||
uint32_t flag;
|
||||
};
|
||||
|
||||
struct ListCommon {
|
||||
uint64_t nextHost;
|
||||
uint64_t preHost;
|
||||
uint64_t nextDevice;
|
||||
uint64_t preDevice;
|
||||
};
|
||||
|
||||
struct HcclStreamInfo {
|
||||
int32_t streamIds;
|
||||
uint32_t sqIds;
|
||||
uint32_t cqIds; // 记录物理cqId
|
||||
uint32_t logicCqids; // 记录逻辑cqId
|
||||
};
|
||||
|
||||
struct LocalResInfoV2 {
|
||||
uint32_t streamNum;
|
||||
uint32_t signalNum;
|
||||
HcclSignalInfo localSignals[LOCAL_NOTIFY_MAX_NUM];
|
||||
HcclStreamInfo streamInfo[LOCAL_STREAM_MAX_NUM];
|
||||
HcclStreamInfo mainStreamInfo;
|
||||
HcclSignalInfo aicpuOpNotify[AICPU_OP_NOTIFY_MAX_NUM]; // 集合通信AICPU展开资源
|
||||
ListCommon nextTagRes; // HccltagLocalResV2
|
||||
};
|
||||
|
||||
enum class rtFloatOverflowMode_t {
|
||||
RT_OVERFLOW_MODE_SATURATION = 0,
|
||||
RT_OVERFLOW_MODE_INFNAN,
|
||||
RT_OVERFLOW_MODE_UNDEF,
|
||||
};
|
||||
|
||||
struct AlgoTopoInfo {
|
||||
uint32_t userRank; // 通信域 RankID
|
||||
uint32_t userRankSize; // 通信域的Rank数量
|
||||
int32_t deviceLogicId;
|
||||
bool isSingleMeshAggregation;
|
||||
uint32_t deviceNumPerAggregation; // 每个Module中的Device数量
|
||||
uint32_t superPodNum; // 集群中总的超节点数
|
||||
uint32_t devicePhyId;
|
||||
uint32_t topoType; // TopoType
|
||||
uint32_t deviceType;
|
||||
uint32_t serverNum;
|
||||
uint32_t meshAggregationRankSize;
|
||||
uint32_t multiModuleDiffDeviceNumMode;
|
||||
uint32_t multiSuperPodDiffServerNumMode;
|
||||
uint32_t realUserRank;
|
||||
bool isDiffDeviceModule;
|
||||
bool isDiffDeviceType;
|
||||
uint32_t gcdDeviceNumPerAggregation;
|
||||
uint32_t moduleNum;
|
||||
uint32_t isUsedRdmaRankPairNum;
|
||||
uint64_t isUsedRdmaRankPair;
|
||||
uint32_t pairLinkCounterNum;
|
||||
uint64_t pairLinkCounter;
|
||||
uint32_t nicNum;
|
||||
uint64_t nicList; // niclist数组指针
|
||||
uint64_t complanRankLength; // complanRank占用的字节数
|
||||
uint64_t complanRank; // 指针
|
||||
uint64_t bridgeRankNum; // bridgeRank占用的个数
|
||||
uint64_t bridgeRank; // 指针
|
||||
uint64_t serverAndsuperPodRankLength; // serverAndsuperPodRank占用的字节数
|
||||
uint64_t serverAndsuperPodRank; // 指针
|
||||
};
|
||||
|
||||
struct HcclOpConfig {
|
||||
uint8_t deterministic; //确定性计算开关
|
||||
uint8_t retryEnable; // 是否重执行
|
||||
uint8_t highPerfEnable;
|
||||
uint8_t padding[5]; // 大小需要64By对齐,未来添加参数时减小padding
|
||||
uint8_t linkTimeOut[8]; // 发送超时时长
|
||||
uint64_t notifyWaitTime; // 超时时长,同HCCL_EXEC_TIMEOUT
|
||||
uint32_t retryHoldTime;
|
||||
uint32_t retryIntervalTime;
|
||||
bool interHccsDisable = false; //使能rdma开关
|
||||
rtFloatOverflowMode_t floatOverflowMode = rtFloatOverflowMode_t::RT_OVERFLOW_MODE_UNDEF;
|
||||
uint32_t multiQpThreshold = 512; // 多QP每个QP分担数据量最小阈值
|
||||
};
|
||||
|
||||
struct HcclMC2WorkSpace {
|
||||
uint64_t workSpace;
|
||||
uint64_t workSpaceSize;
|
||||
};
|
||||
|
||||
struct RemoteResPtr {
|
||||
uint64_t nextHostPtr;
|
||||
uint64_t nextDevicePtr;
|
||||
};
|
||||
|
||||
struct HDCommunicateParams {
|
||||
uint64_t hostAddr { 0 };
|
||||
uint64_t deviceAddr { 0 };
|
||||
uint64_t readCacheAddr { 0 };
|
||||
uint32_t devMemSize{ 0 };
|
||||
uint32_t buffLen{ 0 };
|
||||
uint32_t flag{ 0 };
|
||||
};
|
||||
|
||||
struct HcclRankRelationResV2 {
|
||||
uint32_t remoteUsrRankId;
|
||||
uint32_t remoteWorldRank;
|
||||
uint64_t windowsIn;
|
||||
uint64_t windowsOut;
|
||||
uint64_t windowsExp;
|
||||
ListCommon nextTagRes;
|
||||
};
|
||||
|
||||
struct HcclOpResParam {
|
||||
// 本地资源
|
||||
HcclMC2WorkSpace mc2WorkSpace;
|
||||
uint32_t localUsrRankId; // usrrankid
|
||||
uint32_t rankSize; // 通信域内total rank个数
|
||||
uint64_t winSize; // 每个win大小,静态图时,可能是0,如果通信域内也有动态图,则可能为非0
|
||||
uint64_t localWindowsIn; // 全F为无效值
|
||||
uint64_t localWindowsOut; // 全F为无效值
|
||||
char hcomId[128];
|
||||
// aicore识别remote window
|
||||
uint64_t winExpSize;
|
||||
uint64_t localWindowsExp;
|
||||
uint32_t rWinStart; // 为HcclRankRelationRes起始位置
|
||||
uint32_t rWinOffset; // 为HcclRemoteRes的大小
|
||||
uint64_t version;
|
||||
LocalResInfoV2 localRes;
|
||||
AlgoTopoInfo topoInfo;
|
||||
|
||||
// 外部配置参数
|
||||
HcclOpConfig config;
|
||||
uint64_t hostStateInfo;
|
||||
uint64_t aicpuStateInfo;
|
||||
uint64_t lockAddr;
|
||||
uint32_t rsv[16];
|
||||
uint32_t notifysize; // RDMA场景使用,910B/910_93为4B,其余芯片为8B
|
||||
uint32_t remoteResNum; // 有效的remoteResNum
|
||||
RemoteResPtr remoteRes[AICPU_MAX_RANK_NUM]; //数组指针,指向HcclRankRelationResV2,下标为remoteUserRankId
|
||||
|
||||
// communicate retry
|
||||
HDCommunicateParams kfcControlTransferH2DParams;
|
||||
HDCommunicateParams kfcStatusTransferD2HParams;
|
||||
uint64_t tinyMem; // for all2all
|
||||
uint64_t tinyMemSize;
|
||||
// 零拷贝场景使用
|
||||
uint64_t zeroCopyHeadPtr;
|
||||
uint64_t zeroCopyTailPtr;
|
||||
uint64_t zeroCopyRingBuffer;
|
||||
uint64_t zeroCopyIpcPtrs[16]; // 保存集合通信时每个对端的输入输出内存地址
|
||||
uint32_t zeroCopyDevicePhyId[16]; // 保存每个rank对应的物理卡Id
|
||||
|
||||
bool utraceStatusFlag;
|
||||
};
|
||||
|
||||
// Transport 内存类型
|
||||
enum class HcclAiRMAMemType : uint32_t {
|
||||
LOCAL_INPUT = 0,
|
||||
REMOTE_INPUT,
|
||||
|
||||
LOCAL_OUTPUT,
|
||||
REMOTE_OUTPUT,
|
||||
|
||||
// 可透传更多的内存,可在MAX_NUM之前追加,例如:
|
||||
// LOCAL_EXP,
|
||||
// REMOTE_EXP,
|
||||
MAX_NUM
|
||||
};
|
||||
|
||||
// Transport 内存信息
|
||||
struct HcclAiRMAMemInfo {
|
||||
uint32_t memMaxNum{0}; // 最大内存数量,等于 HcclAiRMAMemType::MAX_NUM
|
||||
uint32_t sizeOfMemDetails{0}; // sizeof(MemDetails),用于内存校验和偏移计算
|
||||
uint64_t memDetailPtr{0}; // MemDetails数组首地址, 个数: HcclAiRMAMemType::MAX_NUM
|
||||
// 可往后追加字段
|
||||
};
|
||||
|
||||
// 全部 Transport QP/Mem 信息
|
||||
struct HcclAiRMAInfo {
|
||||
uint32_t curRankId{0}; // 当前rankId
|
||||
uint32_t rankNum{0}; // rank数量
|
||||
uint32_t qpNum{0}; // 单个Transport的QP数量
|
||||
|
||||
uint32_t sizeOfAiRMAWQ{0}; // sizeof(HcclAiRMAWQ)
|
||||
uint32_t sizeOfAiRMACQ{0}; // sizeof(HcclAiRMACQ)
|
||||
uint32_t sizeOfAiRMAMem{0}; // sizeof(HcclAiRMAMemInfo)
|
||||
|
||||
// HcclAiRMAWQ二维数组首地址
|
||||
// QP个数: rankNum * qpNum
|
||||
// 计算偏移获取SQ指针:sqPtr + (dstRankId * qpNum + qpIndex) * sizeOfAiRMAWQ
|
||||
// 0 <= qpIndex < qpNum
|
||||
uint64_t sqPtr{0};
|
||||
|
||||
// HcclAiRMACQ二维数组首地址
|
||||
// QP个数: rankNum * qpNum
|
||||
// 计算偏移获取SCQ指针:scqPtr + (dstRankId * qpNum + qpIndex) * sizeOfAiRMACQ
|
||||
// 0 <= qpIndex < qpNum
|
||||
uint64_t scqPtr{0};
|
||||
|
||||
// HcclAiRMAWQ二维数组首地址
|
||||
// QP个数: rankNum * qpNum
|
||||
// 计算偏移获取RQ指针:rqPtr + (dstRankId * qpNum + qpIndex) * sizeOfAiRMAWQ
|
||||
// 0 <= qpIndex < qpNum
|
||||
uint64_t rqPtr{0};
|
||||
|
||||
// HcclAiRMACQ二维数组首地址
|
||||
// QP个数: rankNum * qpNum
|
||||
// 计算偏移获取RCQ指针: rcqPtr + (dstRankId * qpNum + qpIndex) * sizeOfAiRMACQ
|
||||
// 0 <= qpIndex < qpNum
|
||||
uint64_t rcqPtr{0};
|
||||
|
||||
// HcclAivMemInfo一维数组
|
||||
// 内存信息个数: rankNum
|
||||
// 计算偏移获取内存信息指针: memPtr + rankId * sizeOfAiRMAMem
|
||||
// srcRankId 获取自身内存信息,dstRankId 获取 Transport 内存信息
|
||||
uint64_t memPtr{0};
|
||||
// 可往后追加字段
|
||||
};
|
||||
struct CombinedCapability {
|
||||
uint64_t dataplaneModeBitmap;
|
||||
};
|
||||
|
||||
struct HcclA2CombineOpParam {
|
||||
uint64_t workSpace; // Address for communication between client and server,
|
||||
// hccl requests and clears
|
||||
uint64_t workSpaceSize; // Space for communication between client and server
|
||||
uint32_t rankId; // id of this rank
|
||||
uint32_t rankNum; // num of ranks in this comm group
|
||||
uint64_t winSize; // size of each windows memory
|
||||
uint64_t windowsIn[AscendC::HCCL_MAX_RANK_NUM]; // windows address for input, windowsIn[rankId] corresponds
|
||||
// to the local card address,
|
||||
// and others are cross-card mapping addresses.
|
||||
uint64_t windowsOut[AscendC::HCCL_MAX_RANK_NUM]; // windows address for output, windowsOut[rankId] corresponds
|
||||
// to the local card address,
|
||||
// and others are cross-card mapping addresses.
|
||||
uint8_t res[8328];
|
||||
uint8_t multiFlag;
|
||||
__gm__ AscendC::IbVerbsData *data;
|
||||
uint64_t dataSize;
|
||||
// 追加字段
|
||||
uint64_t sizeOfAiRMAInfo; // sizeof(HcclAiRMAInfo)
|
||||
uint64_t aiRMAInfo; // HcclAiRMAInfo* 单个结构体指针
|
||||
|
||||
CombinedCapability* capability; // address of the communication capability information structure on the Device
|
||||
uint64_t capabilitySize; // size of the communication capability information structure
|
||||
};
|
||||
enum class DataplaneMode : uint32_t {
|
||||
HOST = 0,
|
||||
AICPU = 1,
|
||||
AIV = 2,
|
||||
};
|
||||
|
||||
enum class DBMode : int32_t {
|
||||
INVALID_DB = -1,
|
||||
HW_DB = 0,
|
||||
SW_DB
|
||||
};
|
||||
|
||||
struct HcclAiRMAWQ {
|
||||
uint32_t wqn{0};
|
||||
uint64_t bufAddr{0};
|
||||
uint32_t wqeSize{0};
|
||||
uint32_t depth{0};
|
||||
uint64_t headAddr{0};
|
||||
uint64_t tailAddr{0};
|
||||
DBMode dbMode{DBMode::INVALID_DB}; // 0-hw/1-sw
|
||||
uint64_t dbAddr{0};
|
||||
uint32_t sl{0};
|
||||
};
|
||||
|
||||
struct HcclAiRMACQ {
|
||||
uint32_t cqn{0};
|
||||
uint64_t bufAddr{0};
|
||||
uint32_t cqeSize{0};
|
||||
uint32_t depth{0};
|
||||
uint64_t headAddr{0};
|
||||
uint64_t tailAddr{0};
|
||||
DBMode dbMode{DBMode::INVALID_DB}; // 0-hw/1-sw
|
||||
uint64_t dbAddr{0};
|
||||
};
|
||||
|
||||
struct hns_roce_rc_sq_wqe {
|
||||
uint32_t byte_4;
|
||||
uint32_t msg_len;
|
||||
uint32_t immtdata;
|
||||
uint32_t byte_16;
|
||||
uint32_t byte_20;
|
||||
uint32_t rkey;
|
||||
uint64_t remoteVA;
|
||||
};
|
||||
|
||||
|
||||
struct hns_roce_lite_wqe_data_seg {
|
||||
uint32_t len;
|
||||
uint32_t lkey;
|
||||
uint64_t localVA;
|
||||
};
|
||||
|
||||
__aicore__ inline void cacheWriteThrough(__gm__ uint8_t* sourceAddr, uint64_t length) {
|
||||
__gm__ uint8_t* start =
|
||||
(__gm__ uint8_t*)((uint64_t)sourceAddr / AscendC::CACHE_LINE_SIZE * AscendC::CACHE_LINE_SIZE);
|
||||
__gm__ uint8_t* end =
|
||||
(__gm__ uint8_t*)(((uint64_t)sourceAddr + length) / AscendC::CACHE_LINE_SIZE * AscendC::CACHE_LINE_SIZE);
|
||||
AscendC::GlobalTensor<uint8_t> global;
|
||||
global.SetGlobalBuffer(start);
|
||||
for (uint32_t i = 0; i <= end - start; i += AscendC::CACHE_LINE_SIZE) {
|
||||
AscendC::DataCacheCleanAndInvalid<uint8_t, AscendC::CacheLine::SINGLE_CACHE_LINE,
|
||||
AscendC::DcciDst::CACHELINE_OUT>(global[i]);
|
||||
}
|
||||
}
|
||||
__aicore__ inline DataplaneMode GetDataplaneMode(GM_ADDR contextGM0) {
|
||||
__gm__ HcclA2CombineOpParam *winContext_ = (__gm__ HcclA2CombineOpParam *)contextGM0;
|
||||
CombinedCapability* capability = winContext_->capability;
|
||||
uint64_t capabilitySize = winContext_->capabilitySize;
|
||||
DataplaneMode dataplaneMode = DataplaneMode::AICPU;
|
||||
if (capability == 0) {
|
||||
return dataplaneMode;
|
||||
}
|
||||
uint64_t dataplaneModeBitmap = capability->dataplaneModeBitmap;
|
||||
if ((dataplaneModeBitmap & 0x04) == 0x04) {
|
||||
dataplaneMode = DataplaneMode::AIV;
|
||||
}
|
||||
return dataplaneMode;
|
||||
}
|
||||
|
||||
__aicore__ inline int64_t GetCurrentTimestampUs()
|
||||
{
|
||||
return AscendC::GetSystemCycle() / TIME_CYCLE;
|
||||
}
|
||||
|
||||
__aicore__ inline void RecordRankCommDuration(AscendC::LocalTensor<int32_t> performanceInfoU32Tensor, uint32_t rankId, int64_t startTime)
|
||||
{
|
||||
int64_t endTime = GetCurrentTimestampUs();
|
||||
int32_t duration = static_cast<int32_t>(endTime - startTime); // int32_t可以表示2^31(us),约35min在实际场景下满足需要
|
||||
performanceInfoU32Tensor.SetValue(rankId * sizeof(int64_t) / sizeof(int32_t), duration); // 使用int32_t是因为atomicAdd不支持int64_t类型,这里只赋值到int64_t的低32位。
|
||||
}
|
||||
#endif // MOE_DISTRIBUTE_BASE_H
|
||||
483
csrc/utils/inc/kernel/pse.h
Normal file
483
csrc/utils/inc/kernel/pse.h
Normal file
@@ -0,0 +1,483 @@
|
||||
/**
|
||||
* Copyright (c) 2024 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.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 pse.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef FLASH_ATTENTION_SCORE_PSE_H
|
||||
#define FLASH_ATTENTION_SCORE_PSE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "util.h"
|
||||
|
||||
constexpr static int64_t pseS1S2 = 0;
|
||||
constexpr static int64_t pse1S2 = 1;
|
||||
constexpr static int64_t pseSlopeBn = 2;
|
||||
constexpr static int64_t pseSlopeN = 3;
|
||||
|
||||
constexpr static uint8_t pseEncodeALibiS2Full = 0x11;
|
||||
|
||||
enum class PseTypeEnum {
|
||||
PSE_OUTER_MUL_ADD_TYPE = 0, // default
|
||||
PSE_OUTER_ADD_MUL_TYPE,
|
||||
PSE_INNER_MUL_ADD_TYPE,
|
||||
PSE_INNER_MUL_ADD_SQRT_TYPE,
|
||||
PSE_INVALID_TYPE
|
||||
};
|
||||
|
||||
struct PseInfo {
|
||||
int64_t blockCount;
|
||||
int64_t bSSOffset; // boidx * s1 * s2
|
||||
int64_t boIdx;
|
||||
int64_t gSize;
|
||||
int64_t goIdx;
|
||||
int64_t loopIdx;
|
||||
int64_t n2G;
|
||||
int64_t n2oIdx;
|
||||
int64_t pseBSize;
|
||||
int64_t pseS1Size; // for alibi
|
||||
int64_t pseS2ComputeSize; // for alibi, do not need assignment
|
||||
int64_t pseS2Size; // for alibi
|
||||
uint32_t pseShapeType;
|
||||
int64_t readS2Size; // for alibi, do not need assignment
|
||||
int64_t s1BaseSize;
|
||||
int64_t s1Size;
|
||||
int64_t s1oIdx;
|
||||
int64_t s2AlignedSize;
|
||||
int64_t s2BaseNratioSize;
|
||||
int64_t s2LoopCount;
|
||||
int64_t s2RealSize;
|
||||
int64_t s2Size;
|
||||
int64_t s2SizeAcc; // accumulated sum of s2 size
|
||||
int64_t s2StartIdx;
|
||||
int64_t vec1S1BaseSize;
|
||||
int64_t vec1S1RealSize;
|
||||
uint32_t pseEncodeType; // for distinguish alibi
|
||||
uint32_t pseType; // 0: outer, mul-add 1:outer, add-mul 2:inner, mul-add 3:inner, mul-add-sqrt
|
||||
int64_t pseAlibiBaseS1;
|
||||
int64_t pseAlibiBaseS2;
|
||||
int64_t qStartIdx;
|
||||
int64_t kvStartIdx;
|
||||
int64_t vecCoreOffset = 0;
|
||||
bool needCast;
|
||||
bool align8 = false;
|
||||
bool pseEndogenous = false;
|
||||
};
|
||||
|
||||
template <typename INPUT_T, bool hasPse>
|
||||
__aicore__ inline void DataCopyInCommon(LocalTensor<INPUT_T> &dstTensor, GlobalTensor<INPUT_T> &srcTensor, int64_t offset,
|
||||
int64_t s1Size, int64_t s2Size, int64_t actualS2Len, int32_t dtypeSize,
|
||||
int32_t alignedS2Size)
|
||||
{
|
||||
if constexpr (hasPse == true) {
|
||||
uint32_t shapeArray[] = {static_cast<uint32_t>(s1Size), static_cast<uint32_t>(alignedS2Size)};
|
||||
dstTensor.SetShapeInfo(ShapeInfo(2, shapeArray, DataFormat::ND));
|
||||
dstTensor.SetSize(s1Size * alignedS2Size);
|
||||
DataCopyParams dataCopyParams;
|
||||
dataCopyParams.blockCount = s1Size;
|
||||
dataCopyParams.blockLen = CeilDiv(s2Size * dtypeSize, blockBytes); // 单位32B
|
||||
dataCopyParams.dstStride = alignedS2Size * dtypeSize / blockBytes - dataCopyParams.blockLen; // gap
|
||||
if (actualS2Len * dtypeSize % blockBytes == 0) {
|
||||
dataCopyParams.srcStride =
|
||||
(actualS2Len * dtypeSize - dataCopyParams.blockLen * blockBytes) / blockBytes; // srcGap
|
||||
DataCopy(dstTensor, srcTensor[offset], dataCopyParams);
|
||||
} else {
|
||||
dataCopyParams.blockLen = s2Size * dtypeSize; // 单位Byte
|
||||
dataCopyParams.srcStride = (actualS2Len * dtypeSize - dataCopyParams.blockLen);
|
||||
dataCopyParams.dstStride = (alignedS2Size - s2Size) * dtypeSize / blockBytes;
|
||||
DataCopyPadParams dataCopyPadParams;
|
||||
dataCopyPadParams.isPad = false;
|
||||
DataCopyPad(dstTensor, srcTensor[offset], dataCopyParams, dataCopyPadParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename INPUT_T, bool hasPse>
|
||||
__aicore__ inline void DataCopyIn(LocalTensor<INPUT_T> &dstTensor, GlobalTensor<INPUT_T> &srcTensor, int64_t offset,
|
||||
int64_t s1Size, int64_t s2Size, int64_t actualS2Len, int64_t alignedSize = 16)
|
||||
{
|
||||
if constexpr (hasPse == true) {
|
||||
int32_t dtypeSize = sizeof(INPUT_T);
|
||||
int32_t alignedS2Size = CeilDiv(s2Size, alignedSize) * alignedSize;
|
||||
DataCopyInCommon<INPUT_T, hasPse>(dstTensor, srcTensor, offset, s1Size, s2Size,
|
||||
actualS2Len, dtypeSize, alignedS2Size);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename INPUT_T, bool hasPse>
|
||||
__aicore__ inline void DataCopyInAlign8(LocalTensor<INPUT_T> &dstTensor, GlobalTensor<INPUT_T> &srcTensor, int64_t offset,
|
||||
int64_t s1Size, int64_t s2Size, int64_t actualS2Len)
|
||||
{
|
||||
if constexpr (hasPse == true) {
|
||||
int32_t dtypeSize = sizeof(INPUT_T);
|
||||
if (dtypeSize == 0){
|
||||
return;
|
||||
}
|
||||
int32_t alignedS2Size = CeilDiv(s2Size, 32 / dtypeSize) * (32 / dtypeSize);
|
||||
DataCopyInCommon<INPUT_T, hasPse>(dstTensor, srcTensor, offset, s1Size, s2Size,
|
||||
actualS2Len, dtypeSize, alignedS2Size);
|
||||
}
|
||||
}
|
||||
|
||||
/*
|
||||
dst = BroadcastAdd(src0, src1)
|
||||
src0 shape: (s1, s2)
|
||||
src1 shape: (1, s2)
|
||||
dst shape: (s1, s2)
|
||||
*/
|
||||
template <typename T, bool hasPse>
|
||||
__aicore__ inline void BroadcastAdd(const LocalTensor<T> &src0Tensor, const LocalTensor<T> &src1Tensor,
|
||||
int64_t src0Offset, int32_t src1Size, int32_t repeatTimes)
|
||||
{
|
||||
if constexpr (hasPse == true) {
|
||||
/* Total data number of single step should be smaller than 256bytes.
|
||||
* If larger, we need to do add multiple times. */
|
||||
int32_t innerLoop = src1Size / repeatMaxSize; // s2轴整块计算次数
|
||||
int32_t innerRemain = src1Size % repeatMaxSize; // s2轴尾块计算量
|
||||
BinaryRepeatParams binaryRepeatParams;
|
||||
binaryRepeatParams.src0BlkStride = 1;
|
||||
binaryRepeatParams.src0RepStride = src1Size / blockSize;
|
||||
binaryRepeatParams.src1BlkStride = 1;
|
||||
binaryRepeatParams.src1RepStride = 0;
|
||||
binaryRepeatParams.dstRepStride = binaryRepeatParams.src0RepStride;
|
||||
binaryRepeatParams.blockNumber = binaryRepeatParams.src0RepStride;
|
||||
|
||||
for (int32_t j = 0; j < innerLoop; j++) {
|
||||
auto innerOffset = j * repeatMaxSize;
|
||||
auto ubOffset = src0Offset + innerOffset;
|
||||
Add(src0Tensor[ubOffset], src0Tensor[ubOffset], src1Tensor[innerOffset], repeatMaxSize, repeatTimes,
|
||||
binaryRepeatParams);
|
||||
}
|
||||
if (innerRemain > 0) {
|
||||
auto innerOffset = innerLoop * repeatMaxSize;
|
||||
auto ubOffset = src0Offset + innerOffset;
|
||||
Add(src0Tensor[ubOffset], src0Tensor[ubOffset], src1Tensor[innerOffset], innerRemain, repeatTimes,
|
||||
binaryRepeatParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, bool hasPse>
|
||||
__aicore__ inline void PseBroadcastAdd(int32_t s1Size, int32_t s2Size, int32_t computeSize, const LocalTensor<T> &pseUb,
|
||||
const LocalTensor<T> &dstTensor, uint32_t pseShapeType)
|
||||
{
|
||||
if constexpr (hasPse == true) {
|
||||
if (pseShapeType == pseS1S2 || pseShapeType == pseSlopeBn || pseShapeType == pseSlopeN) {
|
||||
Add(dstTensor, dstTensor, pseUb, computeSize);
|
||||
} else {
|
||||
/* Total repeated times should be <= repeatMaxTimes. If larger,
|
||||
* we need to do multiple inner loops. */
|
||||
int32_t s1OuterLoop = s1Size / repeatMaxTimes;
|
||||
int32_t s1OuterRemain = s1Size % repeatMaxTimes;
|
||||
for (int32_t s1OuterIdx = 0; s1OuterIdx < s1OuterLoop; s1OuterIdx++) {
|
||||
int32_t s1OuterOffset = s1OuterIdx * repeatMaxTimes * s2Size;
|
||||
BroadcastAdd<T, hasPse>(dstTensor, pseUb, s1OuterOffset, s2Size, repeatMaxTimes);
|
||||
}
|
||||
if (s1OuterRemain > 0) {
|
||||
int32_t s1OuterOffset = s1OuterLoop * repeatMaxTimes * s2Size;
|
||||
BroadcastAdd<T, hasPse>(dstTensor, pseUb, s1OuterOffset, s2Size, s1OuterRemain);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
template <bool hasPse> __aicore__ inline int64_t PseComputeOffset(PseInfo &pseInfo)
|
||||
{
|
||||
if constexpr (hasPse == true) {
|
||||
int64_t bOffset = 0;
|
||||
int64_t n2Offset = 0;
|
||||
int64_t s1Offset = 0;
|
||||
int64_t s2Offset = pseInfo.s2StartIdx + pseInfo.s2LoopCount * pseInfo.s2BaseNratioSize;
|
||||
int64_t gOffset = 0;
|
||||
if (pseInfo.pseShapeType == pseS1S2) {
|
||||
// b, n2, g, s1, s2
|
||||
bOffset = pseInfo.bSSOffset * pseInfo.n2G;
|
||||
n2Offset = pseInfo.n2oIdx * pseInfo.gSize * pseInfo.s1Size * pseInfo.s2Size;
|
||||
gOffset = pseInfo.goIdx * pseInfo.s1Size * pseInfo.s2Size;
|
||||
s1Offset = (pseInfo.s1oIdx * pseInfo.s1BaseSize + pseInfo.vecCoreOffset +
|
||||
pseInfo.loopIdx * pseInfo.vec1S1BaseSize) * pseInfo.s2Size;
|
||||
} else if (pseInfo.pseShapeType == pse1S2) {
|
||||
// b, n2, g, 1, s2
|
||||
bOffset = pseInfo.s2SizeAcc * pseInfo.n2G;
|
||||
n2Offset = pseInfo.n2oIdx * pseInfo.gSize * pseInfo.s2Size;
|
||||
gOffset = pseInfo.goIdx * pseInfo.s2Size;
|
||||
}
|
||||
if (pseInfo.pseBSize == 1) {
|
||||
bOffset = 0;
|
||||
}
|
||||
return bOffset + n2Offset + gOffset + s1Offset + s2Offset;
|
||||
} else {
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
|
||||
template <LayOutTypeEnum layOutType, bool hasPse> __aicore__ inline int64_t PseAlibiComputeOffset(PseInfo &pseInfo)
|
||||
{
|
||||
if constexpr (hasPse == true) {
|
||||
int64_t bOffset = (pseInfo.boIdx % pseInfo.pseBSize) * pseInfo.n2G * pseInfo.pseS2Size * pseInfo.pseS1Size;
|
||||
int64_t n2Offset = pseInfo.n2oIdx * pseInfo.gSize * pseInfo.pseS2Size * pseInfo.pseS1Size;
|
||||
int64_t gOffset = pseInfo.goIdx * pseInfo.pseS2Size * pseInfo.pseS1Size;
|
||||
int64_t row = pseInfo.s1oIdx * pseInfo.s1BaseSize + pseInfo.vecCoreOffset +
|
||||
pseInfo.loopIdx * pseInfo.vec1S1BaseSize;
|
||||
int64_t column = pseInfo.s2StartIdx + pseInfo.s2LoopCount * pseInfo.s2BaseNratioSize;
|
||||
int64_t m = 0;
|
||||
int64_t k = 0;
|
||||
if constexpr (layOutType != LayOutTypeEnum::LAYOUT_TND) {
|
||||
int64_t threshold = pseInfo.s1Size - pseInfo.pseS1Size;
|
||||
if (row >= threshold) {
|
||||
m = row - threshold;
|
||||
k = column;
|
||||
} else {
|
||||
m = row % pseInfo.pseS1Size;
|
||||
k = pseInfo.pseS2Size - (row - column) - (pseInfo.pseS1Size - m);
|
||||
}
|
||||
} else {
|
||||
int64_t threshold = pseInfo.pseS2Size - pseInfo.pseS1Size;
|
||||
int64_t posVal = row - column - threshold;
|
||||
if (threshold >= 0) {
|
||||
if (posVal >= 0) {
|
||||
m = posVal;
|
||||
k = 0;
|
||||
} else {
|
||||
m = 0;
|
||||
k = -posVal;
|
||||
}
|
||||
} else {
|
||||
m = posVal;
|
||||
k = 0;
|
||||
}
|
||||
}
|
||||
int64_t s1Offset = m * pseInfo.pseS2Size;
|
||||
int64_t s2Offset = k;
|
||||
pseInfo.readS2Size = Min(pseInfo.s2AlignedSize, pseInfo.pseS2Size - k);
|
||||
pseInfo.pseS2ComputeSize = Align(pseInfo.readS2Size);
|
||||
|
||||
return bOffset + n2Offset + gOffset + s1Offset + s2Offset;
|
||||
} else {
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
|
||||
template <bool hasPse> __aicore__ inline bool NeedPseAlibiCompute(PseInfo &pseInfo)
|
||||
{
|
||||
if constexpr (hasPse == true) {
|
||||
// Alibi编码只计算下三角
|
||||
if (pseInfo.s1oIdx * pseInfo.s1BaseSize + pseInfo.vecCoreOffset +
|
||||
(pseInfo.loopIdx + 1) * pseInfo.vec1S1BaseSize <=
|
||||
pseInfo.s2StartIdx + pseInfo.s2LoopCount * pseInfo.s2BaseNratioSize) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
} else {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename INPUT_T, typename T, LayOutTypeEnum layOutType, bool hasPse>
|
||||
__aicore__ inline void PseAlibiCopyIn(LocalTensor<T> &dstTensor, LocalTensor<INPUT_T> &tmpTensor,
|
||||
GlobalTensor<INPUT_T> &srcTensor, PseInfo &pseInfo, int64_t alignedSize = 16)
|
||||
{
|
||||
if constexpr (hasPse == true) {
|
||||
if (!NeedPseAlibiCompute<hasPse>(pseInfo)) {
|
||||
return;
|
||||
}
|
||||
int64_t offset = PseAlibiComputeOffset<layOutType, hasPse>(pseInfo);
|
||||
if constexpr (IsSameType<INPUT_T, T>::value) {
|
||||
if (!pseInfo.align8){
|
||||
DataCopyIn<INPUT_T, hasPse>(dstTensor, srcTensor, offset, pseInfo.vec1S1RealSize, pseInfo.readS2Size,
|
||||
pseInfo.pseS2Size, alignedSize);
|
||||
} else {
|
||||
DataCopyInAlign8<INPUT_T, hasPse>(dstTensor, srcTensor, offset, pseInfo.vec1S1RealSize,
|
||||
pseInfo.readS2Size, pseInfo.pseS2Size);
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
DataCopyIn<INPUT_T, hasPse>(tmpTensor, srcTensor, offset, pseInfo.vec1S1RealSize, pseInfo.readS2Size,
|
||||
pseInfo.pseS2Size, alignedSize);
|
||||
if (pseInfo.needCast) {
|
||||
event_t eventIdMte2ToV = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V));
|
||||
SetFlag<HardEvent::MTE2_V>(eventIdMte2ToV);
|
||||
WaitFlag<HardEvent::MTE2_V>(eventIdMte2ToV);
|
||||
Cast(dstTensor, tmpTensor, RoundMode::CAST_NONE, pseInfo.vec1S1RealSize * pseInfo.pseS2ComputeSize);
|
||||
}
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, bool hasPse>
|
||||
__aicore__ inline void PseSlopeCopyIn(LocalTensor<T> &dstTensor, LocalTensor<half> &helpTensor,
|
||||
__gm__ uint8_t *pseSlope, GlobalTensor<half> &alibiGm, PseInfo &pseInfo,
|
||||
int64_t alignedSize = 16) {
|
||||
if constexpr (hasPse == true) {
|
||||
int64_t bOffset = 0;
|
||||
int64_t n2Offset = pseInfo.n2oIdx * pseInfo.gSize;
|
||||
int64_t gOffset = pseInfo.goIdx;
|
||||
|
||||
if (pseInfo.pseShapeType == pseSlopeBn) {
|
||||
bOffset = pseInfo.boIdx * pseInfo.n2G;
|
||||
}
|
||||
int64_t offset = bOffset + n2Offset + gOffset;
|
||||
|
||||
DataCopyIn<half, hasPse>(helpTensor, alibiGm, 0, pseInfo.vec1S1RealSize,
|
||||
pseInfo.s2RealSize, pseInfo.pseAlibiBaseS2, alignedSize);
|
||||
event_t eventIdMte2ToV = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V));
|
||||
SetFlag<HardEvent::MTE2_V>(eventIdMte2ToV);
|
||||
WaitFlag<HardEvent::MTE2_V>(eventIdMte2ToV);
|
||||
|
||||
if (pseInfo.needCast) {
|
||||
int64_t computeSize = pseInfo.vec1S1RealSize * pseInfo.s2AlignedSize;
|
||||
Cast(dstTensor, helpTensor, RoundMode::CAST_NONE, computeSize);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
int64_t s1Offset = pseInfo.s1oIdx * pseInfo.s1BaseSize + pseInfo.vecCoreOffset +
|
||||
pseInfo.loopIdx * pseInfo.vec1S1BaseSize;
|
||||
int64_t s2Offset = pseInfo.s2StartIdx + pseInfo.s2LoopCount * pseInfo.s2BaseNratioSize;
|
||||
|
||||
float posShift = float(s2Offset + pseInfo.kvStartIdx - s1Offset - pseInfo.qStartIdx);
|
||||
|
||||
Adds(dstTensor, dstTensor, posShift, computeSize);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Abs(dstTensor, dstTensor, computeSize);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
float slopes = ((__gm__ T *)pseSlope)[offset] * -1;
|
||||
if (pseInfo.pseType == (uint32_t)PseTypeEnum::PSE_INNER_MUL_ADD_SQRT_TYPE) {
|
||||
Sqrt(dstTensor, dstTensor, computeSize);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
Muls(dstTensor, dstTensor, slopes, computeSize);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, bool hasPse>
|
||||
__aicore__ inline void PseSlopeCast(LocalTensor<T> &dstTensor, LocalTensor<half> &helpTensor,
|
||||
__gm__ uint8_t *pseSlope, PseInfo &pseInfo) {
|
||||
if constexpr (hasPse == true) {
|
||||
int64_t bOffset = 0;
|
||||
int64_t n2Offset = pseInfo.n2oIdx * pseInfo.gSize;
|
||||
int64_t gOffset = pseInfo.goIdx;
|
||||
|
||||
if (pseInfo.pseShapeType == pseSlopeBn) {
|
||||
bOffset = pseInfo.boIdx * pseInfo.n2G;
|
||||
}
|
||||
int64_t offset = bOffset + n2Offset + gOffset;
|
||||
int64_t computeSize = pseInfo.vec1S1RealSize * pseInfo.s2AlignedSize;
|
||||
Cast(dstTensor, helpTensor, RoundMode::CAST_NONE, computeSize);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
int64_t s1Offset = pseInfo.s1oIdx * pseInfo.s1BaseSize + pseInfo.vecCoreOffset +
|
||||
pseInfo.loopIdx * pseInfo.vec1S1BaseSize;
|
||||
int64_t s2Offset = pseInfo.s2StartIdx + pseInfo.s2LoopCount * pseInfo.s2BaseNratioSize;
|
||||
|
||||
float posShift = float(s2Offset + pseInfo.kvStartIdx - s1Offset - pseInfo.qStartIdx);
|
||||
|
||||
Adds(dstTensor, dstTensor, posShift, computeSize);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Abs(dstTensor, dstTensor, computeSize);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
float slopes = ((__gm__ T *)pseSlope)[offset] * -1;
|
||||
if (pseInfo.pseType == (uint32_t)PseTypeEnum::PSE_INNER_MUL_ADD_SQRT_TYPE) {
|
||||
Sqrt(dstTensor, dstTensor, computeSize);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
Muls(dstTensor, dstTensor, slopes, computeSize);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
}
|
||||
|
||||
template <typename INPUT_T, typename T, LayOutTypeEnum layOutType, bool hasPse>
|
||||
__aicore__ inline void PseCopyIn(LocalTensor<T> &dstTensor, LocalTensor<INPUT_T> &tmpTensor,
|
||||
GlobalTensor<INPUT_T> &srcTensor, PseInfo &pseInfo, int64_t alignedSize = 16)
|
||||
{
|
||||
if constexpr (hasPse == true) {
|
||||
if (pseInfo.pseEncodeType == pseEncodeALibiS2Full) {
|
||||
return PseAlibiCopyIn<INPUT_T, T, layOutType, hasPse>(dstTensor, tmpTensor, srcTensor, pseInfo, alignedSize);
|
||||
}
|
||||
int64_t offset = PseComputeOffset<hasPse>(pseInfo);
|
||||
int64_t s1Size = pseInfo.pseShapeType == pse1S2 ? (pseInfo.blockCount == 0 ? 1 : pseInfo.blockCount) :
|
||||
pseInfo.vec1S1RealSize;
|
||||
|
||||
if constexpr (IsSameType<INPUT_T, T>::value) {
|
||||
if (!pseInfo.align8){
|
||||
DataCopyIn<INPUT_T, hasPse>(dstTensor, srcTensor, offset, s1Size, pseInfo.s2RealSize,
|
||||
pseInfo.s2Size, alignedSize);
|
||||
} else {
|
||||
DataCopyInAlign8<INPUT_T, hasPse>(dstTensor, srcTensor, offset, s1Size, pseInfo.s2RealSize, pseInfo.s2Size);
|
||||
}
|
||||
return;
|
||||
}
|
||||
DataCopyIn<INPUT_T, hasPse>(tmpTensor, srcTensor, offset, s1Size, pseInfo.s2RealSize, pseInfo.s2Size,
|
||||
alignedSize);
|
||||
if (pseInfo.needCast) {
|
||||
event_t eventIdMte2ToV = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V));
|
||||
SetFlag<HardEvent::MTE2_V>(eventIdMte2ToV);
|
||||
WaitFlag<HardEvent::MTE2_V>(eventIdMte2ToV);
|
||||
Cast(dstTensor, tmpTensor, RoundMode::CAST_NONE, s1Size * pseInfo.s2AlignedSize);
|
||||
}
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, bool hasPse>
|
||||
__aicore__ inline void PseAlibiCompute(LocalTensor<T> &dstTensor, LocalTensor<T> &pseTensor, PseInfo &pseInfo)
|
||||
{
|
||||
if constexpr (hasPse == true) {
|
||||
if (!NeedPseAlibiCompute<hasPse>(pseInfo)) {
|
||||
return;
|
||||
}
|
||||
Add(dstTensor, dstTensor, pseTensor, pseInfo.vec1S1RealSize * pseInfo.pseS2ComputeSize);
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, bool hasPse>
|
||||
__aicore__ inline void PseCompute(LocalTensor<T> &dstTensor, LocalTensor<T> &pseTensor, PseInfo &pseInfo)
|
||||
{
|
||||
if constexpr (hasPse == true) {
|
||||
if (pseInfo.pseEncodeType == pseEncodeALibiS2Full) {
|
||||
return PseAlibiCompute<T, hasPse>(dstTensor, pseTensor, pseInfo);
|
||||
}
|
||||
int64_t computeSize = (pseInfo.pseShapeType == pseS1S2 || pseInfo.pseShapeType == pseSlopeBn ||
|
||||
pseInfo.pseShapeType == pseSlopeN)
|
||||
? pseInfo.vec1S1RealSize * pseInfo.s2AlignedSize
|
||||
: pseInfo.s2AlignedSize;
|
||||
PseBroadcastAdd<T, hasPse>(pseInfo.vec1S1RealSize, pseInfo.s2AlignedSize, computeSize, pseTensor,
|
||||
dstTensor, pseInfo.pseShapeType);
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
template <bool hasPse>
|
||||
__aicore__ inline void PseInnerAlibiCreate(GlobalTensor<half> &dstTensor, LocalTensor<half> &helpTensor, PseInfo &pseInfo) {
|
||||
if constexpr (hasPse == true) {
|
||||
if (pseInfo.pseType != (uint32_t)PseTypeEnum::PSE_INNER_MUL_ADD_TYPE && pseInfo.pseType != (uint32_t)PseTypeEnum::PSE_INNER_MUL_ADD_SQRT_TYPE) {
|
||||
return;
|
||||
}
|
||||
event_t eventIdMte3ToV = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_V));
|
||||
event_t eventIdMte3ToS = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_S));
|
||||
event_t eventIdVToMte3 = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3));
|
||||
float tmpValue = -1.0;
|
||||
|
||||
for (int64_t i = 0; i < pseInfo.pseAlibiBaseS1; i++) {
|
||||
CreateVecIndex(helpTensor, (half)(i * tmpValue), pseInfo.pseAlibiBaseS2);
|
||||
SetFlag<HardEvent::V_MTE3>(eventIdVToMte3);
|
||||
WaitFlag<HardEvent::V_MTE3>(eventIdVToMte3);
|
||||
DataCopy(dstTensor[i * pseInfo.pseAlibiBaseS2], helpTensor, pseInfo.pseAlibiBaseS2);
|
||||
SetFlag<HardEvent::MTE3_V>(eventIdMte3ToV);
|
||||
WaitFlag<HardEvent::MTE3_V>(eventIdMte3ToV);
|
||||
SetFlag<HardEvent::MTE3_S>(eventIdMte3ToS);
|
||||
WaitFlag<HardEvent::MTE3_S>(eventIdMte3ToS);
|
||||
}
|
||||
}
|
||||
}
|
||||
#endif
|
||||
426
csrc/utils/inc/kernel/sync_collectives.h
Normal file
426
csrc/utils/inc/kernel/sync_collectives.h
Normal file
@@ -0,0 +1,426 @@
|
||||
#ifndef SYNC_COLLECTIVES_H
|
||||
#define SYNC_COLLECTIVES_H
|
||||
|
||||
#include "comm_args.h"
|
||||
|
||||
using namespace AscendC;
|
||||
using namespace Moe;
|
||||
|
||||
// Synchronization flag occupies length
|
||||
constexpr int64_t FLAG_UNIT_INT_NUM = 4;
|
||||
// Memory size occupied by each synchronization unit (Bytes)
|
||||
constexpr int64_t SYNC_UNIT_SIZE = FLAG_UNIT_INT_NUM * sizeof(int64_t);
|
||||
// High-order offset when using magic as a comparison value
|
||||
constexpr int64_t MAGIC_OFFSET = 32;
|
||||
constexpr int64_t MAGIC_MASK = ~((1LL << MAGIC_OFFSET) - 1);
|
||||
|
||||
class SyncCollectives {
|
||||
public:
|
||||
__aicore__ inline SyncCollectives() {}
|
||||
|
||||
__aicore__ inline void Init(int rank, int rankSize, GM_ADDR *shareAddrs, TBuf<QuePosition::VECCALC> &tBuf)
|
||||
{
|
||||
this->rank = rank;
|
||||
this->rankSize = rankSize;
|
||||
this->shareAddrs = shareAddrs;
|
||||
this->blockIdx = GetBlockIdx();
|
||||
this->blockNum = GetBlockNum();
|
||||
// Length of a single indicator segment
|
||||
segmentCount = GetBlockNum() * FLAG_UNIT_INT_NUM;
|
||||
// Initialize the intra-card/inter-card synchronization address corresponding to the current core.
|
||||
localSyncAddr = (__gm__ int64_t*)(shareAddrs[rank]);
|
||||
basicSyncAddr = (__gm__ int64_t*)(shareAddrs[rank]) + GetBlockIdx() * FLAG_UNIT_INT_NUM;
|
||||
blockOuterSyncAddr = (__gm__ int64_t*)(shareAddrs[rank]) + segmentCount + GetBlockIdx() * FLAG_UNIT_INT_NUM;
|
||||
this->tBuf = tBuf;
|
||||
}
|
||||
|
||||
__aicore__ inline void SetSyncFlag(int32_t magic, int32_t value, int32_t eventID)
|
||||
{
|
||||
int64_t v = MergeMagicWithValue(magic, value);
|
||||
SetFlag(localSyncAddr + eventID * FLAG_UNIT_INT_NUM, v);
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Set the flag for the specified eventID of the designated card, with the value being a combination of magic and value.
|
||||
* @param magic The operator batch, which will be combined into the high 32 bits of the flag value to be set.
|
||||
* @param value The specific value to be set, which will be the low 32 bits of the flag value to be set.
|
||||
* @param eventID Physically, it is an offset from the shared memory base address (requires scaling, not an absolute value).
|
||||
* @param rank This rank is the rankId corresponding to the peerMems array in the CommArgs structure, not a global or local id.
|
||||
* (Local is not applicable in the 91093 scenario, and global is not applicable in the 910B multi-machine scenario.)
|
||||
*/
|
||||
__aicore__ inline void SetSyncFlag(int32_t magic, int32_t value, int32_t eventID, int32_t rank)
|
||||
{
|
||||
int64_t v = MergeMagicWithValue(magic, value);
|
||||
SetFlag((__gm__ int64_t*)(shareAddrs[rank]) + eventID * FLAG_UNIT_INT_NUM, v);
|
||||
}
|
||||
|
||||
__aicore__ inline int32_t CalEventIdByMulBlockNum(int32_t blockMultiplier, int32_t targetCoreId)
|
||||
{
|
||||
return (blockMultiplier * blockNum) + targetCoreId;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Wait for the flag of the specified eventID on the specified card to become a value
|
||||
* composed of the combination of magic and value.
|
||||
* @param magic The operator batch, which will be combined into the high 32 bits of the flag
|
||||
* value to be wait.
|
||||
* @param value The specific value to be wait, which will be the low 32 bits of the flag
|
||||
* value to be wait.
|
||||
* @param eventID Physically, it is an offset from the shared memory base address (requires
|
||||
* scaling, not an absolute value).
|
||||
* @param rank This rank is the rankId corresponding to the peerMems array in the CommArgs
|
||||
* structure, not a global or local id. (Local is not applicable in the 91093
|
||||
* scenario, and global is not applicable in the 910B multi-machine scenario.)
|
||||
*/
|
||||
__aicore__ inline void WaitSyncFlag(int32_t magic, int32_t value, int32_t eventID, int32_t rank)
|
||||
{
|
||||
int64_t v = MergeMagicWithValue(magic, value);
|
||||
WaitOneRankPartFlag((__gm__ int64_t*)(shareAddrs[rank]) + eventID * FLAG_UNIT_INT_NUM, 1, v);
|
||||
}
|
||||
|
||||
__aicore__ inline void WaitSyncFlag(int32_t magic, int32_t value, int32_t eventID)
|
||||
{
|
||||
int64_t v = MergeMagicWithValue(magic, value);
|
||||
WaitOneRankPartFlag((__gm__ int64_t*)(shareAddrs[this->rank]) + eventID * FLAG_UNIT_INT_NUM, 1, v);
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Wait for the flags starting from the specified eventID on the specified card to become
|
||||
* a value composed of the combination of magic and value.<br>
|
||||
* Note: [eventID, eventID + flagNum)
|
||||
*/
|
||||
__aicore__ inline void WaitSyncFlag(int32_t magic, int32_t value, int32_t eventID, int32_t rank, int64_t flagNum)
|
||||
{
|
||||
int64_t v = MergeMagicWithValue(magic, value);
|
||||
WaitOneRankPartFlag((__gm__ int64_t*)(shareAddrs[rank]) + eventID * FLAG_UNIT_INT_NUM, flagNum, v);
|
||||
}
|
||||
|
||||
// Set inner-card synchronization flag (memory A)
|
||||
__aicore__ inline void SetInnerFlag(int32_t magic, int32_t eventID)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, eventID);
|
||||
SetFlag(basicSyncAddr, value);
|
||||
}
|
||||
|
||||
__aicore__ inline void SetInnerFlag(int32_t magic, int32_t eventID, int64_t setRank, int64_t setBlock)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, eventID);
|
||||
SetFlag((__gm__ int64_t*)(shareAddrs[setRank]) + setBlock * FLAG_UNIT_INT_NUM, value);
|
||||
}
|
||||
|
||||
// Wait for a single inner-card synchronization flag (memory A)
|
||||
__aicore__ inline void WaitInnerFlag(int32_t magic, int32_t eventID, int64_t waitRank, int64_t waitBlock)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, eventID);
|
||||
WaitOneRankPartFlag((__gm__ int64_t*)(shareAddrs[waitRank]) + waitBlock * FLAG_UNIT_INT_NUM, 1, value);
|
||||
}
|
||||
|
||||
// Wait for all inner-card synchronization flags within the entire rank (memory A)
|
||||
__aicore__ inline void WaitRankInnerFlag(int32_t magic, int32_t eventID, int64_t waitRank)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, eventID);
|
||||
WaitOneRankAllFlag((__gm__ int64_t*)(shareAddrs[waitRank]), value);
|
||||
}
|
||||
|
||||
// Check all inner-card synchronization flags within the entire rank (memory A)
|
||||
__aicore__ inline bool CheckRankInnerFlag(int32_t magic, int32_t eventID, int64_t waitRank)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, eventID);
|
||||
return CheckOneRankAllFlag((__gm__ int64_t*)(shareAddrs[waitRank]), value);
|
||||
}
|
||||
|
||||
// Set inter-card synchronization flag (memory B)
|
||||
__aicore__ inline void SetOuterFlag(int32_t magic, int32_t eventID)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, eventID);
|
||||
SetFlag(blockOuterSyncAddr, value);
|
||||
}
|
||||
|
||||
__aicore__ inline void SetOuterFlag(int32_t magic, int32_t eventID, int64_t setRank, int64_t setBlock)
|
||||
{
|
||||
__gm__ int64_t* flagAddr = GetOuterFlagAddr(setRank, setBlock);
|
||||
int64_t value = MergeMagicWithValue(magic, eventID);
|
||||
SetFlag(flagAddr, value);
|
||||
}
|
||||
|
||||
// Wait for a single inter-card synchronization flag (memory B)
|
||||
__aicore__ inline void WaitOuterFlag(int32_t magic, int32_t eventID, int64_t waitRank, int64_t waitBlock)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, eventID);
|
||||
__gm__ int64_t* flagAddr = GetOuterFlagAddr(waitRank, waitBlock);
|
||||
WaitOneRankPartFlag(flagAddr, 1, value);
|
||||
}
|
||||
|
||||
// Wait for all inter-card synchronization flags within the entire rank (memory B)
|
||||
__aicore__ inline void WaitOneRankOuterFlag(int32_t magic, int32_t eventID, int64_t rank)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, eventID);
|
||||
__gm__ int64_t* flagAddr;
|
||||
flagAddr = GetOuterFlagAddr(rank, 0);
|
||||
WaitOneRankPartFlag(flagAddr, blockNum, value);
|
||||
}
|
||||
|
||||
// Wait for flagNum inter-card synchronization flags starting from startBlock for all ranks (memory B)
|
||||
__aicore__ inline void WaitAllRankPartOuterFlag(int32_t magic, int32_t eventID, int64_t startBlock, int64_t flagNum)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, eventID);
|
||||
__gm__ int64_t* flagAddr;
|
||||
int waitRank;
|
||||
for (auto r = 0; r < rankSize; ++r) {
|
||||
waitRank = (rank + r) % rankSize; // Offset reading of rank flags to prevent performance impact from concurrent copying by multiple cores
|
||||
flagAddr = GetOuterFlagAddr(waitRank, startBlock);
|
||||
WaitOneRankPartFlag(flagAddr, flagNum, value);
|
||||
}
|
||||
}
|
||||
|
||||
// Check flagNum inter-card synchronization flags starting from startBlock for all ranks (memory B)
|
||||
__aicore__ inline bool CheckAllRankPartOuterFlag(int32_t magic, int32_t eventID, int64_t startBlock,
|
||||
int64_t flagNum)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, eventID);
|
||||
__gm__ int64_t* flagAddr;
|
||||
int waitRank;
|
||||
for (auto r = 0; r < rankSize; ++r) {
|
||||
waitRank = (rank + r) % rankSize; // Offset reading of rank flags to prevent performance impact from concurrent copying by multiple cores
|
||||
flagAddr = GetOuterFlagAddr(waitRank, startBlock);
|
||||
if (!CheckOneRankPartFlag(flagAddr, flagNum, value)) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
// Wait for all inter-card synchronization flags for all ranks, full rank synchronization (memory B)
|
||||
__aicore__ inline void WaitAllRankOuterFlag(int32_t magic, int32_t eventID)
|
||||
{
|
||||
WaitAllRankPartOuterFlag(magic, eventID, 0, blockNum);
|
||||
}
|
||||
|
||||
// Check all inter-card synchronization flags for all ranks, full rank synchronization (memory B)
|
||||
__aicore__ inline bool CheckAllRankOuterFlag(int32_t magic, int32_t eventID)
|
||||
{
|
||||
return CheckAllRankPartOuterFlag(magic, eventID, 0, blockNum);
|
||||
}
|
||||
|
||||
// Low-level interface, set synchronization flag
|
||||
__aicore__ inline void SetFlag(__gm__ int64_t* setAddr, int64_t setValue)
|
||||
{
|
||||
AscendC::SetFlag<HardEvent::MTE3_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<HardEvent::MTE3_S>(EVENT_ID0);
|
||||
AscendC::SetFlag<HardEvent::MTE2_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<HardEvent::MTE2_S>(EVENT_ID0);
|
||||
GlobalTensor<int64_t> globalSet;
|
||||
globalSet.SetGlobalBuffer(setAddr, FLAG_UNIT_INT_NUM);
|
||||
LocalTensor<int64_t> localSet = tBuf.GetWithOffset<int64_t>(1, 0);
|
||||
localSet.SetValue(0, setValue);
|
||||
|
||||
// Copy global synchronization flag to local
|
||||
AscendC::SetFlag<HardEvent::S_MTE3>(EVENT_ID0);
|
||||
AscendC::WaitFlag<HardEvent::S_MTE3>(EVENT_ID0); // Wait for SetValue to complete
|
||||
DataCopy(globalSet, localSet, FLAG_UNIT_INT_NUM);
|
||||
AscendC::SetFlag<HardEvent::MTE3_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<HardEvent::MTE3_S>(EVENT_ID0); // Wait for UB->GM to complete
|
||||
}
|
||||
|
||||
// Low-level interface, wait for synchronization flag
|
||||
__aicore__ inline void WaitFlag(__gm__ int64_t* waitAddr, int64_t waitValue)
|
||||
{
|
||||
WaitOneRankPartFlag(waitAddr, 1, waitValue);
|
||||
}
|
||||
|
||||
// Read a flag, return an immediate number
|
||||
__aicore__ inline int64_t GetFlag(__gm__ int64_t* waitAddr)
|
||||
{
|
||||
GlobalTensor<int64_t> globalWait;
|
||||
globalWait.SetGlobalBuffer(waitAddr, FLAG_UNIT_INT_NUM);
|
||||
LocalTensor<int64_t> localWait = tBuf.GetWithOffset<int64_t>(1, 0);
|
||||
// Copy global to local
|
||||
DataCopy(localWait, globalWait, FLAG_UNIT_INT_NUM);
|
||||
AscendC::SetFlag<HardEvent::MTE2_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<HardEvent::MTE2_S>(EVENT_ID0); // Wait for GM->UB
|
||||
|
||||
int64_t res = localWait.GetValue(0);
|
||||
return res;
|
||||
}
|
||||
|
||||
// Get multiple consecutive synchronization flags within a single card
|
||||
__aicore__ inline void WaitOneRankPartOuterFlag(int32_t magic, int32_t eventID, int64_t waitRank,
|
||||
int64_t startBlock, int64_t flagNum)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, eventID);
|
||||
__gm__ int64_t* flagAddr;
|
||||
flagAddr = GetOuterFlagAddr(waitRank, startBlock);
|
||||
WaitOneRankPartFlag(flagAddr, flagNum, value);
|
||||
}
|
||||
|
||||
// Get synchronization flag within a single card (memory A)
|
||||
__aicore__ inline int64_t GetInnerFlag(int64_t waitRank, int64_t waitBlock)
|
||||
{
|
||||
return GetFlag((__gm__ int64_t*)(shareAddrs[waitRank]) + waitBlock * FLAG_UNIT_INT_NUM);
|
||||
}
|
||||
|
||||
__aicore__ inline int64_t GetOuterFlag(int64_t waitRank, int64_t waitBlock)
|
||||
{
|
||||
return GetFlag((__gm__ int64_t*)(shareAddrs[waitRank]) + segmentCount + waitBlock * FLAG_UNIT_INT_NUM);
|
||||
}
|
||||
|
||||
// In the rank Chunk Flag area, return success if the destRank chunk Flag value is 0, otherwise fail
|
||||
__aicore__ inline int64_t GetChunkFlag(int64_t rank, int64_t destRank, int64_t magic, int64_t timeout)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, 0);
|
||||
int64_t status = GetChunkFlagValue((__gm__ int64_t*)(shareAddrs[rank]) +
|
||||
IPC_CHUNK_FLAG + destRank * FLAG_UNIT_INT_NUM, value, timeout);
|
||||
return status;
|
||||
}
|
||||
|
||||
// Set the destRank chunk Flag value in the rank Chunk Flag area to value
|
||||
__aicore__ inline void SetChunkFlag(int64_t rank, int64_t destRank, int64_t magic, int64_t eventId)
|
||||
{
|
||||
int64_t value = MergeMagicWithValue(magic, eventId);
|
||||
SetFlag((__gm__ int64_t*)(shareAddrs[rank]) + IPC_CHUNK_FLAG + destRank * FLAG_UNIT_INT_NUM, value);
|
||||
}
|
||||
|
||||
__aicore__ inline int64_t GetChunkRecvLen(int64_t rank, int64_t destRank, int64_t magic, int64_t timeout)
|
||||
{
|
||||
int64_t len = GetChunkFlagValue((__gm__ int64_t*)(shareAddrs[rank]) + IPC_CHUNK_FLAG +
|
||||
destRank * FLAG_UNIT_INT_NUM, 0, timeout, true, magic);
|
||||
return len;
|
||||
}
|
||||
|
||||
private:
|
||||
__aicore__ inline int64_t MergeMagicWithValue(int32_t magic, int32_t value)
|
||||
{
|
||||
// Merge magic as the high bits and eventID as the low bits into a value for comparison
|
||||
return (static_cast<int64_t>(static_cast<uint32_t>(magic)) << MAGIC_OFFSET) | static_cast<int64_t>(value);
|
||||
}
|
||||
|
||||
__aicore__ inline __gm__ int64_t* GetInnerFlagAddr(int64_t flagRank, int64_t flagBlock)
|
||||
{
|
||||
return (__gm__ int64_t*)(shareAddrs[flagRank]) + flagBlock * FLAG_UNIT_INT_NUM;
|
||||
}
|
||||
|
||||
__aicore__ inline __gm__ int64_t* GetOuterFlagAddr(int64_t flagRank, int64_t flagBlock)
|
||||
{
|
||||
return (__gm__ int64_t*)(shareAddrs[flagRank]) + segmentCount + flagBlock * FLAG_UNIT_INT_NUM;
|
||||
}
|
||||
|
||||
// Wait for a part of synchronization flags within a rank
|
||||
__aicore__ inline void WaitOneRankPartFlag(__gm__ int64_t* waitAddr, int64_t flagNum, int64_t checkValue)
|
||||
{
|
||||
GlobalTensor<int64_t> globalWait;
|
||||
globalWait.SetGlobalBuffer(waitAddr, flagNum * FLAG_UNIT_INT_NUM);
|
||||
LocalTensor<int64_t> localWait = tBuf.GetWithOffset<int64_t>(flagNum * FLAG_UNIT_INT_NUM, 0);
|
||||
bool isSync = true;
|
||||
int64_t checkedFlagNum = 0;
|
||||
do {
|
||||
// Copy global synchronization flags to local
|
||||
DataCopy(localWait, globalWait[checkedFlagNum * FLAG_UNIT_INT_NUM],
|
||||
(flagNum - checkedFlagNum) * FLAG_UNIT_INT_NUM);
|
||||
AscendC::SetFlag<HardEvent::MTE2_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<HardEvent::MTE2_S>(EVENT_ID0); // Wait for GM->UB
|
||||
|
||||
// Check if the synchronization flags are equal to checkValue
|
||||
isSync = true;
|
||||
int64_t remainToCheck = flagNum - checkedFlagNum;
|
||||
for (auto i = 0; i < remainToCheck; ++i) {
|
||||
// Continue waiting if any core has not reached the checkValue phase
|
||||
int64_t v = localWait.GetValue(i * FLAG_UNIT_INT_NUM);
|
||||
if ((v & MAGIC_MASK) != (checkValue & MAGIC_MASK) || v < checkValue) {
|
||||
isSync = false;
|
||||
checkedFlagNum += i;
|
||||
break;
|
||||
}
|
||||
}
|
||||
} while (!isSync);
|
||||
}
|
||||
|
||||
// Wait for all synchronization flags within a rank
|
||||
__aicore__ inline void WaitOneRankAllFlag(__gm__ int64_t* waitAddr, int64_t checkValue)
|
||||
{
|
||||
WaitOneRankPartFlag(waitAddr, blockNum, checkValue);
|
||||
}
|
||||
|
||||
// Check partial synchronization flags within a rank, copy only once
|
||||
__aicore__ inline bool CheckOneRankPartFlag(__gm__ int64_t* waitAddr, int64_t flagNum, int64_t checkValue)
|
||||
{
|
||||
GlobalTensor<int64_t> globalWait;
|
||||
globalWait.SetGlobalBuffer(waitAddr, flagNum * FLAG_UNIT_INT_NUM);
|
||||
LocalTensor<int64_t> localWait = tBuf.GetWithOffset<int64_t>(flagNum * FLAG_UNIT_INT_NUM, 0);
|
||||
// Copy global synchronization flags to local
|
||||
DataCopy(localWait, globalWait, flagNum * FLAG_UNIT_INT_NUM);
|
||||
AscendC::SetFlag<HardEvent::MTE2_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<HardEvent::MTE2_S>(EVENT_ID0); // Wait for GM->UB
|
||||
// Check if the synchronization flags are equal to checkValue
|
||||
bool isSync = true;
|
||||
for (auto i = 0; i < flagNum; ++i) {
|
||||
// Continue waiting if any core has not reached the checkValue phase
|
||||
int64_t v = localWait.GetValue(i * FLAG_UNIT_INT_NUM);
|
||||
if ((v & MAGIC_MASK) != (checkValue & MAGIC_MASK) || v < checkValue) {
|
||||
isSync = false;
|
||||
break;
|
||||
}
|
||||
}
|
||||
return isSync;
|
||||
}
|
||||
|
||||
__aicore__ inline int64_t GetChunkFlagValue(__gm__ int64_t* waitAddr, int64_t checkValue, int64_t timeout,
|
||||
bool checkNonZero = false, int64_t magic = 0)
|
||||
{
|
||||
GlobalTensor<int64_t> globalWait;
|
||||
globalWait.SetGlobalBuffer(waitAddr, FLAG_UNIT_INT_NUM);
|
||||
LocalTensor<int64_t> localWait = tBuf.GetWithOffset<int64_t>(FLAG_UNIT_INT_NUM, 0);
|
||||
bool isSync = true;
|
||||
|
||||
int64_t waitTimes = 0;
|
||||
int64_t v = 0;
|
||||
|
||||
do {
|
||||
// Copy global sync flag to local
|
||||
DataCopy(localWait, globalWait[0], FLAG_UNIT_INT_NUM);
|
||||
AscendC::SetFlag<HardEvent::MTE2_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<HardEvent::MTE2_S>(EVENT_ID0); // Wait for GM->UB
|
||||
|
||||
isSync = true;
|
||||
v = localWait.GetValue(0);
|
||||
if (checkNonZero) {
|
||||
// Non-zero check mode
|
||||
if (((v & MAGIC_MASK) == (static_cast<int64_t>(magic) << MAGIC_OFFSET)) && (v & 0xFFFFFFFF)) {
|
||||
return v & 0xFFFFFFFF; // Return lower 32 bits when non-zero
|
||||
}
|
||||
} else {
|
||||
// Exact value check mode
|
||||
if (v == checkValue) {
|
||||
return WAIT_SUCCESS;
|
||||
}
|
||||
}
|
||||
|
||||
isSync = false;
|
||||
waitTimes++;
|
||||
|
||||
if (timeout > INT64_MAX / MAX_WAIT_ROUND_UNIT || waitTimes >= (timeout * MAX_WAIT_ROUND_UNIT)) {
|
||||
isSync = true;
|
||||
return v; // Return the read flag value
|
||||
}
|
||||
} while (!isSync);
|
||||
|
||||
return checkNonZero ? 0 : v;
|
||||
}
|
||||
|
||||
// Check all sync flags within a rank, copy only once
|
||||
__aicore__ inline bool CheckOneRankAllFlag(__gm__ int64_t* waitAddr, int64_t checkValue)
|
||||
{
|
||||
return CheckOneRankPartFlag(waitAddr, blockNum, checkValue);
|
||||
}
|
||||
int rank;
|
||||
int rankSize;
|
||||
int blockIdx;
|
||||
int blockNum;
|
||||
GM_ADDR *shareAddrs;
|
||||
int64_t segmentCount; // Length of a single sync flag segment (count in int64_t)
|
||||
__gm__ int64_t* localSyncAddr;
|
||||
__gm__ int64_t* basicSyncAddr; // Intra-card sync flag address for the current block
|
||||
__gm__ int64_t* blockOuterSyncAddr; // Inter-card sync flag address for the current block
|
||||
TBuf<QuePosition::VECCALC> tBuf;
|
||||
};
|
||||
|
||||
#endif // SYNC_COLLECTIVES_H
|
||||
144
csrc/utils/inc/kernel/util.h
Normal file
144
csrc/utils/inc/kernel/util.h
Normal file
@@ -0,0 +1,144 @@
|
||||
/**
|
||||
* Copyright (c) 2024 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.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 util.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef FLASH_ATTENTION_UTIL_H
|
||||
#define FLASH_ATTENTION_UTIL_H
|
||||
|
||||
constexpr int32_t blockBytes = 32;
|
||||
constexpr int32_t byteBitRatio = 8;
|
||||
constexpr int64_t prefixAttenMaskDownHeight = 1024;
|
||||
constexpr static int32_t blockSize = blockBytes / 4; // 4 means sizeof(T)
|
||||
constexpr static int32_t repeatMaxBytes = 256;
|
||||
constexpr static int32_t repeatMaxTimes = 255;
|
||||
constexpr static int32_t repeatMaxSize = repeatMaxBytes / 4; // 4 means sizeof(T)
|
||||
|
||||
using AscendC::LocalTensor;
|
||||
using AscendC::GlobalTensor;
|
||||
using AscendC::DataFormat;
|
||||
using AscendC::ShapeInfo;
|
||||
using AscendC::DataCopyParams;
|
||||
using AscendC::DataCopyPadParams;
|
||||
using AscendC::BinaryRepeatParams;
|
||||
using AscendC::IsSameType;
|
||||
using AscendC::HardEvent;
|
||||
using AscendC::SetFlag;
|
||||
using AscendC::WaitFlag;
|
||||
|
||||
enum class LayOutTypeEnum { None = 0, LAYOUT_BSH = 1, LAYOUT_SBH = 2, LAYOUT_BNSD = 3, LAYOUT_TND = 4, LAYOUT_NTD_TND = 5};
|
||||
|
||||
namespace math {
|
||||
template <typename T> __aicore__ inline T Ceil(T a, T b)
|
||||
{
|
||||
if (b == 0) {
|
||||
return 0;
|
||||
}
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
template <typename T> __aicore__ inline T Align(T a, T b)
|
||||
{
|
||||
if (b == 0) {
|
||||
return 0;
|
||||
}
|
||||
return (a + b - 1) / b * b;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T1, typename T2>
|
||||
__aicore__ inline T1 CeilDiv(T1 a, T2 b)
|
||||
{
|
||||
if (b == 0) {
|
||||
return 0;
|
||||
}
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
template <typename T1, typename T2>
|
||||
__aicore__ inline T1 Max(T1 a, T2 b)
|
||||
{
|
||||
return (a > b) ? (a) : (b);
|
||||
}
|
||||
|
||||
template <typename T1, typename T2>
|
||||
__aicore__ inline T1 Min(T1 a, T2 b)
|
||||
{
|
||||
return (a > b) ? (b) : (a);
|
||||
}
|
||||
|
||||
__aicore__ inline void BoolCopyIn(LocalTensor<uint8_t> &dstTensor, GlobalTensor<uint8_t> &srcTensor,
|
||||
int64_t srcOffset, uint32_t s1Size, uint32_t s2Size, int64_t totalS2Size, int64_t alignedSize = blockBytes)
|
||||
{
|
||||
uint32_t alignedS2Size = CeilDiv(s2Size, alignedSize) * alignedSize;
|
||||
uint32_t shapeArray[] = {s1Size, alignedS2Size};
|
||||
dstTensor.SetShapeInfo(ShapeInfo(2, shapeArray, DataFormat::ND));
|
||||
dstTensor.SetSize(s1Size * alignedS2Size);
|
||||
DataCopyParams dataCopyParams;
|
||||
dataCopyParams.blockCount = s1Size;
|
||||
dataCopyParams.dstStride = 0;
|
||||
if (totalS2Size == blockBytes && alignedSize == 64) { // totalS2Size < 64 && totalS2Size % blockBytes == 0
|
||||
dataCopyParams.dstStride = 1;
|
||||
alignedSize = blockBytes;
|
||||
alignedS2Size = CeilDiv(s2Size, blockBytes) * blockBytes;
|
||||
}
|
||||
if (totalS2Size % alignedSize == 0) {
|
||||
dataCopyParams.blockLen = alignedS2Size / blockBytes;
|
||||
dataCopyParams.srcStride = (totalS2Size - alignedS2Size) / blockBytes;
|
||||
DataCopy(dstTensor, srcTensor[srcOffset], dataCopyParams);
|
||||
} else {
|
||||
dataCopyParams.blockLen = s2Size;
|
||||
dataCopyParams.srcStride = totalS2Size - s2Size;
|
||||
DataCopyPadParams dataCopyPadParams;
|
||||
dataCopyPadParams.isPad = true;
|
||||
dataCopyPadParams.rightPadding = Min(alignedS2Size - s2Size, blockBytes);
|
||||
dataCopyPadParams.paddingValue = 1;
|
||||
DataCopyPad(dstTensor, srcTensor[srcOffset], dataCopyParams, dataCopyPadParams);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void Bit2Int8CopyIn(LocalTensor<uint8_t> &dstTensor, GlobalTensor<uint8_t> &srcTensor,
|
||||
int64_t srcOffset, uint32_t batchSize, uint32_t s1BaseSize, uint32_t s2BaseSize, int64_t s2TotalSize,
|
||||
int64_t alignedSize = blockBytes)
|
||||
{
|
||||
uint32_t alignedS2Size = CeilDiv(s2BaseSize / byteBitRatio, alignedSize) * alignedSize;
|
||||
uint32_t shapeArray[] = {batchSize * s1BaseSize, alignedS2Size};
|
||||
dstTensor.SetShapeInfo(ShapeInfo(2, shapeArray, DataFormat::ND));
|
||||
dstTensor.SetSize(batchSize * s1BaseSize * alignedS2Size);
|
||||
DataCopyParams dataCopyParams;
|
||||
dataCopyParams.blockCount = batchSize * s1BaseSize;
|
||||
dataCopyParams.blockLen = CeilDiv(s2BaseSize / byteBitRatio, blockBytes);
|
||||
dataCopyParams.dstStride = 0;
|
||||
if (s2TotalSize / byteBitRatio % alignedSize == 0 && s2BaseSize / byteBitRatio % alignedSize == 0) {
|
||||
dataCopyParams.srcStride =
|
||||
(s2TotalSize / byteBitRatio - dataCopyParams.blockLen * blockBytes) / blockBytes;
|
||||
DataCopy(dstTensor, srcTensor[srcOffset / byteBitRatio], dataCopyParams);
|
||||
} else {
|
||||
dataCopyParams.blockLen = CeilDiv(s2BaseSize , byteBitRatio);
|
||||
dataCopyParams.srcStride = (s2TotalSize - s2BaseSize) / byteBitRatio;
|
||||
DataCopyPadParams dataCopyPadParams;
|
||||
dataCopyPadParams.isPad = true;
|
||||
dataCopyPadParams.rightPadding = 0;
|
||||
dataCopyPadParams.paddingValue = 0;
|
||||
DataCopyPad(dstTensor, srcTensor[srcOffset / byteBitRatio], dataCopyParams, dataCopyPadParams);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline int32_t Align(int32_t shape)
|
||||
{
|
||||
int32_t alignFactor = 16;
|
||||
int32_t alignedSize = CeilDiv<int32_t, int32_t>(shape, alignFactor) * alignFactor;
|
||||
return alignedSize;
|
||||
}
|
||||
|
||||
#endif // FLASH_ATTENTION_UTIL_H
|
||||
Reference in New Issue
Block a user