@@ -0,0 +1,208 @@
|
||||
/*
|
||||
* 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 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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_ROW_HPP
|
||||
#define CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_ROW_HPP
|
||||
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/epilogue/dispatch_policy.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/detail/callback.hpp"
|
||||
#include "catlass/epilogue/block/block_epilogue.hpp"
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
|
||||
// float scale, dequant per expert
|
||||
template <
|
||||
uint32_t UB_STAGES_,
|
||||
class CType_,
|
||||
class LayoutPerTokenScale_,
|
||||
class DType_,
|
||||
class TileCopy_
|
||||
>
|
||||
class BlockEpilogue <
|
||||
EpilogueAtlasA2PerTokenDequant<UB_STAGES_>,
|
||||
CType_,
|
||||
Gemm::GemmType<float, LayoutPerTokenScale_>,
|
||||
DType_,
|
||||
TileCopy_
|
||||
> {
|
||||
public:
|
||||
using DispatchPolicy = EpilogueAtlasA2PerTokenDequant<UB_STAGES_>;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
|
||||
// Data infos
|
||||
using ElementC = typename CType_::Element;
|
||||
using LayoutC = typename CType_::Layout;
|
||||
using ElementPerTokenScale = float;
|
||||
using LayoutPerTokenScale = LayoutPerTokenScale_;
|
||||
using ElementD = typename DType_::Element;
|
||||
using LayoutD = typename DType_::Layout;
|
||||
|
||||
// Check data infos
|
||||
static_assert(
|
||||
(std::is_same_v<ElementC, half> || std::is_same_v<ElementC, bfloat16_t>) &&
|
||||
(std::is_same_v<ElementD, half> || std::is_same_v<ElementD, bfloat16_t>),
|
||||
"The element type template parameters of BlockEpilogue are wrong"
|
||||
);
|
||||
static_assert(
|
||||
std::is_same_v<LayoutC, layout::RowMajor> &&
|
||||
std::is_same_v<LayoutPerTokenScale, layout::VectorLayout> && std::is_same_v<LayoutD, layout::RowMajor>,
|
||||
"The layout template parameters of BlockEpilogue are wrong"
|
||||
);
|
||||
|
||||
|
||||
// Tile copy
|
||||
using CopyGmToUbC = typename TileCopy_::CopyGmToUbC;
|
||||
using CopyUbToGmD = typename TileCopy_::CopyUbToGmD;
|
||||
|
||||
struct Params {
|
||||
__gm__ int32_t *ptrTokenPerExpert{nullptr};
|
||||
int32_t EP;
|
||||
int32_t expertPerRank;
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params() {};
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params(int32_t EP_, int32_t expertPerRank_, __gm__ int32_t *ptrTokenPerExpert_) : ptrTokenPerExpert(ptrTokenPerExpert_), EP(EP_), expertPerRank(expertPerRank_) {}
|
||||
};
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> const &resource, Params const ¶ms = Params{}) : params(params)
|
||||
{
|
||||
size_t ubOffset = 4096;
|
||||
int32_t eventVMTE2 = 0;
|
||||
int32_t eventMTE2V = 0;
|
||||
int32_t eventMTE3V = 0;
|
||||
int32_t eventVMTE3 = 0;
|
||||
constexpr int32_t blockN = 12000;
|
||||
for (uint32_t i = 0; i < UB_STAGES; ++i) {
|
||||
ubCList[i] = resource.ubBuf.template GetBufferByByte<ElementC>(ubOffset);
|
||||
ubOffset += blockN * sizeof(ElementC);
|
||||
ubDList[i] = resource.ubBuf.template GetBufferByByte<ElementD>(ubOffset);
|
||||
ubOffset += blockN * sizeof(ElementD);
|
||||
|
||||
eventUbCVMTE2List[i] = eventVMTE2++;
|
||||
eventUbCMTE2VList[i] = eventMTE2V++;
|
||||
eventUbDMTE3VList[i] = eventMTE3V++;
|
||||
eventUbDVMTE3List[i] = eventVMTE3++;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[i]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[i]);
|
||||
ubCFp32List[i] = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
ubOffset += blockN * sizeof(float);
|
||||
}
|
||||
}
|
||||
CATLASS_DEVICE
|
||||
void Finalize()
|
||||
{
|
||||
for (uint32_t i = 0; i < UB_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[i]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[i]);
|
||||
}
|
||||
}
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue()
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void UpdateParams(Params const ¶ms_)
|
||||
{
|
||||
params = params_;
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator() (
|
||||
AscendC::GlobalTensor<ElementC> const &gmC,
|
||||
MatrixCoord const &shapeC,
|
||||
AscendC::GlobalTensor<ElementPerTokenScale> const &gmPerTokenScale,
|
||||
AscendC::GlobalTensor<ElementD> const &gmD
|
||||
)
|
||||
{
|
||||
uint32_t blockM = shapeC.row();
|
||||
uint32_t blockN = shapeC.column();
|
||||
|
||||
uint32_t tileLoops = blockM;
|
||||
|
||||
for (uint32_t loopIdx = 0; loopIdx < tileLoops; loopIdx ++) {
|
||||
auto gmTileC = gmC[loopIdx * blockN];
|
||||
auto &ubC = ubCList[ubListId];
|
||||
auto &ubCFp32 = ubCFp32List[ubListId];
|
||||
auto &ubMul = ubMulList[ubListId];
|
||||
auto &ubD = ubDList[ubListId];
|
||||
auto gmTileD = gmD[loopIdx * blockN];
|
||||
LayoutC layoutUbC{1, blockN};
|
||||
|
||||
// Move C from GM workspace to UB
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
copyGmToUbC(ubC, gmTileC, layoutUbC, layoutUbC);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
|
||||
// Cast C to FP32 in UB
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
AscendC::Cast(ubCFp32, ubC, AscendC::RoundMode::CAST_NONE, blockN);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
|
||||
// Get per-token scale from row loopIdx of gmPerTokenScale
|
||||
ElementPerTokenScale perTokenScale = gmPerTokenScale(loopIdx);
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::S_V>(0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::S_V>(0);
|
||||
// Multiply FP32 C by the per-token scale
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Muls(ubCFp32, ubCFp32, perTokenScale, blockN);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
// Cast the muls result back to fp16/bf16
|
||||
LayoutD layoutUbD{1, blockN};
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
|
||||
AscendC::Cast(ubD, ubCFp32, AscendC::RoundMode::CAST_RINT, blockN);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(eventUbDVMTE3List[ubListId]);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(eventUbDVMTE3List[ubListId]);
|
||||
copyUbToGmD(gmTileD, ubD, layoutUbD, layoutUbD);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
|
||||
ubListId = (ubListId + 1 < UB_STAGES) ? (ubListId + 1) : 0;
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
Params params;
|
||||
|
||||
AscendC::LocalTensor<ElementC> ubCList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementD> ubDList[UB_STAGES];
|
||||
|
||||
int32_t eventUbCVMTE2List[UB_STAGES];
|
||||
int32_t eventUbCMTE2VList[UB_STAGES];
|
||||
int32_t eventUbDMTE3VList[UB_STAGES];
|
||||
int32_t eventUbDVMTE3List[UB_STAGES];
|
||||
|
||||
uint32_t ubListId{0};
|
||||
|
||||
AscendC::LocalTensor<float> ubCFp32List[UB_STAGES];
|
||||
AscendC::LocalTensor<float> ubMulList[UB_STAGES];
|
||||
|
||||
|
||||
CopyGmToUbC copyGmToUbC;
|
||||
CopyUbToGmD copyUbToGmD;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Epilogue::Block
|
||||
|
||||
#endif // CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_ROW_HPP
|
||||
@@ -0,0 +1,402 @@
|
||||
/*
|
||||
* 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 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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_SWIGLU_HPP
|
||||
#define CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_SWIGLU_HPP
|
||||
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/epilogue/dispatch_policy.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/detail/callback.hpp"
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
|
||||
// float scale, dequant per expert
|
||||
template <
|
||||
uint32_t UB_STAGES_,
|
||||
class CType_,
|
||||
class LayoutPerTokenScale_,
|
||||
class DType_,
|
||||
class TileElemWiseMuls_,
|
||||
class TileCopy_
|
||||
>
|
||||
class BlockEpilogue <
|
||||
EpilogueAtlasA2PerTokenDequantSwigluQuant<UB_STAGES_>,
|
||||
CType_,
|
||||
Gemm::GemmType<float, LayoutPerTokenScale_>,
|
||||
DType_,
|
||||
TileElemWiseMuls_,
|
||||
TileCopy_
|
||||
> {
|
||||
public:
|
||||
using DispatchPolicy = EpilogueAtlasA2PerTokenDequantSwigluQuant<UB_STAGES_>;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
|
||||
// Data infos
|
||||
using ElementC = typename CType_::Element;
|
||||
using LayoutC = typename CType_::Layout;
|
||||
using ElementPerTokenScale = float;
|
||||
using LayoutPerTokenScale = LayoutPerTokenScale_;
|
||||
using ElementD = typename DType_::Element;
|
||||
using LayoutD = typename DType_::Layout;
|
||||
|
||||
// Check data infos
|
||||
static_assert(
|
||||
(std::is_same_v<ElementC, half> || std::is_same_v<ElementC, bfloat16_t>) &&
|
||||
(std::is_same_v<ElementD, float> || std::is_same_v<ElementD, int8_t> || std::is_same_v<ElementD, half> || std::is_same_v<ElementD, bfloat16_t>),
|
||||
"The element type template parameters of BlockEpilogue are wrong"
|
||||
);
|
||||
static_assert(
|
||||
std::is_same_v<LayoutC, layout::RowMajor> &&
|
||||
std::is_same_v<LayoutPerTokenScale, layout::VectorLayout> && std::is_same_v<LayoutD, layout::RowMajor>,
|
||||
"The layout template parameters of BlockEpilogue are wrong"
|
||||
);
|
||||
|
||||
// Tile copy
|
||||
using CopyGmToUbC = typename TileCopy_::CopyGmToUbC;
|
||||
using CopyUbToGmD = typename TileCopy_::CopyUbToGmD;
|
||||
using CopyUbToGmDequantScale = Epilogue::Tile::CopyUb2Gm<ArchTag, Gemm::GemmType<ElementPerTokenScale, LayoutPerTokenScale>>;
|
||||
|
||||
struct Params {
|
||||
__gm__ ElementPerTokenScale *ptrPerTokenScale{nullptr};
|
||||
LayoutPerTokenScale layoutPerTokenScale{};
|
||||
__gm__ ElementD *ptrD{nullptr};
|
||||
LayoutD layoutD{};
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params() {};
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params(__gm__ ElementPerTokenScale *ptrPerTokenScale_, LayoutPerTokenScale const &layoutPerTokenScale_,
|
||||
__gm__ ElementD *ptrD_, LayoutD const &layoutD_
|
||||
) : ptrPerTokenScale(ptrPerTokenScale_), layoutPerTokenScale(layoutPerTokenScale_),
|
||||
ptrD(ptrD_), layoutD(layoutD_) {}
|
||||
};
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> const &resource, int32_t n, Params const ¶ms = Params{}) : params(params)
|
||||
{
|
||||
size_t ubOffset = 0;
|
||||
int32_t eventVMTE2 = 0;
|
||||
int32_t eventMTE2V = 0;
|
||||
int32_t eventMTE3V = 0;
|
||||
int32_t eventVMTE3 = 0;
|
||||
uint32_t blockN = n;
|
||||
uint32_t ChunkTileLen = blockN / 2;
|
||||
uint32_t HalfChunkTileLen = ChunkTileLen / 2;
|
||||
|
||||
for (uint32_t i = 0; i < UB_STAGES; ++i) {
|
||||
ubCList[i] = resource.ubBuf.template GetBufferByByte<ElementC>(ubOffset);
|
||||
ubOffset += blockN * sizeof(ElementC);
|
||||
ubDList[i] = resource.ubBuf.template GetBufferByByte<ElementD>(ubOffset);
|
||||
ubOffset += blockN * sizeof(ElementD);
|
||||
ubCFp32List[i] = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
ubOffset += blockN * sizeof(float);
|
||||
ubCFp32ChunkNList[i] = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
ubOffset += ChunkTileLen * sizeof(float);
|
||||
ubCFp32ChunkNAbsList[i] = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
ubOffset += ChunkTileLen * sizeof(float);
|
||||
ubCFp32ChunkNMaxList[i] = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
ubOffset += HalfChunkTileLen * sizeof(float);
|
||||
ubQuantS32List[i] = ubCFp32ChunkNAbsList[i].template ReinterpretCast<int32_t>();
|
||||
ubQuantF16List[i] = ubCFp32ChunkNAbsList[i].template ReinterpretCast<half>();
|
||||
|
||||
eventUbCVMTE2List[i] = eventVMTE2++;
|
||||
eventUbCMTE2VList[i] = eventMTE2V++;
|
||||
eventUbDMTE3VList[i] = eventMTE3V++;
|
||||
eventUbDVMTE3List[i] = eventVMTE3++;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[i]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[i]);
|
||||
}
|
||||
|
||||
ubPerTokenScaleOutput = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
}
|
||||
CATLASS_DEVICE
|
||||
void Finalize()
|
||||
{
|
||||
for (uint32_t i = 0; i < UB_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[i]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[i]);
|
||||
}
|
||||
}
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue()
|
||||
{
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void UpdateParams(Params const ¶ms_)
|
||||
{
|
||||
params = params_;
|
||||
}
|
||||
// 每个tile就是1*7168,每个block是一个expert的所有token=[group[i], 7168]
|
||||
CATLASS_DEVICE
|
||||
void operator() (
|
||||
AscendC::GlobalTensor<ElementC> const &gmC,
|
||||
MatrixCoord const &shapeC,
|
||||
AscendC::GlobalTensor<ElementPerTokenScale> const &gmPerTokenScale1,
|
||||
AscendC::GlobalTensor<ElementD> const &gmD,
|
||||
AscendC::GlobalTensor<ElementPerTokenScale> const &gmPerTokenScale2,
|
||||
|
||||
uint32_t epilogueCoreNum = 40,
|
||||
Callback &&callback = Callback{}
|
||||
)
|
||||
{
|
||||
callback();
|
||||
uint32_t blockM = shapeC.row();
|
||||
uint32_t blockN = shapeC.column();
|
||||
|
||||
uint32_t tileLoops = blockM;
|
||||
uint32_t subblockIdx = get_block_idx() + get_subblockid() * get_block_num();
|
||||
|
||||
uint32_t subblockNum = get_block_num() * 2;
|
||||
uint32_t moveDataCoreNum = subblockNum - epilogueCoreNum;
|
||||
|
||||
if (subblockIdx < moveDataCoreNum) {
|
||||
return;
|
||||
}
|
||||
uint32_t epilogueCoreIdx = subblockIdx - moveDataCoreNum;
|
||||
|
||||
uint32_t perCoreData = blockM / epilogueCoreNum;
|
||||
uint32_t remainderData = blockM % epilogueCoreNum;
|
||||
|
||||
uint32_t tasksForIdx = epilogueCoreIdx < remainderData ? perCoreData + 1 : perCoreData;
|
||||
uint32_t loopStartIdx = epilogueCoreIdx * perCoreData + (epilogueCoreIdx < remainderData? epilogueCoreIdx : remainderData);
|
||||
|
||||
uint32_t alignedPerCoreData = RoundUp<BYTE_PER_BLK / sizeof(ElementPerTokenScale)>(perCoreData + 1);
|
||||
|
||||
uint32_t ChunkTileLen = blockN / 2;
|
||||
uint32_t HalfChunkTileLen = ChunkTileLen / 2;
|
||||
|
||||
|
||||
for (uint32_t loopIdx = loopStartIdx; loopIdx < loopStartIdx + tasksForIdx; ++loopIdx) {
|
||||
|
||||
auto gmTileC = gmC[loopIdx * blockN];
|
||||
|
||||
auto &ubC = ubCList[ubListId];
|
||||
auto &ubD = ubDList[ubListId];
|
||||
|
||||
auto &ubCFp32 = ubCFp32List[ubListId];
|
||||
auto &ubCFp32ChunkN = ubCFp32ChunkNList[ubListId];
|
||||
auto &ubAbs = ubCFp32ChunkNAbsList[ubListId];
|
||||
// auto &ubMax = ubCFp32ChunkNMaxList[ubListId];
|
||||
auto &ubReduceMax = ubCFp32ChunkNMaxList[ubListId];
|
||||
auto &ubOutputTmp = ubAbs;
|
||||
auto &sharedUbTmpBuffer = ubReduceMax;
|
||||
auto &ubQuantS32 = ubQuantS32List[ubListId];
|
||||
auto &ubQuantF16 = ubQuantF16List[ubListId];
|
||||
|
||||
auto gmTileD = gmD[loopIdx * ChunkTileLen];
|
||||
LayoutC layoutUbC{1, blockN};
|
||||
|
||||
// 把C从GM workspace搬到UB
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
copyGmToUbC(ubC, gmTileC, layoutUbC, layoutUbC);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
|
||||
// 在UB上做把C cast成FP32
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
AscendC::Cast(ubCFp32, ubC, AscendC::RoundMode::CAST_NONE, blockN);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
|
||||
// 获取pertoken scale值,gmPerTokenScale的第loopIdx行
|
||||
ElementPerTokenScale perTokenScale = gmPerTokenScale1(loopIdx);
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::S_V>(0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::S_V>(0);
|
||||
// pertoken scale值与FP32的C做Muls乘法
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Muls(ubCFp32, ubCFp32, perTokenScale, blockN);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
//swiglue计算过程
|
||||
AscendC::Muls(ubCFp32ChunkN, ubCFp32, -1.0f, ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Exp(ubCFp32ChunkN, ubCFp32ChunkN, ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Adds(ubCFp32ChunkN, ubCFp32ChunkN, 1.0f, ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
//TODO除的时候是否会对之后的数据有影响;
|
||||
AscendC::Div(ubCFp32ChunkN, ubCFp32, ubCFp32ChunkN, ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Mul(ubCFp32ChunkN, ubCFp32ChunkN, ubCFp32[ChunkTileLen], ChunkTileLen);
|
||||
|
||||
//quant过程,两种方式区别;
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Abs(ubAbs, ubCFp32ChunkN, ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::ReduceMax<float>(ubReduceMax, ubAbs, sharedUbTmpBuffer, ChunkTileLen, false);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_S>(0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_S>(0);
|
||||
|
||||
//TODO两种计算方法的效率比较
|
||||
ElementPerTokenScale GMubDequantScale = ubReduceMax.GetValue(0);
|
||||
AscendC::SetFlag<AscendC::HardEvent::S_V>(0);
|
||||
|
||||
auto ubPerTokenScaleOutputOffset = loopIdx - loopStartIdx;
|
||||
ubPerTokenScaleOutput.SetValue(ubPerTokenScaleOutputOffset, GMubDequantScale / 127.f);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::S_V>(0);
|
||||
AscendC::Muls(ubOutputTmp, ubCFp32ChunkN, 127.f / GMubDequantScale, ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::Cast(ubQuantS32, ubOutputTmp, AscendC::RoundMode::CAST_RINT, ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::SetDeqScale(static_cast<half>(1.0));
|
||||
AscendC::Cast(ubQuantF16, ubQuantS32, AscendC::RoundMode::CAST_RINT, ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDVMTE3List[ubListId]);
|
||||
AscendC::Cast(ubD, ubQuantF16, AscendC::RoundMode::CAST_RINT, ChunkTileLen);
|
||||
// AscendC::Muls(ubD, ubCFp32ChunkN, 127.f / GMubDequantScale, ChunkTileLen);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(eventUbDMTE3VList[ubListId]);
|
||||
|
||||
LayoutD layoutUbD{1, ChunkTileLen};
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(eventUbDVMTE3List[ubListId]);
|
||||
copyUbToGmD(gmTileD, ubD, layoutUbD, layoutUbD);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
ubListId = (ubListId + 1 < UB_STAGES) ? (ubListId + 1) : 0;
|
||||
}
|
||||
|
||||
if(tasksForIdx > 0){
|
||||
LayoutPerTokenScale layoutGmPerTokenScale2{tasksForIdx};
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::S_MTE3>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::S_MTE3>(EVENT_ID0);
|
||||
|
||||
copyUbToGmDequantScale(gmPerTokenScale2[loopStartIdx], ubPerTokenScaleOutput[0], layoutGmPerTokenScale2, layoutGmPerTokenScale2);
|
||||
}
|
||||
|
||||
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator() (
|
||||
AscendC::GlobalTensor<ElementC> const &gmC,
|
||||
MatrixCoord const &shapeC,
|
||||
AscendC::GlobalTensor<ElementD> const &gmD,
|
||||
uint32_t epilogueCoreNum = 40,
|
||||
Callback &&callback = Callback{}
|
||||
)
|
||||
{
|
||||
callback();
|
||||
uint32_t blockM = shapeC.row();
|
||||
uint32_t blockN = shapeC.column();
|
||||
|
||||
uint32_t tileLoops = blockM;
|
||||
uint32_t subblockIdx = get_block_idx() + get_subblockid() * get_block_num();
|
||||
//uint32_t subblockIdx = get_block_idx() * 2 + get_subblockid();
|
||||
|
||||
uint32_t subblockNum = get_block_num() * 2;
|
||||
uint32_t moveDataCoreNum = subblockNum - epilogueCoreNum;
|
||||
|
||||
if (subblockIdx < moveDataCoreNum) {
|
||||
return;
|
||||
}
|
||||
uint32_t epilogueCoreIdx = subblockIdx - moveDataCoreNum;
|
||||
|
||||
|
||||
uint32_t perCoreData = blockM / epilogueCoreNum;
|
||||
uint32_t remainderData = blockM % epilogueCoreNum;
|
||||
|
||||
uint32_t tasksForIdx = epilogueCoreIdx < remainderData ? perCoreData + 1 : perCoreData;
|
||||
uint32_t loopStartIdx = epilogueCoreIdx * perCoreData + (epilogueCoreIdx < remainderData? epilogueCoreIdx : remainderData);
|
||||
|
||||
uint32_t alignedPerCoreData = RoundUp<BYTE_PER_BLK / sizeof(ElementPerTokenScale)>(perCoreData + 1);
|
||||
|
||||
uint32_t ChunkTileLen = blockN / 2;
|
||||
uint32_t HalfChunkTileLen = ChunkTileLen / 2;
|
||||
|
||||
|
||||
for (uint32_t loopIdx = loopStartIdx; loopIdx < loopStartIdx + tasksForIdx; ++loopIdx) {
|
||||
|
||||
auto gmTileC = gmC[loopIdx * blockN];
|
||||
|
||||
auto &ubC = ubCList[ubListId];
|
||||
auto &ubD = ubDList[ubListId];
|
||||
|
||||
auto &ubCFp32 = ubCFp32List[ubListId];
|
||||
auto &ubCFp32ChunkN = ubCFp32ChunkNList[ubListId];
|
||||
|
||||
auto gmTileD = gmD[loopIdx * ChunkTileLen];
|
||||
LayoutC layoutUbC{1, blockN};
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
copyGmToUbC(ubC, gmTileC, layoutUbC, layoutUbC);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
AscendC::Cast(ubCFp32, ubC, AscendC::RoundMode::CAST_NONE, blockN);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
|
||||
AscendC::Muls(ubCFp32ChunkN, ubCFp32, -1.0f, ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Exp(ubCFp32ChunkN, ubCFp32ChunkN, ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Adds(ubCFp32ChunkN, ubCFp32ChunkN, 1.0f, ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Div(ubCFp32ChunkN, ubCFp32, ubCFp32ChunkN, ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Mul(ubCFp32ChunkN, ubCFp32ChunkN, ubCFp32[ChunkTileLen], ChunkTileLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDVMTE3List[ubListId]);
|
||||
AscendC::Cast(ubD, ubCFp32ChunkN, AscendC::RoundMode::CAST_ROUND, ChunkTileLen);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(eventUbDMTE3VList[ubListId]);
|
||||
|
||||
LayoutD layoutUbD{1, ChunkTileLen};
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(eventUbDVMTE3List[ubListId]);
|
||||
// copyUbToGmD(gmTileD, ubCFp32ChunkN, layoutUbD, layoutUbD);
|
||||
copyUbToGmD(gmTileD, ubD, layoutUbD, layoutUbD);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
ubListId = (ubListId + 1 < UB_STAGES) ? (ubListId + 1) : 0;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
|
||||
private:
|
||||
Params params;
|
||||
|
||||
AscendC::LocalTensor<ElementC> ubCList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementD> ubDList[UB_STAGES];
|
||||
|
||||
int32_t eventUbCVMTE2List[UB_STAGES];
|
||||
int32_t eventUbCMTE2VList[UB_STAGES];
|
||||
int32_t eventUbDMTE3VList[UB_STAGES];
|
||||
int32_t eventUbDVMTE3List[UB_STAGES];
|
||||
|
||||
uint32_t ubListId{0};
|
||||
|
||||
AscendC::LocalTensor<float> ubCFp32List[UB_STAGES];
|
||||
AscendC::LocalTensor<float> ubCFp32ChunkNList[UB_STAGES];
|
||||
AscendC::LocalTensor<float> ubCFp32ChunkNAbsList[UB_STAGES];
|
||||
AscendC::LocalTensor<float> ubCFp32ChunkNMaxList[UB_STAGES];
|
||||
AscendC::LocalTensor<int32_t> ubQuantS32List[UB_STAGES];
|
||||
AscendC::LocalTensor<half> ubQuantF16List[UB_STAGES];
|
||||
AscendC::LocalTensor<float> ubPerTokenScaleOutput;
|
||||
|
||||
|
||||
CopyGmToUbC copyGmToUbC;
|
||||
CopyUbToGmD copyUbToGmD;
|
||||
CopyUbToGmDequantScale copyUbToGmDequantScale;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Epilogue::Block
|
||||
|
||||
#endif // CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_SWIGLU_HPP
|
||||
@@ -0,0 +1,201 @@
|
||||
#ifndef CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_V2_ONLY_HPP
|
||||
#define CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_V2_ONLY_HPP
|
||||
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/epilogue/dispatch_policy.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/detail/callback.hpp"
|
||||
|
||||
#include "hccl_shmem.hpp"
|
||||
#include "layout3d.hpp"
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
template <
|
||||
uint32_t UB_STAGES_,
|
||||
class CType_,
|
||||
class LayoutPerTokenScale_,
|
||||
class DType_,
|
||||
class TileCopy_
|
||||
>
|
||||
class BlockEpilogue <
|
||||
EpilogueAtlasA2PerTokenDequantV2<UB_STAGES_>,
|
||||
CType_,
|
||||
Gemm::GemmType<float, LayoutPerTokenScale_>,
|
||||
DType_,
|
||||
TileCopy_
|
||||
> {
|
||||
public:
|
||||
using DispatchPolicy = EpilogueAtlasA2PerTokenDequantV2<UB_STAGES_>;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
|
||||
// Data infos
|
||||
using ElementC = typename CType_::Element;
|
||||
using LayoutC = typename CType_::Layout;
|
||||
using ElementPerTokenScale = float;
|
||||
using LayoutPerTokenScale = LayoutPerTokenScale_;
|
||||
using ElementD = typename DType_::Element;
|
||||
using LayoutD = typename DType_::Layout;
|
||||
|
||||
//using CopyScaleGmToUb = Epilogue::Tile::CopyGm2Ub<ArchTag, Gemm::GemmType<float, layout::RowMajor>>;
|
||||
using CopyScaleGmToUb = Epilogue::Tile::CopyGm2Ub<ArchTag, Gemm::GemmType<float, layout::VectorLayout>>;
|
||||
// Tile copy
|
||||
using CopyGmToUbC = typename TileCopy_::CopyGmToUbC;
|
||||
using CopyUbToGmD = typename TileCopy_::CopyUbToGmD;
|
||||
|
||||
struct Params {
|
||||
__gm__ int32_t *ptrTokenPerExpert{nullptr};
|
||||
int32_t EP;
|
||||
int32_t expertPerRank;
|
||||
int32_t n2;
|
||||
LayoutC layoutC;
|
||||
int32_t n0;
|
||||
int32_t rank;
|
||||
HcclShmem shmem;
|
||||
int32_t offsetD;
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params() {};
|
||||
CATLASS_DEVICE
|
||||
Params(int32_t EP_, int32_t expertPerRank_, int32_t rank_, __gm__ int32_t *ptrTokenPerExpert_,
|
||||
LayoutC layoutC_, int32_t n2_, int32_t n0_, HcclShmem& shmem_, int32_t offsetD_) :
|
||||
ptrTokenPerExpert(ptrTokenPerExpert_), EP(EP_),
|
||||
expertPerRank(expertPerRank_),rank(rank_), layoutC(layoutC_), n2(n2_), n0(n0_),
|
||||
shmem(shmem_), offsetD(offsetD_)
|
||||
{}
|
||||
};
|
||||
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> const &resource, Params const ¶ms = Params{}) : params(params)
|
||||
{
|
||||
//ub:192KB
|
||||
n0 = params.n0;
|
||||
size_t ubOffset = 0;
|
||||
for(int32_t i = 0; i < 2; i++) {
|
||||
ubCList[i] = resource.ubBuf.template GetBufferByByte<ElementC>(ubOffset);
|
||||
ubOffset += max_len * sizeof(ElementC);
|
||||
ubDList[i] = resource.ubBuf.template GetBufferByByte<ElementD>(ubOffset);
|
||||
ubOffset += max_len * sizeof(ElementD);
|
||||
ubFp32List[i] = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
ubOffset += max_len * sizeof(float);
|
||||
scaleUbList[i] = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
ubOffset += (max_len / n0) * sizeof(float);
|
||||
source_scale_offset[i] = -1;
|
||||
}
|
||||
tokenPerExpert.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(params.ptrTokenPerExpert));
|
||||
tokenPerExpertLayout = Layout3D(AlignUp(params.EP * params.expertPerRank, ALIGN_128), params.expertPerRank);
|
||||
is_ping = true;
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void Finalize()
|
||||
{
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID1);
|
||||
}
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue()
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator() (
|
||||
AscendC::GlobalTensor<ElementC> const &gmC,
|
||||
GemmCoord& blockCoord,
|
||||
GemmCoord& actualBlockShape,
|
||||
int32_t groupIdx,
|
||||
int32_t preSrcExpertSum,
|
||||
AscendC::GlobalTensor<int32_t> preSumBeforeRank
|
||||
){
|
||||
is_ping = !is_ping;
|
||||
auto event_id = is_ping ? EVENT_ID0 : EVENT_ID1;
|
||||
|
||||
auto &ubC = ubCList[is_ping];
|
||||
int32_t gmCOffset = preSrcExpertSum * params.n2 + blockCoord.m() * params.n2 + blockCoord.n();
|
||||
auto gmTileC = gmC[gmCOffset];
|
||||
|
||||
LayoutC layoutGM{actualBlockShape.m(), actualBlockShape.n(), params.n2};
|
||||
LayoutC layoutUB{actualBlockShape.m(), actualBlockShape.n(), n0};
|
||||
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(event_id); //for debug
|
||||
copyGmToUbC(ubC, gmTileC, layoutUB, layoutGM);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE3>(event_id); //for debug
|
||||
|
||||
int32_t lenTile = actualBlockShape.m();
|
||||
int32_t stTile = blockCoord.m();
|
||||
int32_t edTile = stTile + lenTile;
|
||||
int32_t preSumRankInExpert = 0;
|
||||
int32_t tileOffset = 0;
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE3>(event_id);
|
||||
for (int32_t dstEpIdx = 0; dstEpIdx < params.EP; dstEpIdx ++) {
|
||||
int32_t lenRankInExpert = tokenPerExpert(tokenPerExpertLayout(dstEpIdx, params.rank, groupIdx));
|
||||
int32_t dstExpertOffset = preSumBeforeRank(dstEpIdx * params.expertPerRank + groupIdx);
|
||||
int32_t stRankInExpert = preSumRankInExpert;
|
||||
int32_t edRankInExpert = stRankInExpert + lenRankInExpert;
|
||||
preSumRankInExpert += lenRankInExpert;
|
||||
if (stRankInExpert >= edTile) {
|
||||
break;
|
||||
}
|
||||
else if (edRankInExpert <= stTile) {
|
||||
continue;
|
||||
}
|
||||
int32_t stData = max(stRankInExpert, stTile);
|
||||
int32_t edData = min(edRankInExpert, edTile);
|
||||
uint32_t lenData = edData - stData;
|
||||
if (lenData <= 0){
|
||||
continue;
|
||||
}
|
||||
|
||||
uint32_t dstOffsetInExpert = 0;
|
||||
if (stTile > stRankInExpert) {
|
||||
dstOffsetInExpert = stTile - stRankInExpert;
|
||||
}
|
||||
AscendC::GlobalTensor<ElementD> gmRemotePeer;
|
||||
__gm__ void* dstPeermemPtr = params.shmem(params.offsetD, dstEpIdx);
|
||||
gmRemotePeer.SetGlobalBuffer(reinterpret_cast<__gm__ ElementD*>(dstPeermemPtr));
|
||||
MatrixCoord dstOffset{dstOffsetInExpert + dstExpertOffset, blockCoord.n()};
|
||||
int64_t gmDstOffset = params.layoutC.GetOffset(dstOffset);
|
||||
auto gmTileD = gmRemotePeer[gmDstOffset];
|
||||
LayoutC layoutGM2{lenData, actualBlockShape.n(), params.n2};
|
||||
LayoutC layoutUB2{lenData, actualBlockShape.n(), n0};
|
||||
copyUbToGmD(gmTileD, ubC[tileOffset * n0], layoutGM2, layoutUB2);
|
||||
tileOffset += lenData;
|
||||
}
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(event_id);
|
||||
|
||||
}
|
||||
|
||||
private:
|
||||
|
||||
Params params;
|
||||
AscendC::LocalTensor<ElementC> ubCList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementD> ubDList[UB_STAGES];
|
||||
AscendC::LocalTensor<float> ubFp32List[UB_STAGES];
|
||||
AscendC::LocalTensor<float> scaleUbList[UB_STAGES];
|
||||
int32_t source_scale_offset[UB_STAGES];
|
||||
|
||||
int32_t max_len = 8 * 32 / 4 * 128;
|
||||
int32_t n0;
|
||||
bool is_ping = false;
|
||||
|
||||
|
||||
int32_t repeat = 128;
|
||||
|
||||
|
||||
CopyGmToUbC copyGmToUbC;
|
||||
CopyUbToGmD copyUbToGmD;
|
||||
|
||||
CopyScaleGmToUb copyScaleGmToUb;
|
||||
AscendC::GlobalTensor<int32_t> tokenPerExpert;
|
||||
Layout3D tokenPerExpertLayout;
|
||||
};
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,556 @@
|
||||
/*
|
||||
* 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 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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_GEMM_BLOCK_BLOCK_MMAD_PRELOAD_FIXPIPE_QUANT_HPP
|
||||
#define CATLASS_GEMM_BLOCK_BLOCK_MMAD_PRELOAD_FIXPIPE_QUANT_HPP
|
||||
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/coord.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/gemm/dispatch_policy.hpp"
|
||||
#include "catlass/gemm/helper.hpp"
|
||||
#include "dispatch_policy_custom.hpp"
|
||||
|
||||
|
||||
namespace Catlass::Gemm::Block {
|
||||
|
||||
template<AscendC::HardEvent event>
|
||||
__aicore__ inline void SyncFlagFunc(int32_t eventID)
|
||||
{
|
||||
AscendC::SetFlag<event>(eventID);
|
||||
AscendC::WaitFlag<event>(eventID);
|
||||
}
|
||||
|
||||
template <
|
||||
uint32_t PRELOAD_STAGES_,
|
||||
uint32_t L1_STAGES_,
|
||||
uint32_t L0A_STAGES_,
|
||||
uint32_t L0B_STAGES_,
|
||||
uint32_t L0C_STAGES_,
|
||||
bool ENABLE_UNIT_FLAG_,
|
||||
bool ENABLE_SHUFFLE_K_,
|
||||
class L1TileShape_,
|
||||
class L0TileShape_,
|
||||
class AType_,
|
||||
class BType_,
|
||||
class CType_,
|
||||
class BiasType_,
|
||||
class TileCopy_,
|
||||
class TileMmad_
|
||||
>
|
||||
struct BlockMmad <
|
||||
MmadAtlasA2PreloadAsyncFixpipe<
|
||||
PRELOAD_STAGES_,
|
||||
L1_STAGES_,
|
||||
L0A_STAGES_,
|
||||
L0B_STAGES_,
|
||||
L0C_STAGES_,
|
||||
ENABLE_UNIT_FLAG_,
|
||||
ENABLE_SHUFFLE_K_
|
||||
>,
|
||||
L1TileShape_,
|
||||
L0TileShape_,
|
||||
AType_,
|
||||
BType_,
|
||||
CType_,
|
||||
BiasType_,
|
||||
TileCopy_,
|
||||
TileMmad_
|
||||
> {
|
||||
public:
|
||||
// Type Aliases
|
||||
using DispatchPolicy = MmadAtlasA2PreloadAsyncFixpipe<
|
||||
PRELOAD_STAGES_,
|
||||
L1_STAGES_,
|
||||
L0A_STAGES_,
|
||||
L0B_STAGES_,
|
||||
L0C_STAGES_,
|
||||
ENABLE_UNIT_FLAG_,
|
||||
ENABLE_SHUFFLE_K_
|
||||
>;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
using L1TileShape = L1TileShape_;
|
||||
using L0TileShape = L0TileShape_;
|
||||
using ElementA = typename AType_::Element;
|
||||
using LayoutA = typename AType_::Layout;
|
||||
using ElementB = typename BType_::Element;
|
||||
using LayoutB = typename BType_::Layout;
|
||||
using ElementC = typename CType_::Element;
|
||||
using LayoutC = typename CType_::Layout;
|
||||
using TileMmad = TileMmad_;
|
||||
using CopyGmToL1A = typename TileCopy_::CopyGmToL1A;
|
||||
using CopyGmToL1B = typename TileCopy_::CopyGmToL1B;
|
||||
using CopyGmToL1S = Gemm::Tile::CopyGmToL1<ArchTag, Gemm::GemmType<uint64_t, layout::VectorLayout>>;
|
||||
using CopyL1ToFP = typename Gemm::Tile::QuantTileCopy<ArchTag, AType_, BType_, CType_, void, Catlass::Gemm::Tile::ScaleGranularity::PER_CHANNEL>::CopyL1ToFP;
|
||||
using CopyL1ToL0A = typename TileCopy_::CopyL1ToL0A;
|
||||
using CopyL1ToL0B = typename TileCopy_::CopyL1ToL0B;
|
||||
|
||||
using ElementAccumulator =
|
||||
typename Gemm::helper::ElementAccumulatorSelector<ElementA, ElementB>::ElementAccumulator;
|
||||
using CopyL0CToGm = typename std::conditional<
|
||||
std::is_same_v<ElementA, int8_t>,
|
||||
Gemm::Tile::CopyL0CToGm<ArchTag, ElementAccumulator, CType_, Gemm::Tile::ScaleGranularity::PER_CHANNEL>,
|
||||
typename TileCopy_::CopyL0CToGm
|
||||
>::type;
|
||||
using LayoutAInL1 = typename CopyL1ToL0A::LayoutSrc;
|
||||
using LayoutBInL1 = typename CopyL1ToL0B::LayoutSrc;
|
||||
using LayoutAInL0 = typename CopyL1ToL0A::LayoutDst;
|
||||
using LayoutBInL0 = typename CopyL1ToL0B::LayoutDst;
|
||||
using LayoutCInL0 = layout::zN;
|
||||
|
||||
using L1AAlignHelper = Gemm::helper::L1AlignHelper<ElementA, LayoutA>;
|
||||
using L1BAlignHelper = Gemm::helper::L1AlignHelper<ElementB, LayoutB>;
|
||||
|
||||
static constexpr uint32_t PRELOAD_STAGES = DispatchPolicy::PRELOAD_STAGES;
|
||||
static constexpr uint32_t L1_STAGES = DispatchPolicy::L1_STAGES;
|
||||
static constexpr uint32_t L0A_STAGES = DispatchPolicy::L0A_STAGES;
|
||||
static constexpr uint32_t L0B_STAGES = DispatchPolicy::L0B_STAGES;
|
||||
static constexpr uint32_t L0C_STAGES = DispatchPolicy::L0C_STAGES;
|
||||
|
||||
static constexpr bool ENABLE_UNIT_FLAG = DispatchPolicy::ENABLE_UNIT_FLAG;
|
||||
static constexpr bool ENABLE_SHUFFLE_K = DispatchPolicy::ENABLE_SHUFFLE_K;
|
||||
|
||||
// L1 tile size
|
||||
static constexpr uint32_t L1A_TILE_SIZE = L1TileShape::M * L1TileShape::K * sizeof(ElementA);
|
||||
static constexpr uint32_t L1B_TILE_SIZE = L1TileShape::N * L1TileShape::K * sizeof(ElementB);
|
||||
static constexpr uint32_t L1S_TILE_SIZE = L1TileShape::N * sizeof(int64_t);
|
||||
// L0 tile size
|
||||
static constexpr uint32_t L0A_TILE_SIZE = L0TileShape::M * L0TileShape::K * sizeof(ElementA);
|
||||
static constexpr uint32_t L0B_TILE_SIZE = L0TileShape::K * L0TileShape::N * sizeof(ElementB);
|
||||
static constexpr uint32_t L0C_TILE_SIZE = L1TileShape::M * L1TileShape::N * sizeof(ElementAccumulator);
|
||||
|
||||
// Check LayoutC
|
||||
static_assert(std::is_same_v<LayoutC, layout::RowMajor>, "LayoutC only support RowMajor yet!");
|
||||
|
||||
// Check L1TileShape
|
||||
static_assert(
|
||||
(std::is_same_v<ElementA, int8_t>
|
||||
? (L1A_TILE_SIZE + L1B_TILE_SIZE + L1S_TILE_SIZE) * L1_STAGES <= ArchTag::L1_SIZE
|
||||
: (L1A_TILE_SIZE + L1B_TILE_SIZE) * L1_STAGES <= ArchTag::L1_SIZE),
|
||||
"L1TileShape exceeding the L1 space for the given data type"
|
||||
);
|
||||
|
||||
// Check L0TileShape
|
||||
static_assert(L0A_TILE_SIZE * L0A_STAGES <= ArchTag::L0A_SIZE, "L0TileShape exceeding the L0A space!");
|
||||
static_assert(L0B_TILE_SIZE * L0B_STAGES <= ArchTag::L0B_SIZE, "L0TileShape exceeding the L0B space!");
|
||||
static_assert(L0C_TILE_SIZE * L0C_STAGES <= ArchTag::L0C_SIZE, "L0TileShape exceeding the L0C space!");
|
||||
|
||||
static_assert(L1TileShape::M == L0TileShape::M && L1TileShape::N == L0TileShape::N,
|
||||
"The situation where the basic blocks of L1 and L0 differ on the m and n axes is not supported yet");
|
||||
|
||||
static constexpr auto L1A_LAYOUT = LayoutAInL1::template MakeLayout<ElementA>(
|
||||
L1TileShape::M, L1TileShape::K);
|
||||
static constexpr auto L1B_LAYOUT = LayoutBInL1::template MakeLayout<ElementB>(
|
||||
L1TileShape::K, L1TileShape::N);
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockMmad(Arch::Resource<ArchTag> &resource, __gm__ int32_t* flagPtr = nullptr, int32_t expertPerRank = 0, uint32_t l1BufAddrStart = 0, uint32_t FpAddrStart = 0)
|
||||
{
|
||||
syncGroupIdx = 0;
|
||||
ptrSoftFlagBase_ = flagPtr;
|
||||
expertPerRank_ = expertPerRank;
|
||||
InitL1(resource, l1BufAddrStart);
|
||||
InitFpBuf(resource, FpAddrStart);
|
||||
InitL0A(resource);
|
||||
InitL0B(resource);
|
||||
InitL0C(resource);
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
~BlockMmad()
|
||||
{
|
||||
SynchronizeBlock();
|
||||
for (uint32_t i = 0; i < L1_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[i]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[i]);
|
||||
}
|
||||
for (uint32_t i = 0; i < L0A_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0AEventList[i]);
|
||||
}
|
||||
for (uint32_t i = 0; i < L0B_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0BEventList[i]);
|
||||
}
|
||||
for (uint32_t i = 0; i < L0C_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::FIX_M>(l0CEventList[i]);
|
||||
}
|
||||
if constexpr (std::is_same_v<ElementA, int8_t>) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::FIX_MTE2>(0);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(
|
||||
AscendC::GlobalTensor<ElementA> const &gmBlockA, LayoutA const &layoutA,
|
||||
AscendC::GlobalTensor<ElementB> const &gmBlockB, LayoutB const &layoutB,
|
||||
AscendC::GlobalTensor<ElementC> const &gmBlockC, LayoutC const &layoutC,
|
||||
AscendC::GlobalTensor<uint64_t> const &gmBlockS, layout::VectorLayout const &layoutScale,
|
||||
GemmCoord const &actualShape, int32_t syncLoopIdx = -1, int32_t flag = 0
|
||||
)
|
||||
{
|
||||
uint32_t kTileCount = CeilDiv<L1TileShape::K>(actualShape.k());
|
||||
|
||||
uint32_t mRound = RoundUp<L1AAlignHelper::M_ALIGNED>(actualShape.m());
|
||||
uint32_t nRound = RoundUp<L1BAlignHelper::N_ALIGNED>(actualShape.n());
|
||||
|
||||
uint32_t startTileIdx = 0;
|
||||
if constexpr (ENABLE_SHUFFLE_K) {
|
||||
startTileIdx = AscendC::GetBlockIdx() % kTileCount;
|
||||
}
|
||||
|
||||
for (uint32_t kLoopIdx = 0; kLoopIdx < kTileCount; ++kLoopIdx) {
|
||||
uint32_t kTileIdx = (startTileIdx + kLoopIdx < kTileCount) ?
|
||||
(startTileIdx + kLoopIdx) : (startTileIdx + kLoopIdx - kTileCount);
|
||||
|
||||
uint32_t kActual = (kTileIdx < kTileCount - 1) ?
|
||||
L1TileShape::K : (actualShape.k() - kTileIdx * L1TileShape::K);
|
||||
|
||||
// Emission load instruction from GM to L1
|
||||
MatrixCoord gmTileAOffset{0, kTileIdx * L1TileShape::K};
|
||||
MatrixCoord gmTileBOffset{kTileIdx * L1TileShape::K, 0};
|
||||
auto gmTileA = gmBlockA[layoutA.GetOffset(gmTileAOffset)];
|
||||
auto gmTileB = gmBlockB[layoutB.GetOffset(gmTileBOffset)];
|
||||
// Load first matrix A tile from GM to L1
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[l1ListId]);
|
||||
auto layoutTileA = layoutA.GetTileLayout(MakeCoord(actualShape.m(), kActual));
|
||||
copyGmToL1A(l1ATensorList[l1ListId], gmTileA, L1A_LAYOUT, layoutTileA);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE1>(l1AEventList[l1ListId]);
|
||||
// Load first matrix B tile from GM to L1
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[l1ListId]);
|
||||
auto layoutTileB = layoutB.GetTileLayout(MakeCoord(kActual, actualShape.n()));
|
||||
copyGmToL1B(l1BTensorList[l1ListId], gmTileB, L1B_LAYOUT, layoutTileB);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE1>(l1BEventList[l1ListId]);
|
||||
|
||||
// If the number of preload instructions reaches the upper limit, perform an mmad calculation on L1 tile
|
||||
if (preloadCount == PRELOAD_STAGES) {
|
||||
L1TileMmad(l1TileMmadParamsList[l1TileMmadParamsId]);
|
||||
}
|
||||
|
||||
// Store the current load status
|
||||
uint32_t preloadL1TileMmadParamsId = (l1TileMmadParamsId + preloadCount < PRELOAD_STAGES) ?
|
||||
(l1TileMmadParamsId + preloadCount) : (l1TileMmadParamsId + preloadCount - PRELOAD_STAGES);
|
||||
auto &l1TileMmadParams = l1TileMmadParamsList[preloadL1TileMmadParamsId];
|
||||
l1TileMmadParams.l1ListId = l1ListId;
|
||||
l1TileMmadParams.mRound = mRound;
|
||||
l1TileMmadParams.nRound = nRound;
|
||||
l1TileMmadParams.kActual = kActual;
|
||||
l1TileMmadParams.isKLoopFirst = (kLoopIdx == 0);
|
||||
l1TileMmadParams.isKLoopLast = (kLoopIdx == kTileCount - 1);
|
||||
l1TileMmadParams.flag = flag;
|
||||
if (kLoopIdx == kTileCount - 1) {
|
||||
l1TileMmadParams.gmBlockC = gmBlockC;
|
||||
l1TileMmadParams.gmBlockS = gmBlockS;
|
||||
l1TileMmadParams.layoutCInGm = layoutC.GetTileLayout(actualShape.GetCoordMN());
|
||||
l1TileMmadParams.layoutScale = layoutScale;
|
||||
l1TileMmadParams.syncLoopIdx = syncLoopIdx;
|
||||
}
|
||||
|
||||
if (preloadCount < PRELOAD_STAGES) {
|
||||
++preloadCount;
|
||||
} else {
|
||||
l1TileMmadParamsId = (l1TileMmadParamsId + 1 < PRELOAD_STAGES) ? (l1TileMmadParamsId + 1) : 0;
|
||||
}
|
||||
l1ListId = (l1ListId + 1 < L1_STAGES) ? (l1ListId + 1) : 0;
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void SynchronizeBlock()
|
||||
{
|
||||
while (preloadCount > 0) {
|
||||
L1TileMmad(l1TileMmadParamsList[l1TileMmadParamsId]);
|
||||
l1TileMmadParamsId = (l1TileMmadParamsId + 1 < PRELOAD_STAGES) ? (l1TileMmadParamsId + 1) : 0;
|
||||
--preloadCount;
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void Finalize(int32_t target, int32_t flag = 0)
|
||||
{
|
||||
if (ptrSoftFlagBase_ != nullptr) {
|
||||
if (target < 0) {
|
||||
return;
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::FIX_MTE3>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::FIX_MTE3>(EVENT_ID0);
|
||||
AscendC::GlobalTensor<int32_t> flagGlobal;
|
||||
flagGlobal.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(ptrSoftFlagBase_) + (expertPerRank_ + AscendC::GetBlockIdx()) * FLAGSTRIDE);
|
||||
AscendC::DataCopy(flagGlobal, l1FTensor[target * 16], FLAGSTRIDE);
|
||||
}
|
||||
else {
|
||||
for(;syncGroupIdx <= target; syncGroupIdx++) {
|
||||
int32_t flagId = syncGroupIdx / 15 + flag;
|
||||
AscendC::CrossCoreSetFlag<0x2, PIPE_FIX>(flagId);
|
||||
}
|
||||
}
|
||||
}
|
||||
private:
|
||||
struct L1TileMmadParams {
|
||||
uint32_t l1ListId;
|
||||
uint32_t mRound;
|
||||
uint32_t nRound;
|
||||
uint32_t kActual;
|
||||
bool isKLoopFirst;
|
||||
bool isKLoopLast;
|
||||
AscendC::GlobalTensor<ElementC> gmBlockC;
|
||||
AscendC::GlobalTensor<uint64_t> gmBlockS;
|
||||
LayoutC layoutCInGm;
|
||||
layout::VectorLayout layoutScale;
|
||||
int32_t syncLoopIdx;
|
||||
int32_t flag;
|
||||
CATLASS_DEVICE
|
||||
L1TileMmadParams() = default;
|
||||
};
|
||||
|
||||
CATLASS_DEVICE
|
||||
void InitL1(Arch::Resource<ArchTag> &resource, uint32_t l1BufAddrStart)
|
||||
{
|
||||
uint32_t l1AOffset = l1BufAddrStart;
|
||||
uint32_t l1BOffset = l1BufAddrStart + L1A_TILE_SIZE * L1_STAGES;
|
||||
|
||||
for (uint32_t i = 0; i < L1_STAGES; ++i) {
|
||||
l1ATensorList[i] = resource.l1Buf.template GetBufferByByte<ElementA>(l1AOffset + L1A_TILE_SIZE * i);
|
||||
l1BTensorList[i] = resource.l1Buf.template GetBufferByByte<ElementB>(l1BOffset + L1B_TILE_SIZE * i);
|
||||
l1AEventList[i] = i;
|
||||
l1BEventList[i] = i + L1_STAGES;
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[i]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[i]);
|
||||
}
|
||||
uint32_t l1SOffset = l1BOffset + L1B_TILE_SIZE * L1_STAGES;
|
||||
if constexpr (std::is_same_v<ElementA, int8_t>) {
|
||||
l1STensor = resource.l1Buf.template GetBufferByByte<uint64_t>(l1SOffset);
|
||||
AscendC::SetFlag<AscendC::HardEvent::FIX_MTE2>(0);
|
||||
}
|
||||
if (ptrSoftFlagBase_ != nullptr) {
|
||||
// Initialize the flag matrix (structure as below):
|
||||
// 1 0 0 0 0 0 0 0
|
||||
// 2 0 0 0 0 0 0 0
|
||||
// ...
|
||||
// 16 0 0 0 0 0 0 0
|
||||
// Then move it to L1
|
||||
uint32_t l1FOffset = l1SOffset + L1S_TILE_SIZE;
|
||||
l1FTensor = resource.l1Buf.template GetBufferByByte<int32_t>(l1FOffset);
|
||||
AscendC::GlobalTensor<int32_t> flagBase;
|
||||
flagBase.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(ptrSoftFlagBase_));
|
||||
AscendC::DataCopy(l1FTensor, flagBase, expertPerRank_ * FLAGSTRIDE);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void InitFpBuf(Arch::Resource<ArchTag> &resource, uint32_t FpAddrStart)
|
||||
{
|
||||
uint32_t FpOffset = FpAddrStart;
|
||||
fixpipeBuf = resource.fpBuf.template GetBufferByByte<uint64_t>(FpOffset);
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void InitL0A(Arch::Resource<ArchTag> &resource)
|
||||
{
|
||||
for (uint32_t i = 0; i < L0A_STAGES; ++i) {
|
||||
l0ATensorList[i] = resource.l0ABuf.template GetBufferByByte<ElementA>(L0A_TILE_SIZE * i);
|
||||
l0AEventList[i] = i;
|
||||
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0AEventList[i]);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void InitL0B(Arch::Resource<ArchTag> &resource)
|
||||
{
|
||||
for (uint32_t i = 0; i < L0B_STAGES; ++i) {
|
||||
l0BTensorList[i] = resource.l0BBuf.template GetBufferByByte<ElementB>(L0B_TILE_SIZE * i);
|
||||
l0BEventList[i] = i + L0A_STAGES;
|
||||
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0BEventList[i]);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void InitL0C(Arch::Resource<ArchTag> &resource)
|
||||
{
|
||||
for (uint32_t i = 0; i < L0C_STAGES; ++i) {
|
||||
l0CTensorList[i] = resource.l0CBuf.template GetBufferByByte<ElementAccumulator>(L0C_TILE_SIZE * i);
|
||||
l0CEventList[i] = i;
|
||||
AscendC::SetFlag<AscendC::HardEvent::FIX_M>(l0CEventList[i]);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void L1TileMmad(L1TileMmadParams const ¶ms)
|
||||
{
|
||||
uint32_t mPartLoop = CeilDiv<L0TileShape::M>(params.mRound);
|
||||
uint32_t nPartLoop = CeilDiv<L0TileShape::N>(params.nRound);
|
||||
uint32_t kPartLoop = CeilDiv<L0TileShape::K>(params.kActual);
|
||||
auto &l1ATensor = l1ATensorList[params.l1ListId];
|
||||
auto &l1BTensor = l1BTensorList[params.l1ListId];
|
||||
|
||||
auto &l0CTensor = l0CTensorList[l0CListId];
|
||||
LayoutCInL0 layoutCInL0 = LayoutCInL0::MakeLayoutInL0C(MakeCoord(params.mRound, params.nRound));
|
||||
|
||||
if constexpr (!ENABLE_UNIT_FLAG) {
|
||||
if (params.isKLoopFirst) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::FIX_M>(l0CEventList[l0CListId]);
|
||||
}
|
||||
}
|
||||
|
||||
for (uint32_t mPartIdx = 0; mPartIdx < mPartLoop; ++mPartIdx) {
|
||||
uint32_t mPartActual = (mPartIdx < mPartLoop - 1) ?
|
||||
L0TileShape::M : (params.mRound - mPartIdx * L0TileShape::M);
|
||||
|
||||
for (uint32_t kPartIdx = 0; kPartIdx < kPartLoop; ++kPartIdx) {
|
||||
uint32_t kPartActual = (kPartIdx < kPartLoop - 1) ?
|
||||
L0TileShape::K : (params.kActual - kPartIdx * L0TileShape::K);
|
||||
|
||||
auto &l0ATile = l0ATensorList[l0AListId];
|
||||
auto layoutAInL0 = LayoutAInL0::template MakeLayout<ElementA>(mPartActual, kPartActual);
|
||||
auto l1AOffset = MakeCoord(mPartIdx, kPartIdx) * L0TileShape::ToCoordMK();
|
||||
auto l1ATile = l1ATensor[L1A_LAYOUT.GetOffset(l1AOffset)];
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0AEventList[l0AListId]);
|
||||
if ((mPartIdx == 0) && (kPartIdx == 0)) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE1>(l1AEventList[params.l1ListId]);
|
||||
}
|
||||
copyL1ToL0A(l0ATile, l1ATile, layoutAInL0, L1A_LAYOUT);
|
||||
if ((mPartIdx == mPartLoop - 1) && (kPartIdx == kPartLoop - 1)) {
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[params.l1ListId]);
|
||||
}
|
||||
|
||||
for (uint32_t nPartIdx = 0; nPartIdx < nPartLoop; ++nPartIdx) {
|
||||
uint32_t nPartActual = (nPartIdx < nPartLoop - 1) ?
|
||||
L0TileShape::N : (params.nRound - nPartIdx * L0TileShape::N);
|
||||
|
||||
auto &l0BTile = l0BTensorList[l0BListId];
|
||||
auto layoutBInL0 = LayoutBInL0::template MakeLayout<ElementB>(kPartActual, nPartActual);
|
||||
auto l1BOffset = MakeCoord(kPartIdx, nPartIdx) * L0TileShape::ToCoordKN();
|
||||
auto l1BTile = l1BTensor[L1B_LAYOUT.GetOffset(l1BOffset)];
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0BEventList[l0BListId]);
|
||||
if ((kPartIdx == 0) && (nPartIdx == 0)) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE1>(l1BEventList[params.l1ListId]);
|
||||
}
|
||||
copyL1ToL0B(l0BTile, l1BTile, layoutBInL0, L1B_LAYOUT);
|
||||
if ((kPartIdx == kPartLoop - 1) && (nPartIdx == nPartLoop - 1)) {
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[params.l1ListId]);
|
||||
}
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE1_M>(EVENT_ID0);
|
||||
|
||||
auto l0COffset = MakeCoord(mPartIdx, nPartIdx) * L0TileShape::ToCoordMN();
|
||||
auto l0CTile = l0CTensor[layoutCInL0.GetOffset(l0COffset)];
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE1_M>(EVENT_ID0);
|
||||
// If the current tile is the first tile on the k axis, the accumulator needs to be reset to 0
|
||||
bool initC = (params.isKLoopFirst && (kPartIdx == 0));
|
||||
// If the unit flag is enabled, the unit flag is set according to the calculation progress
|
||||
uint8_t unitFlag = 0b00;
|
||||
if constexpr (ENABLE_UNIT_FLAG) {
|
||||
if (params.isKLoopLast &&
|
||||
(mPartIdx == mPartLoop - 1) && (kPartIdx == kPartLoop - 1) && (nPartIdx == nPartLoop - 1)) {
|
||||
unitFlag = 0b11;
|
||||
} else {
|
||||
unitFlag = 0b10;
|
||||
}
|
||||
}
|
||||
tileMmad(l0CTile, l0ATile, l0BTile, mPartActual, nPartActual, kPartActual, initC, unitFlag);
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0BEventList[l0BListId]);
|
||||
l0BListId = (l0BListId + 1 < L0B_STAGES) ? (l0BListId + 1) : 0;
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0AEventList[l0AListId]);
|
||||
l0AListId = (l0AListId + 1 < L0A_STAGES) ? (l0AListId + 1) : 0;
|
||||
}
|
||||
}
|
||||
|
||||
if (params.isKLoopLast) {
|
||||
auto layoutCInGm = params.layoutCInGm;
|
||||
if constexpr (std::is_same_v<ElementA, int8_t>) {
|
||||
auto layoutScale = params.layoutScale;
|
||||
auto layoutTileS = layoutScale.GetTileLayout(MakeCoord(layoutCInGm.shape(1)));
|
||||
AscendC::WaitFlag<AscendC::HardEvent::FIX_MTE2>(0);
|
||||
copyGmToL1S(l1STensor, params.gmBlockS, layoutTileS, layoutTileS);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_FIX>(0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_FIX>(0);
|
||||
|
||||
copyL1ToFP(fixpipeBuf, l1STensor, layoutTileS, layoutTileS);
|
||||
AscendC::SetFlag<AscendC::HardEvent::FIX_MTE2>(0);
|
||||
AscendC::PipeBarrier<PIPE_FIX>();
|
||||
}
|
||||
if constexpr (!ENABLE_UNIT_FLAG) {
|
||||
AscendC::SetFlag<AscendC::HardEvent::M_FIX>(l0CEventList[l0CListId]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::M_FIX>(l0CEventList[l0CListId]);
|
||||
if constexpr (std::is_same_v<ElementA, int8_t>) {
|
||||
copyL0CToGm(params.gmBlockC, l0CTensor, l1STensor, layoutCInGm, layoutCInL0);
|
||||
} else {
|
||||
copyL0CToGm(params.gmBlockC, l0CTensor, layoutCInGm, layoutCInL0);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::FIX_M>(l0CEventList[l0CListId]);
|
||||
} else {
|
||||
if constexpr (std::is_same_v<ElementA, int8_t>) {
|
||||
copyL0CToGm(params.gmBlockC, l0CTensor, l1STensor, layoutCInGm, layoutCInL0, 0b11);
|
||||
} else {
|
||||
copyL0CToGm(params.gmBlockC, l0CTensor, layoutCInGm, layoutCInL0, 0b11);
|
||||
}
|
||||
}
|
||||
l0CListId = (l0CListId + 1 < L0C_STAGES) ? (l0CListId + 1) : 0;
|
||||
if constexpr (std::is_same_v<ElementA, int8_t>) {
|
||||
AscendC::SetFlag<AscendC::HardEvent::FIX_MTE2>(0);
|
||||
}
|
||||
#ifdef __TILE_SYNC__
|
||||
if (params.flag > 0) {
|
||||
int32_t flagId = params.flag + params.syncLoopIdx / 8;
|
||||
AscendC::CrossCoreSetFlag<0x2, PIPE_FIX>(flagId);
|
||||
}
|
||||
#else
|
||||
Finalize(params.syncLoopIdx, params.flag);
|
||||
#endif
|
||||
}
|
||||
}
|
||||
|
||||
AscendC::LocalTensor<uint64_t> fixpipeBuf;
|
||||
|
||||
AscendC::LocalTensor<ElementA> l1ATensorList[L1_STAGES];
|
||||
AscendC::LocalTensor<ElementB> l1BTensorList[L1_STAGES];
|
||||
AscendC::LocalTensor<uint64_t> l1STensor;
|
||||
AscendC::LocalTensor<int32_t> l1FTensor;
|
||||
int32_t syncGroupIdx;
|
||||
int32_t l1AEventList[L1_STAGES];
|
||||
int32_t l1BEventList[L1_STAGES];
|
||||
uint32_t l1ListId{0};
|
||||
|
||||
AscendC::LocalTensor<ElementA> l0ATensorList[L0A_STAGES];
|
||||
int32_t l0AEventList[L0A_STAGES];
|
||||
uint32_t l0AListId{0};
|
||||
|
||||
AscendC::LocalTensor<ElementB> l0BTensorList[L0B_STAGES];
|
||||
int32_t l0BEventList[L0B_STAGES];
|
||||
uint32_t l0BListId{0};
|
||||
|
||||
AscendC::LocalTensor<ElementAccumulator> l0CTensorList[L0C_STAGES_];
|
||||
int32_t l0CEventList[L0C_STAGES_];
|
||||
uint32_t l0CListId{0};
|
||||
|
||||
L1TileMmadParams l1TileMmadParamsList[PRELOAD_STAGES];
|
||||
uint32_t l1TileMmadParamsId{0};
|
||||
uint32_t preloadCount{0};
|
||||
|
||||
TileMmad tileMmad;
|
||||
CopyGmToL1A copyGmToL1A;
|
||||
CopyGmToL1B copyGmToL1B;
|
||||
CopyGmToL1S copyGmToL1S;
|
||||
CopyL1ToL0A copyL1ToL0A;
|
||||
CopyL1ToL0B copyL1ToL0B;
|
||||
CopyL0CToGm copyL0CToGm;
|
||||
|
||||
__gm__ int32_t* ptrSoftFlagBase_ = nullptr;
|
||||
int32_t expertPerRank_;
|
||||
CopyL1ToFP copyL1ToFP;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Gemm::Block
|
||||
|
||||
#endif // CATLASS_GEMM_BLOCK_BLOCK_MMAD_PRELOAD_FIXPIPE_QUANT_HPP
|
||||
@@ -0,0 +1,13 @@
|
||||
|
||||
#ifndef CONST_ARGS_HPP
|
||||
#define CONST_ARGS_HPP
|
||||
constexpr static uint64_t MB_SIZE = 1024 * 1024UL;
|
||||
constexpr static int32_t NUMS_PER_FLAG = 16;
|
||||
constexpr static int32_t CACHE_LINE = 512;
|
||||
constexpr static int32_t FLAGSTRIDE = 16;
|
||||
constexpr static int32_t RESET_VAL = 0xffff;
|
||||
constexpr static int32_t ALIGN_128 = 128;
|
||||
constexpr uint32_t MAX_EXPERTS_PER_RANK = 32;
|
||||
constexpr static int32_t UB_ALIGN = 32;
|
||||
constexpr uint16_t CROSS_CORE_FLAG_MAX_SET_COUNT = 15;
|
||||
#endif
|
||||
@@ -0,0 +1,40 @@
|
||||
#ifndef COPY_GM_TO_L1_CUSTOM_HPP
|
||||
#define COPY_GM_TO_L1_CUSTOM_HPP
|
||||
|
||||
namespace Catlass::Gemm::Tile {
|
||||
/// Partial specialization for nZ in and nZ out.
|
||||
template <
|
||||
class ArchTag,
|
||||
class Element
|
||||
>
|
||||
struct CopyGmToL1<ArchTag, Gemm::GemmType<Element, layout::VectorLayout>> {
|
||||
using LayoutDst = layout::VectorLayout;
|
||||
using LayoutSrc = layout::VectorLayout;
|
||||
|
||||
static constexpr uint32_t ELE_NUM_PER_C0 = BYTE_PER_C0 / sizeof(Element); // int64, 32/8=4
|
||||
|
||||
// Methods
|
||||
|
||||
CATLASS_DEVICE
|
||||
CopyGmToL1() {};
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(
|
||||
AscendC::LocalTensor<Element> const &dstTensor,
|
||||
AscendC::GlobalTensor<Element> const &srcTensor,
|
||||
LayoutDst const &layoutDst, LayoutSrc const &layoutSrc)
|
||||
{
|
||||
uint32_t blockCount = 1;
|
||||
uint32_t blockLen = CeilDiv<ELE_NUM_PER_C0>(layoutSrc.shape(0));
|
||||
|
||||
AscendC::DataCopyParams repeatParams;
|
||||
|
||||
repeatParams.blockCount = blockCount;
|
||||
repeatParams.blockLen = blockLen;
|
||||
repeatParams.srcStride = 0;
|
||||
repeatParams.dstStride = 0;
|
||||
AscendC::DataCopy(dstTensor, srcTensor, repeatParams);
|
||||
}
|
||||
};
|
||||
}
|
||||
#endif // COPY_GM_TO_L1_CUSTOM_HPP
|
||||
@@ -0,0 +1,47 @@
|
||||
#ifndef COPY_L0C_TO_GM_CUSTOM_HPP
|
||||
#define COPY_L0C_TO_GM_CUSTOM_HPP
|
||||
|
||||
namespace Catlass::Gemm::Tile {
|
||||
template <
|
||||
class ElementAccumulator_,
|
||||
class ElementDst_,
|
||||
bool ReluEnable_
|
||||
>
|
||||
struct CopyL0CToGm<Catlass::Arch::AtlasA2,
|
||||
ElementAccumulator_,
|
||||
Gemm::GemmType<ElementDst_, layout::RowMajor>,
|
||||
ScaleGranularity::PER_CHANNEL,
|
||||
ReluEnable_>
|
||||
{
|
||||
using ArchTag = Catlass::Arch::AtlasA2;
|
||||
using ElementDst = ElementDst_;
|
||||
using ElementSrc = ElementAccumulator_;
|
||||
using LayoutSrc = Catlass::layout::zN;
|
||||
using LayoutDst = Catlass::layout::RowMajor;
|
||||
static constexpr auto quantPre = CopyL0CToGmQuantMode<ArchTag, ElementSrc, ElementDst,
|
||||
ScaleGranularity::PER_CHANNEL>::VALUE;
|
||||
static constexpr auto reluEn = ReluEnable_;
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(AscendC::GlobalTensor<ElementDst> const &dst, AscendC::LocalTensor<ElementSrc> const &src, AscendC::LocalTensor<uint64_t> cbufWorkspace,
|
||||
LayoutDst const &dstLayout, LayoutSrc const &srcLayout, uint8_t unitFlag = 0)
|
||||
{
|
||||
AscendC::FixpipeParamsV220 intriParams;
|
||||
|
||||
// Fixpipe layout information
|
||||
intriParams.nSize = dstLayout.shape(1);
|
||||
intriParams.mSize = dstLayout.shape(0);
|
||||
intriParams.srcStride = srcLayout.stride(3) / srcLayout.stride(0);
|
||||
intriParams.dstStride = dstLayout.stride(0);
|
||||
|
||||
// Fixpipe auxiliary arguments
|
||||
intriParams.quantPre = quantPre;
|
||||
intriParams.reluEn = reluEn;
|
||||
intriParams.unitFlag = unitFlag;
|
||||
|
||||
// Call AscendC Fixpipe
|
||||
AscendC::Fixpipe<ElementDst, ElementSrc, AscendC::CFG_ROW_MAJOR>(dst, src, cbufWorkspace, intriParams);
|
||||
}
|
||||
};
|
||||
}
|
||||
#endif // COPY_L0C_TO_GM_CUSTOM_HPP
|
||||
@@ -0,0 +1,53 @@
|
||||
#ifndef DISPATH_POLICY_CUSTOM_HPP
|
||||
#define DISPATH_POLICY_CUSTOM_HPP
|
||||
|
||||
namespace Catlass::Gemm {
|
||||
template <bool ENABLE_UNIT_FLAG_ = false, bool ENABLE_SHUFFLE_K_ = false>
|
||||
struct MmadAtlasA2PreloadFixpipeQuant : public MmadAtlasA2 {
|
||||
static constexpr uint32_t STAGES = 2;
|
||||
static constexpr bool ENABLE_UNIT_FLAG = ENABLE_UNIT_FLAG_;
|
||||
static constexpr bool ENABLE_SHUFFLE_K = ENABLE_SHUFFLE_K_;
|
||||
};
|
||||
|
||||
template <uint32_t PRELOAD_STAGES_, uint32_t L1_STAGES_, uint32_t L0A_STAGES_, uint32_t L0B_STAGES_,
|
||||
uint32_t L0C_STAGES_, bool ENABLE_UNIT_FLAG_, bool ENABLE_SHUFFLE_K_>
|
||||
struct MmadAtlasA2PreloadAsyncFixpipe :
|
||||
public MmadAtlasA2PreloadAsync<
|
||||
PRELOAD_STAGES_,
|
||||
L1_STAGES_,
|
||||
L0A_STAGES_,
|
||||
L0B_STAGES_,
|
||||
L0C_STAGES_,
|
||||
ENABLE_UNIT_FLAG_,
|
||||
ENABLE_SHUFFLE_K_
|
||||
> {
|
||||
};
|
||||
}
|
||||
|
||||
namespace Catlass::Epilogue {
|
||||
|
||||
template <uint32_t UB_STAGES_>
|
||||
struct EpilogueAtlasA2UnQuant {
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
};
|
||||
|
||||
template <uint32_t UB_STAGES_>
|
||||
struct EpilogueAtlasA2PerTokenDequantQuant {
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
};
|
||||
|
||||
template <uint32_t UB_STAGES_>
|
||||
struct EpilogueAtlasA2PerTokenDequantSwigluQuant {
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
};
|
||||
|
||||
template <uint32_t UB_STAGES_>
|
||||
struct EpilogueAtlasA2PerTokenDequantV2 {
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
};
|
||||
}
|
||||
#endif // DISPATH_POLICY_CUSTOM_HPP
|
||||
@@ -0,0 +1,16 @@
|
||||
#ifndef GET_TENSOR_ADDR_HPP
|
||||
#define GET_TENSOR_ADDR_HPP
|
||||
#include "kernel_operator.h"
|
||||
|
||||
#define FORCE_INLINE_AICORE inline __attribute__((always_inline)) __aicore__
|
||||
|
||||
template <typename T>
|
||||
FORCE_INLINE_AICORE __gm__ T* GetTensorAddr(uint32_t index, GM_ADDR tensorPtr) {
|
||||
__gm__ uint64_t* dataAddr = reinterpret_cast<__gm__ uint64_t*>(tensorPtr);
|
||||
uint64_t tensorPtrOffset = *dataAddr; // The offset of the data address from the first address.
|
||||
// Moving 3 bits to the right means dividing by sizeof(uint64 t).
|
||||
__gm__ uint64_t* retPtr = dataAddr + (tensorPtrOffset >> 3);
|
||||
return reinterpret_cast<__gm__ T*>(*(retPtr + index));
|
||||
}
|
||||
|
||||
#endif // GET_TENSOR_ADDR_HPP
|
||||
@@ -0,0 +1,327 @@
|
||||
#ifndef SYNC_UTIL_HPP
|
||||
#define SYNC_UTIL_HPP
|
||||
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "const_args.hpp"
|
||||
|
||||
#ifdef HCCL_COMM
|
||||
#include "moe_distribute_base.h"
|
||||
using namespace AscendC::HcclContextDef;
|
||||
|
||||
#else
|
||||
#include "shmem_api.h"
|
||||
#endif
|
||||
|
||||
#define FORCE_INLINE_AICORE inline __attribute__((always_inline)) __aicore__
|
||||
constexpr int32_t MAX_RANK_SIZE = 32;
|
||||
constexpr int32_t SHMEM_MEM = 700 * MB_SIZE;
|
||||
|
||||
constexpr uint16_t SEND_SYNC_EVENT_ID = 9;
|
||||
constexpr uint16_t RECV_SYNC_EVENT_ID = 10;
|
||||
|
||||
constexpr uint32_t SELF_STATE_OFFSET = 256 * 1024;
|
||||
constexpr uint32_t STATE_OFFSET = 512;
|
||||
|
||||
FORCE_INLINE_AICORE void AicSyncAll() {
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_FIX>(8);
|
||||
AscendC::CrossCoreWaitFlag<0x0>(8);
|
||||
}
|
||||
|
||||
template<typename T>
|
||||
FORCE_INLINE_AICORE void gm_store(__gm__ T *addr, T val) {
|
||||
*((__gm__ T *)addr) = val;
|
||||
}
|
||||
|
||||
template<typename T>
|
||||
FORCE_INLINE_AICORE T gm_load(__gm__ T *cache) {
|
||||
return *((__gm__ T *)cache);
|
||||
}
|
||||
|
||||
template<typename T>
|
||||
FORCE_INLINE_AICORE void gm_dcci(__gm__ T * addr) {
|
||||
using namespace AscendC;
|
||||
GlobalTensor<uint8_t> global;
|
||||
global.SetGlobalBuffer(reinterpret_cast<GM_ADDR>(addr));
|
||||
|
||||
// Important: add hint to avoid dcci being optimized by compiler
|
||||
__asm__ __volatile__("");
|
||||
DataCacheCleanAndInvalid<uint8_t, CacheLine::SINGLE_CACHE_LINE, DcciDst::CACHELINE_OUT>(global);
|
||||
__asm__ __volatile__("");
|
||||
}
|
||||
|
||||
FORCE_INLINE_AICORE int32_t gm_signal_wait_until_eq_for_barrier(__gm__ int32_t *sig_addr, int32_t cmp_val) {
|
||||
do {
|
||||
gm_dcci((__gm__ uint8_t *)sig_addr);
|
||||
if (*sig_addr == cmp_val) {
|
||||
return *sig_addr;
|
||||
}
|
||||
if (*sig_addr == cmp_val + 1) {
|
||||
return *sig_addr;
|
||||
}
|
||||
} while (true);
|
||||
return -1;
|
||||
}
|
||||
|
||||
FORCE_INLINE_AICORE void gm_signal_wait_until_ne(__gm__ int32_t *sig_addr, int32_t cmp_val) {
|
||||
do {
|
||||
AscendC::LocalTensor<int32_t> ub;
|
||||
ub.address_.logicPos = static_cast<uint8_t>(AscendC::TPosition::VECIN);
|
||||
ub.address_.bufferAddr = 0;
|
||||
AscendC::GlobalTensor<int32_t> sig;
|
||||
sig.SetGlobalBuffer(sig_addr);
|
||||
AscendC::DataCopy(ub, sig, 8);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_S>(EVENT_ID0);
|
||||
if (ub(0) != cmp_val) {
|
||||
return;
|
||||
}
|
||||
} while (true);
|
||||
return;
|
||||
}
|
||||
|
||||
|
||||
class HcclShmem {
|
||||
public:
|
||||
#ifdef HCCL_COMM // HCCL needs to initialize the HCCL context
|
||||
__gm__ HcclOpResParamCustom *WinContext_{nullptr};
|
||||
Hccl<HCCL_SERVER_TYPE_AICPU> hccl_;
|
||||
AscendC::LocalTensor<int32_t> ub;
|
||||
FORCE_INLINE_AICORE
|
||||
HcclShmem(){
|
||||
auto contextGM0 = AscendC::GetHcclContext<HCCL_GROUP_ID_0>();
|
||||
WinContext_ = (__gm__ HcclOpResParamCustom *)contextGM0;
|
||||
|
||||
m_rank = WinContext_->localUsrRankId;
|
||||
m_rankSize = WinContext_->rankSize;
|
||||
m_segmentSize = WinContext_->winSize;
|
||||
}
|
||||
#else
|
||||
FORCE_INLINE_AICORE
|
||||
HcclShmem(){
|
||||
m_segmentSize = SHMEM_MEM;
|
||||
}
|
||||
FORCE_INLINE_AICORE
|
||||
void initShmem(GM_ADDR symmetricPtr_, size_t rank, size_t rankSize) {
|
||||
symmetricPtr = symmetricPtr_;
|
||||
m_rank = rank;
|
||||
m_rankSize = rankSize;
|
||||
}
|
||||
#endif
|
||||
|
||||
FORCE_INLINE_AICORE
|
||||
GM_ADDR operator() () const { // No parameters: return pointer to local peermem
|
||||
#ifdef HCCL_COMM
|
||||
return (GM_ADDR)(WinContext_->localWindowsIn);
|
||||
#else
|
||||
return reinterpret_cast<GM_ADDR>(shmem_ptr(symmetricPtr, m_rank));
|
||||
#endif
|
||||
}
|
||||
|
||||
FORCE_INLINE_AICORE
|
||||
GM_ADDR operator() (int32_t index) const { // With index parameter: return pointer to the base address of remote peermem
|
||||
#ifdef HCCL_COMM
|
||||
return (GM_ADDR)((index == m_rank) ? WinContext_->localWindowsIn :
|
||||
((HcclRankRelationResV2Custom *)(WinContext_->remoteRes[index].nextDevicePtr))->windowsIn);
|
||||
#else
|
||||
return reinterpret_cast<GM_ADDR>(shmem_ptr(symmetricPtr, index));
|
||||
#endif
|
||||
}
|
||||
|
||||
FORCE_INLINE_AICORE
|
||||
GM_ADDR operator () (int64_t offset, int32_t rankId) const {
|
||||
#ifdef HCCL_COMM
|
||||
if (offset < 0 || offset >= m_segmentSize) {
|
||||
return nullptr;
|
||||
}
|
||||
if (rankId < 0 || rankId >= m_rankSize) {
|
||||
return nullptr;
|
||||
}
|
||||
return (GM_ADDR)((rankId == m_rank) ? WinContext_->localWindowsIn :
|
||||
((HcclRankRelationResV2Custom *)(WinContext_->remoteRes[rankId].nextDevicePtr))->windowsIn) + offset;
|
||||
#else
|
||||
return reinterpret_cast<GM_ADDR>(shmem_ptr((symmetricPtr + offset), rankId));
|
||||
#endif
|
||||
}
|
||||
|
||||
|
||||
|
||||
FORCE_INLINE_AICORE
|
||||
size_t SegmentSize() const {
|
||||
return m_segmentSize;
|
||||
}
|
||||
|
||||
FORCE_INLINE_AICORE
|
||||
int32_t RankSize() const {
|
||||
return m_rankSize;
|
||||
}
|
||||
|
||||
|
||||
FORCE_INLINE_AICORE
|
||||
~HcclShmem() {
|
||||
}
|
||||
|
||||
|
||||
FORCE_INLINE_AICORE
|
||||
void CrossRankSync() {
|
||||
uint64_t flag_offset = (m_segmentSize - MB_SIZE) / sizeof(int32_t);
|
||||
__gm__ int32_t* sync_counter = (__gm__ int32_t*)(*this)() + flag_offset;
|
||||
__gm__ int32_t* sync_base = (__gm__ int32_t*)(*this)() + flag_offset + 2048;
|
||||
int count = gm_load(sync_base) + 1;
|
||||
int vec_id = AscendC::GetBlockIdx();
|
||||
int vec_size = AscendC::GetBlockNum() * AscendC::GetTaskRation();
|
||||
for(int i = vec_id; i < m_rankSize; i += vec_size) {
|
||||
__gm__ int32_t* sync_remote = (__gm__ int32_t*)((*this)(i)) + flag_offset + m_rank * 16;
|
||||
gm_store(sync_remote, count);
|
||||
gm_dcci((__gm__ uint8_t*)sync_remote);
|
||||
auto sync_check = sync_counter + i * 16;
|
||||
gm_signal_wait_until_eq_for_barrier(sync_check, count);
|
||||
}
|
||||
|
||||
AscendC::SyncAll<true>();
|
||||
gm_store(sync_base, count);
|
||||
}
|
||||
|
||||
|
||||
FORCE_INLINE_AICORE
|
||||
void InitStatusTargetSum()
|
||||
{
|
||||
using namespace AscendC;
|
||||
uint64_t flag_offset = (m_segmentSize - MB_SIZE) + SELF_STATE_OFFSET;
|
||||
//uint64_t self_state_offset = (m_segmentSize - 2 * MB_SIZE);
|
||||
// ep state
|
||||
//uint32_t coreIdx = get_block_idx();;
|
||||
uint32_t coreIdx = GetBlockIdx();
|
||||
GlobalTensor<int32_t> selfStatusTensor;
|
||||
selfStatusTensor.SetGlobalBuffer((__gm__ int32_t *)((*this)() + flag_offset));
|
||||
__asm__ __volatile__("");
|
||||
DataCacheCleanAndInvalid<int32_t, CacheLine::SINGLE_CACHE_LINE, DcciDst::CACHELINE_OUT>(selfStatusTensor[coreIdx * UB_ALIGN]);
|
||||
__asm__ __volatile__("");
|
||||
int32_t state = selfStatusTensor(coreIdx * UB_ALIGN);
|
||||
if (state == 0) {
|
||||
sumTarget_ = static_cast<float>(1.0);
|
||||
selfStatusTensor(coreIdx * UB_ALIGN) = 0x3F800000; // 1.0f
|
||||
epStateValue_ = 0x3F800000; // 1.0f
|
||||
} else {
|
||||
sumTarget_ = static_cast<float>(0.0);
|
||||
selfStatusTensor(coreIdx * UB_ALIGN) = 0;
|
||||
epStateValue_ = 0;
|
||||
}
|
||||
__asm__ __volatile__("");
|
||||
DataCacheCleanAndInvalid<int32_t, CacheLine::SINGLE_CACHE_LINE, DcciDst::CACHELINE_OUT>(selfStatusTensor[coreIdx * UB_ALIGN]);
|
||||
__asm__ __volatile__("");
|
||||
}
|
||||
|
||||
FORCE_INLINE_AICORE
|
||||
void CrossRankSyncV2Set(AscendC::LocalTensor<int32_t> ctrBuffer) {
|
||||
//subblockid = 0
|
||||
uint32_t stateOffset_ = STATE_OFFSET;
|
||||
// uint32_t epStateOffsetOnWin_ = m_rank * stateOffset_;
|
||||
|
||||
uint64_t flag_offset = (m_segmentSize - MB_SIZE) + m_rank * stateOffset_;
|
||||
//uint64_t flag_offset = (m_segmentSize - MB_SIZE);
|
||||
int vec_size = get_block_num();
|
||||
int vec_id = get_block_idx();
|
||||
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_MTE3>(RECV_SYNC_EVENT_ID);
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_MTE3>(SEND_SYNC_EVENT_ID);
|
||||
AscendC::CrossCoreWaitFlag(SEND_SYNC_EVENT_ID);
|
||||
pipe_barrier(PIPE_ALL);
|
||||
|
||||
ctrBuffer.SetValue(0, epStateValue_);
|
||||
AscendC::SetFlag<AscendC::HardEvent::S_MTE3>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::S_MTE3>(EVENT_ID0);
|
||||
for (uint32_t dstEpIdx = vec_id; dstEpIdx < m_rankSize; dstEpIdx += vec_size) {
|
||||
AscendC::GlobalTensor<int32_t> gmDstStates;
|
||||
gmDstStates.SetGlobalBuffer((__gm__ int32_t*)((*this)(flag_offset, dstEpIdx)));
|
||||
DataCopy(gmDstStates, ctrBuffer, 8);
|
||||
}
|
||||
AscendC::CrossCoreWaitFlag(RECV_SYNC_EVENT_ID);
|
||||
}
|
||||
|
||||
FORCE_INLINE_AICORE
|
||||
void CrossRankSyncV2Wait(AscendC::LocalTensor<float> statusTensor, AscendC::LocalTensor<float> gatherMaskOutTensor,
|
||||
AscendC::LocalTensor<uint32_t> gatherTmpTensor, AscendC::LocalTensor<float> statusSumOutTensor) {
|
||||
|
||||
uint64_t flag_offset = (m_segmentSize - MB_SIZE);
|
||||
int vec_size = get_block_num();
|
||||
int vec_id = get_block_idx();
|
||||
uint32_t stateOffset_ = STATE_OFFSET;
|
||||
|
||||
uint32_t sendRankNum_ = m_rankSize / vec_size;
|
||||
uint32_t remainderRankNum = m_rankSize % vec_size;
|
||||
uint32_t startRankId_ = sendRankNum_ * vec_id;
|
||||
if (vec_id < remainderRankNum) {
|
||||
sendRankNum_++;
|
||||
startRankId_ += vec_id;
|
||||
} else {
|
||||
startRankId_ += remainderRankNum;
|
||||
}
|
||||
uint32_t endRankId_ = startRankId_ + sendRankNum_;
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_MTE3>(SEND_SYNC_EVENT_ID);
|
||||
|
||||
AscendC::GlobalTensor<float> epStatusSpaceGlobalTensor_;
|
||||
epStatusSpaceGlobalTensor_.SetGlobalBuffer((__gm__ float *)((*this)() + flag_offset));
|
||||
|
||||
if (startRankId_ < m_rankSize) {
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
gatherTmpTensor.SetValue(0, 1);
|
||||
uint32_t mask = 1; // gatherMask + sum
|
||||
uint64_t rsvdCnt = 0;
|
||||
// DataCopyParams intriParams{static_cast<uint16_t>(sendRankNum_), 1,
|
||||
// static_cast<uint16_t>((moeSendNum_ > 512) ? 7 : 15), 0};
|
||||
AscendC::DataCopyParams intriParams{static_cast<uint16_t>(sendRankNum_), 1,
|
||||
static_cast<uint16_t>(15), 0};
|
||||
|
||||
float sumOfFlag = static_cast<float>(-1.0);
|
||||
float minTarget = (sumTarget_ * sendRankNum_) - (float)0.5;
|
||||
float maxTarget = (sumTarget_ * sendRankNum_) + (float)0.5;
|
||||
AscendC::SumParams sumParams{1, sendRankNum_, sendRankNum_};
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::S_V>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::S_V>(EVENT_ID0);
|
||||
|
||||
while ((sumOfFlag < minTarget) || (sumOfFlag > maxTarget)) {
|
||||
AscendC::DataCopy<float>(statusTensor, epStatusSpaceGlobalTensor_[startRankId_ * stateOffset_ / sizeof(float)],
|
||||
intriParams);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0);
|
||||
|
||||
GatherMask(gatherMaskOutTensor, statusTensor, gatherTmpTensor, true, mask,
|
||||
{1, (uint16_t)sendRankNum_, 1, 0}, rsvdCnt);
|
||||
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Sum(statusSumOutTensor, gatherMaskOutTensor, sumParams);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_S>(EVENT_ID0);
|
||||
sumOfFlag = statusSumOutTensor.GetValue(0);
|
||||
}
|
||||
}
|
||||
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_MTE3>(RECV_SYNC_EVENT_ID);
|
||||
AscendC::CrossCoreWaitFlag(RECV_SYNC_EVENT_ID);
|
||||
|
||||
//unpermute
|
||||
AscendC::CrossCoreWaitFlag(SEND_SYNC_EVENT_ID);
|
||||
}
|
||||
|
||||
|
||||
FORCE_INLINE_AICORE
|
||||
__gm__ int32_t* SyncBaseAddr() {
|
||||
uint64_t flag_offset = (m_segmentSize - MB_SIZE) / sizeof(int32_t);
|
||||
return (__gm__ int32_t*)(*this)() + flag_offset + 2048;
|
||||
}
|
||||
|
||||
private:
|
||||
GM_ADDR symmetricPtr;
|
||||
int32_t m_rank;
|
||||
int32_t m_rankSize;
|
||||
size_t m_segmentSize;
|
||||
float sumTarget_{0.0};
|
||||
int32_t epStateValue_;
|
||||
};
|
||||
|
||||
|
||||
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,20 @@
|
||||
#ifndef LAYOUT_3D_HPP
|
||||
#define LAYOUT_3D_HPP
|
||||
#include "kernel_operator.h"
|
||||
#include "catlass/catlass.hpp"
|
||||
class Layout3D {
|
||||
int64_t strides[2];
|
||||
public:
|
||||
CATLASS_DEVICE
|
||||
Layout3D() {}
|
||||
CATLASS_DEVICE
|
||||
Layout3D(int64_t stride0, int64_t stride1) {
|
||||
strides[0] = stride0;
|
||||
strides[1] = stride1;
|
||||
}
|
||||
CATLASS_DEVICE
|
||||
int64_t operator() (int64_t dim0, int64_t dim1, int64_t dim2) {
|
||||
return dim0 * strides[0] + dim1 * strides[1] + dim2;
|
||||
}
|
||||
};
|
||||
#endif // LAYOUT_3D_HPP
|
||||
365
csrc/mc2/dispatch_ffn_combine_bf16/op_kernel/utils/moe_distribute_base.h
Executable file
365
csrc/mc2/dispatch_ffn_combine_bf16/op_kernel/utils/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 HcclRankRelationResV2Custom {
|
||||
uint32_t remoteUsrRankId;
|
||||
uint32_t remoteWorldRank;
|
||||
uint64_t windowsIn;
|
||||
uint64_t windowsOut;
|
||||
uint64_t windowsExp;
|
||||
ListCommon nextTagRes;
|
||||
};
|
||||
|
||||
struct HcclOpResParamCustom {
|
||||
// 本地资源
|
||||
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
|
||||
@@ -0,0 +1,25 @@
|
||||
#ifndef SELECT_HELPER_HPP
|
||||
#define SELECT_HELPER_HPP
|
||||
|
||||
#include "catlass/layout/layout.hpp"
|
||||
using namespace AscendC;
|
||||
using namespace Catlass;
|
||||
|
||||
template <typename Layout, typename ElementType, typename = void>
|
||||
struct LayoutBInitializer {
|
||||
CATLASS_DEVICE
|
||||
static Layout create(uint32_t k, uint32_t n) {
|
||||
return Layout{k, n};
|
||||
}
|
||||
};
|
||||
|
||||
template <typename Layout, typename ElementType>
|
||||
struct LayoutBInitializer<Layout, ElementType,
|
||||
std::enable_if_t<std::is_same_v<Layout, layout::zN>>
|
||||
> {
|
||||
CATLASS_DEVICE
|
||||
static Layout create(uint32_t k, uint32_t n) {
|
||||
return Layout::template MakeLayout<ElementType>(k, n);
|
||||
}
|
||||
};
|
||||
#endif // SELECT_HELPER_HPP
|
||||
Reference in New Issue
Block a user